From 2f03b1a3220378426ba1c0894f1551829f4c60e5 Mon Sep 17 00:00:00 2001 From: pakrym-oai Date: Thu, 12 Mar 2026 09:00:20 -0700 Subject: [PATCH] Dispatch tools when code mode is not awaited directly (#14437) ## Summary - start a code mode worker once per turn and let it pump nested tool calls through a dedicated queue - simplify code mode request/response dispatch around request ids and generic runner-unavailable errors - clean up the code mode process API and runner protocol plumbing ## Testing - not run yet --- codex-rs/core/src/codex.rs | 5 + codex-rs/core/src/tools/code_mode.rs | 543 +++++++++++-------- codex-rs/core/src/tools/code_mode_runner.cjs | 65 ++- codex-rs/core/tests/suite/code_mode.rs | 129 ++++- 4 files changed, 495 insertions(+), 247 deletions(-) diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index ef6d6b8be..4a614ad5f 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -5545,6 +5545,11 @@ pub(crate) async fn run_turn( // Although from the perspective of codex.rs, TurnDiffTracker has the lifecycle of a Task which contains // many turns, from the perspective of the user, it is a single turn. let turn_diff_tracker = Arc::new(tokio::sync::Mutex::new(TurnDiffTracker::new())); + let _code_mode_worker = sess + .services + .code_mode_service + .start_turn_worker(&sess, &turn_context, &turn_diff_tracker) + .await; let mut server_model_warning_emitted_for_turn = false; // `ModelClientSession` is turn-scoped and caches WebSocket + sticky routing state, so we reuse diff --git a/codex-rs/core/src/tools/code_mode.rs b/codex-rs/core/src/tools/code_mode.rs index a6e6227be..0f8ac1a8e 100644 --- a/codex-rs/core/src/tools/code_mode.rs +++ b/codex-rs/core/src/tools/code_mode.rs @@ -1,5 +1,4 @@ use std::collections::HashMap; -use std::collections::VecDeque; use std::path::PathBuf; use std::sync::Arc; use std::time::Duration; @@ -33,6 +32,8 @@ use tokio::io::AsyncReadExt; use tokio::io::AsyncWriteExt; use tokio::io::BufReader; use tokio::sync::Mutex; +use tokio::sync::mpsc; +use tokio::sync::oneshot; use tokio::task::JoinHandle; use tracing::warn; @@ -51,79 +52,108 @@ struct ExecContext { pub(crate) struct CodeModeProcess { child: tokio::process::Child, - stdin: tokio::process::ChildStdin, - stdout_lines: tokio::io::Lines>, - stderr_task: Option>, - pending_messages: HashMap>, + stdin: Arc>, + stdout_task: JoinHandle<()>, + // A set of current requests waiting for a response from code mode host + response_waiters: Arc>>>, + // When there is an active worker it listens for tool calls from code mode and processes them + tool_call_rx: Arc>>, +} + +pub(crate) struct CodeModeWorker { + shutdown_tx: Option>, +} + +#[derive(Debug, Deserialize)] +#[serde(rename_all = "snake_case")] +struct CodeModeToolCall { + request_id: String, + id: String, + name: String, + #[serde(default)] + input: Option, +} + +impl Drop for CodeModeWorker { + fn drop(&mut self) { + if let Some(shutdown_tx) = self.shutdown_tx.take() { + let _ = shutdown_tx.send(()); + } + } } impl CodeModeProcess { - async fn write(&mut self, message: &HostToNodeMessage) -> Result<(), std::io::Error> { - let line = serde_json::to_string(message).map_err(std::io::Error::other)?; - self.stdin.write_all(line.as_bytes()).await?; - self.stdin.write_all(b"\n").await?; - self.stdin.flush().await?; - Ok(()) - } - - async fn read(&mut self, session_id: i32) -> Result { - if let Some(message) = self - .pending_messages - .get_mut(&session_id) - .and_then(VecDeque::pop_front) - { - return Ok(message); - } - - loop { - let Some(line) = self.stdout_lines.next_line().await? else { - match self.wait_for_exit().await { - Ok(status) => { - self.join_stderr_task().await; - return Err(std::io::Error::other(format!( - "{PUBLIC_TOOL_NAME} runner exited without returning a result (status {status})" - ))); + fn worker(&self, exec: ExecContext) -> CodeModeWorker { + let (shutdown_tx, mut shutdown_rx) = oneshot::channel(); + let stdin = Arc::clone(&self.stdin); + let tool_call_rx = Arc::clone(&self.tool_call_rx); + tokio::spawn(async move { + loop { + let tool_call = tokio::select! { + _ = &mut shutdown_rx => break, + tool_call = async { + let mut tool_call_rx = tool_call_rx.lock().await; + tool_call_rx.recv().await + } => tool_call, + }; + let Some(tool_call) = tool_call else { + break; + }; + let exec = exec.clone(); + let stdin = Arc::clone(&stdin); + tokio::spawn(async move { + let response = HostToNodeMessage::Response { + request_id: tool_call.request_id, + id: tool_call.id, + code_mode_result: call_nested_tool(exec, tool_call.name, tool_call.input) + .await, + }; + if let Err(err) = write_message(&stdin, &response).await { + warn!("failed to write {PUBLIC_TOOL_NAME} tool response: {err}"); } - Err(err) => return Err(std::io::Error::other(err)), - } - }; - if line.trim().is_empty() { - continue; + }); } - let message: NodeToHostMessage = - serde_json::from_str(&line).map_err(std::io::Error::other)?; - let message_session_id = message_session_id(&message); - if message_session_id == session_id { - return Ok(message); - } - self.pending_messages - .entry(message_session_id) - .or_default() - .push_back(message); + }); + + CodeModeWorker { + shutdown_tx: Some(shutdown_tx), } } - fn has_exited(&mut self) -> Result { + async fn send( + &mut self, + request_id: &str, + message: &HostToNodeMessage, + ) -> Result { + if self.stdout_task.is_finished() { + return Err(std::io::Error::other(format!( + "{PUBLIC_TOOL_NAME} runner is not available" + ))); + } + + let (tx, rx) = oneshot::channel(); + self.response_waiters + .lock() + .await + .insert(request_id.to_string(), tx); + if let Err(err) = write_message(&self.stdin, message).await { + self.response_waiters.lock().await.remove(request_id); + return Err(err); + } + + match rx.await { + Ok(message) => Ok(message), + Err(_) => Err(std::io::Error::other(format!( + "{PUBLIC_TOOL_NAME} runner is not available" + ))), + } + } + + 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 join_stderr_task(&mut self) { - let Some(stderr_task) = self.stderr_task.take() else { - return; - }; - if let Err(err) = stderr_task.await { - warn!("failed to join {PUBLIC_TOOL_NAME} stderr task: {err}"); - } + .map_err(std::io::Error::other) } } @@ -154,26 +184,62 @@ impl CodeModeService { async fn ensure_started( &self, - ) -> Result>, String> { + ) -> Result>, std::io::Error> { let mut process_slot = self.process.lock().await; let needs_spawn = match process_slot.as_mut() { Some(process) => !matches!(process.has_exited(), Ok(false)), None => true, }; if needs_spawn { - let node_path = resolve_compatible_node(self.js_repl_node_path.as_deref()).await?; + let node_path = resolve_compatible_node(self.js_repl_node_path.as_deref()) + .await + .map_err(std::io::Error::other)?; *process_slot = Some(spawn_code_mode_process(&node_path).await?); } drop(process_slot); Ok(self.process.clone().lock_owned().await) } + pub(crate) async fn start_turn_worker( + &self, + session: &Arc, + turn: &Arc, + tracker: &SharedTurnDiffTracker, + ) -> Option { + if !turn.features.enabled(Feature::CodeMode) { + return None; + } + let exec = ExecContext { + session: Arc::clone(session), + turn: Arc::clone(turn), + tracker: Arc::clone(tracker), + }; + let mut process_slot = match self.ensure_started().await { + Ok(process_slot) => process_slot, + Err(err) => { + warn!("failed to start {PUBLIC_TOOL_NAME} worker for turn: {err}"); + return None; + } + }; + let Some(process) = process_slot.as_mut() else { + warn!( + "failed to start {PUBLIC_TOOL_NAME} worker for turn: {PUBLIC_TOOL_NAME} runner failed to start" + ); + return None; + }; + Some(process.worker(exec)) + } + pub(crate) async fn allocate_session_id(&self) -> i32 { let mut next_session_id = self.next_session_id.lock().await; let session_id = *next_session_id; *next_session_id = next_session_id.saturating_add(1); session_id } + + pub(crate) async fn allocate_request_id(&self) -> String { + uuid::Uuid::new_v4().to_string() + } } #[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] @@ -198,20 +264,23 @@ struct EnabledTool { #[serde(tag = "type", rename_all = "snake_case")] enum HostToNodeMessage { Start { + request_id: String, session_id: i32, enabled_tools: Vec, stored_values: HashMap, source: String, }, Poll { + request_id: String, session_id: i32, yield_time_ms: u64, }, Terminate { + request_id: String, session_id: i32, }, Response { - session_id: i32, + request_id: String, id: String, code_mode_result: JsonValue, }, @@ -221,22 +290,19 @@ enum HostToNodeMessage { #[serde(tag = "type", rename_all = "snake_case")] enum NodeToHostMessage { ToolCall { - session_id: i32, - id: String, - name: String, - #[serde(default)] - input: Option, + #[serde(flatten)] + tool_call: CodeModeToolCall, }, Yielded { - session_id: i32, + request_id: String, content_items: Vec, }, Terminated { - session_id: i32, + request_id: String, content_items: Vec, }, Result { - session_id: i32, + request_id: String, content_items: Vec, stored_values: HashMap, #[serde(default)] @@ -278,7 +344,7 @@ pub(crate) fn instructions(config: &Config) -> Option { )); section.push_str("- Import nested tools from `tools.js`, for example `import { exec_command } from \"tools.js\"` or `import { ALL_TOOLS } from \"tools.js\"` to inspect the available `{ module, name, description }` entries. Namespaced tools are also available from `tools/.js`; MCP tools use `tools/mcp/.js`, for example `import { append_notebook_logs_chart } from \"tools/mcp/ologs.js\"`. Nested tool calls resolve to their code-mode result values.\n"); section.push_str(&format!( - "- Import `{{ output_text, output_image, set_max_output_tokens_per_exec_call, set_yield_time, store, load }}` from `@openai/code_mode` (or `\"openai/code_mode\"`). `output_text(value)` surfaces text back to the model and stringifies non-string objects with `JSON.stringify(...)` when possible. `output_image(imageUrl)` appends an `input_image` content item for `http(s)` or `data:` URLs. `store(key, value)` persists JSON-serializable values across `{PUBLIC_TOOL_NAME}` calls in the current session, and `load(key)` returns a cloned stored value or `undefined`. `set_max_output_tokens_per_exec_call(value)` sets the token budget used to truncate direct `{PUBLIC_TOOL_NAME}` returns; `{WAIT_TOOL_NAME}` uses its own `max_tokens` argument instead and defaults to `10000`. `set_yield_time(value)` asks `{PUBLIC_TOOL_NAME}` to return early if the script is still running after that many milliseconds so `{WAIT_TOOL_NAME}` can resume it later. The returned content starts with a separate `Script completed`, `Script failed`, or `Script running with session ID …` text item that includes wall time. When truncation happens, the final text may include `Total output lines:` and the usual `…N tokens truncated…` marker.\n", + "- Import `{{ background, output_text, output_image, set_max_output_tokens_per_exec_call, set_yield_time, store, load }}` from `@openai/code_mode` (or `\"openai/code_mode\"`). `output_text(value)` surfaces text back to the model and stringifies non-string objects with `JSON.stringify(...)` when possible. `output_image(imageUrl)` appends an `input_image` content item for `http(s)` or `data:` URLs. `store(key, value)` persists JSON-serializable values across `{PUBLIC_TOOL_NAME}` calls in the current session, and `load(key)` returns a cloned stored value or `undefined`. `set_max_output_tokens_per_exec_call(value)` sets the token budget used to truncate direct `{PUBLIC_TOOL_NAME}` returns; `{WAIT_TOOL_NAME}` uses its own `max_tokens` argument instead and defaults to `10000`. `set_yield_time(value)` asks `{PUBLIC_TOOL_NAME}` to return early if the script is still running after that many milliseconds so `{WAIT_TOOL_NAME}` can resume it later. `background()` returns a yielded `{PUBLIC_TOOL_NAME}` response immediately while the script keeps running in the background. The returned content starts with a separate `Script completed`, `Script failed`, or `Script running with session ID …` text item that includes wall time. When truncation happens, the final text may include `Total output lines:` and the usual `…N tokens truncated…` marker.\n", )); section.push_str(&format!( "- If `{PUBLIC_TOOL_NAME}` returns `Script running with session ID …`, call `{WAIT_TOOL_NAME}` with that `session_id` to keep waiting for more output, completion, or termination.\n", @@ -308,10 +374,19 @@ pub(crate) async fn execute( let stored_values = service.stored_values().await; let source = build_source(&code, &enabled_tools).map_err(FunctionCallError::RespondToModel)?; let session_id = service.allocate_session_id().await; + let request_id = service.allocate_request_id().await; let process_slot = service .ensure_started() .await - .map_err(FunctionCallError::RespondToModel)?; + .map_err(|err| FunctionCallError::RespondToModel(err.to_string()))?; + let started_at = std::time::Instant::now(); + let message = HostToNodeMessage::Start { + request_id: request_id.clone(), + session_id, + enabled_tools, + stored_values, + source, + }; let result = { let mut process_slot = process_slot; let Some(process) = process_slot.as_mut() else { @@ -319,19 +394,15 @@ pub(crate) async fn execute( "{PUBLIC_TOOL_NAME} runner failed to start" ))); }; - drive_code_mode_session( - &exec, - process, - HostToNodeMessage::Start { - session_id, - enabled_tools, - stored_values, - source, - }, - None, - false, - ) - .await + let message = process + .send(&request_id, &message) + .await + .map_err(|err| err.to_string()); + let message = match message { + Ok(message) => message, + Err(error) => return Err(FunctionCallError::RespondToModel(error)), + }; + handle_node_message(&exec, session_id, message, None, started_at).await }; match result { Ok(CodeModeSessionProgress::Finished(output)) @@ -354,13 +425,32 @@ pub(crate) async fn wait( turn, tracker, }; + let request_id = exec + .session + .services + .code_mode_service + .allocate_request_id() + .await; + let started_at = std::time::Instant::now(); + let message = if terminate { + HostToNodeMessage::Terminate { + request_id: request_id.clone(), + session_id, + } + } else { + HostToNodeMessage::Poll { + request_id: request_id.clone(), + session_id, + yield_time_ms, + } + }; let process_slot = exec .session .services .code_mode_service .ensure_started() .await - .map_err(FunctionCallError::RespondToModel)?; + .map_err(|err| FunctionCallError::RespondToModel(err.to_string()))?; let result = { let mut process_slot = process_slot; let Some(process) = process_slot.as_mut() else { @@ -373,19 +463,20 @@ pub(crate) async fn wait( "{PUBLIC_TOOL_NAME} runner failed to start" ))); } - drive_code_mode_session( + let message = process + .send(&request_id, &message) + .await + .map_err(|err| err.to_string()); + let message = match message { + Ok(message) => message, + Err(error) => return Err(FunctionCallError::RespondToModel(error)), + }; + handle_node_message( &exec, - process, - if terminate { - HostToNodeMessage::Terminate { session_id } - } else { - HostToNodeMessage::Poll { - session_id, - yield_time_ms, - } - }, + session_id, + message, Some(max_output_tokens), - terminate, + started_at, ) .await }; @@ -396,131 +487,18 @@ pub(crate) async fn wait( } } -async fn spawn_code_mode_process(node_path: &std::path::Path) -> Result { - let mut cmd = tokio::process::Command::new(node_path); - cmd.arg("--experimental-vm-modules"); - cmd.arg("--eval"); - cmd.arg(CODE_MODE_RUNNER_SOURCE); - cmd.stdin(std::process::Stdio::piped()) - .stdout(std::process::Stdio::piped()) - .stderr(std::process::Stdio::piped()) - .kill_on_drop(true); - - let mut child = cmd - .spawn() - .map_err(|err| format!("failed to start {PUBLIC_TOOL_NAME} Node runtime: {err}"))?; - let stdout = child - .stdout - .take() - .ok_or_else(|| format!("{PUBLIC_TOOL_NAME} runner missing stdout"))?; - let stderr = child - .stderr - .take() - .ok_or_else(|| format!("{PUBLIC_TOOL_NAME} runner missing stderr"))?; - let stdin = child - .stdin - .take() - .ok_or_else(|| format!("{PUBLIC_TOOL_NAME} runner missing stdin"))?; - - let stderr_task = tokio::spawn(async move { - let mut reader = BufReader::new(stderr); - let mut buf = Vec::new(); - match reader.read_to_end(&mut buf).await { - Ok(_) => { - let stderr = String::from_utf8_lossy(&buf).trim().to_string(); - if !stderr.is_empty() { - warn!("{PUBLIC_TOOL_NAME} runner stderr: {stderr}"); - } - } - Err(err) => { - warn!("failed to read {PUBLIC_TOOL_NAME} stderr: {err}"); - } - } - }); - - Ok(CodeModeProcess { - child, - stdin, - stdout_lines: BufReader::new(stdout).lines(), - stderr_task: Some(stderr_task), - pending_messages: HashMap::new(), - }) -} - -async fn drive_code_mode_session( - exec: &ExecContext, - process: &mut CodeModeProcess, - message: HostToNodeMessage, - poll_max_output_tokens: Option>, - is_terminate: bool, -) -> Result { - let started_at = std::time::Instant::now(); - let session_id = match &message { - HostToNodeMessage::Start { session_id, .. } - | HostToNodeMessage::Poll { session_id, .. } - | HostToNodeMessage::Terminate { session_id } - | HostToNodeMessage::Response { session_id, .. } => *session_id, - }; - process - .write(&message) - .await - .map_err(|err| err.to_string())?; - - loop { - let message = process - .read(session_id) - .await - .map_err(|err| err.to_string())?; - if let Some(progress) = handle_node_message( - exec, - process, - session_id, - message, - poll_max_output_tokens, - started_at, - is_terminate, - ) - .await? - { - return Ok(progress); - } - } -} - async fn handle_node_message( exec: &ExecContext, - process: &mut CodeModeProcess, session_id: i32, message: NodeToHostMessage, poll_max_output_tokens: Option>, started_at: std::time::Instant, - is_terminate: bool, -) -> Result, String> { +) -> Result { match message { - NodeToHostMessage::ToolCall { - session_id: message_session_id, - id, - name, - input, - } => { - if is_terminate { - return Ok(None); - } - let response = HostToNodeMessage::Response { - session_id: message_session_id, - id, - code_mode_result: call_nested_tool(exec.clone(), name, input).await, - }; - process - .write(&response) - .await - .map_err(|err| err.to_string())?; - Ok(None) - } + NodeToHostMessage::ToolCall { .. } => Err(format!( + "{PUBLIC_TOOL_NAME} received an unexpected tool call response" + )), NodeToHostMessage::Yielded { content_items, .. } => { - if is_terminate { - return Ok(None); - } let mut delta_items = output_content_items_from_json_values(content_items)?; delta_items = truncate_code_mode_result(delta_items, poll_max_output_tokens.flatten()); prepend_script_status( @@ -528,9 +506,9 @@ async fn handle_node_message( CodeModeExecutionStatus::Running(session_id), started_at.elapsed(), ); - Ok(Some(CodeModeSessionProgress::Yielded { + Ok(CodeModeSessionProgress::Yielded { output: FunctionToolOutput::from_content(delta_items, Some(true)), - })) + }) } NodeToHostMessage::Terminated { content_items, .. } => { let mut delta_items = output_content_items_from_json_values(content_items)?; @@ -540,9 +518,9 @@ async fn handle_node_message( CodeModeExecutionStatus::Terminated, started_at.elapsed(), ); - Ok(Some(CodeModeSessionProgress::Finished( + Ok(CodeModeSessionProgress::Finished( FunctionToolOutput::from_content(delta_items, Some(true)), - ))) + )) } NodeToHostMessage::Result { content_items, @@ -577,19 +555,126 @@ async fn handle_node_message( }, started_at.elapsed(), ); - Ok(Some(CodeModeSessionProgress::Finished( + Ok(CodeModeSessionProgress::Finished( FunctionToolOutput::from_content(delta_items, Some(success)), - ))) + )) } } } -fn message_session_id(message: &NodeToHostMessage) -> i32 { +async fn spawn_code_mode_process( + node_path: &std::path::Path, +) -> Result { + let mut cmd = tokio::process::Command::new(node_path); + cmd.arg("--experimental-vm-modules"); + cmd.arg("--eval"); + cmd.arg(CODE_MODE_RUNNER_SOURCE); + cmd.stdin(std::process::Stdio::piped()) + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()) + .kill_on_drop(true); + + let mut child = cmd.spawn().map_err(std::io::Error::other)?; + let stdout = child.stdout.take().ok_or_else(|| { + std::io::Error::other(format!("{PUBLIC_TOOL_NAME} runner missing stdout")) + })?; + let stderr = child.stderr.take().ok_or_else(|| { + std::io::Error::other(format!("{PUBLIC_TOOL_NAME} runner missing stderr")) + })?; + let stdin = child + .stdin + .take() + .ok_or_else(|| std::io::Error::other(format!("{PUBLIC_TOOL_NAME} runner missing stdin")))?; + let stdin = Arc::new(Mutex::new(stdin)); + let response_waiters = Arc::new(Mutex::new(HashMap::< + String, + oneshot::Sender, + >::new())); + let (tool_call_tx, tool_call_rx) = mpsc::unbounded_channel(); + + tokio::spawn(async move { + let mut reader = BufReader::new(stderr); + let mut buf = Vec::new(); + match reader.read_to_end(&mut buf).await { + Ok(_) => { + let stderr = String::from_utf8_lossy(&buf).trim().to_string(); + if !stderr.is_empty() { + warn!("{PUBLIC_TOOL_NAME} runner stderr: {stderr}"); + } + } + Err(err) => { + warn!("failed to read {PUBLIC_TOOL_NAME} stderr: {err}"); + } + } + }); + let stdout_task = tokio::spawn({ + let response_waiters = Arc::clone(&response_waiters); + async move { + let mut stdout_lines = BufReader::new(stdout).lines(); + loop { + let line = match stdout_lines.next_line().await { + Ok(line) => line, + Err(err) => { + warn!("failed to read {PUBLIC_TOOL_NAME} stdout: {err}"); + break; + } + }; + let Some(line) = line else { + break; + }; + if line.trim().is_empty() { + continue; + } + let message: NodeToHostMessage = match serde_json::from_str(&line) { + Ok(message) => message, + Err(err) => { + warn!("failed to parse {PUBLIC_TOOL_NAME} stdout message: {err}"); + break; + } + }; + match message { + NodeToHostMessage::ToolCall { tool_call } => { + let _ = tool_call_tx.send(tool_call); + } + message => { + let request_id = message_request_id(&message).to_string(); + if let Some(waiter) = response_waiters.lock().await.remove(&request_id) { + let _ = waiter.send(message); + } + } + } + } + response_waiters.lock().await.clear(); + } + }); + + Ok(CodeModeProcess { + child, + stdin, + stdout_task, + response_waiters, + tool_call_rx: Arc::new(Mutex::new(tool_call_rx)), + }) +} + +async fn write_message( + stdin: &Arc>, + message: &HostToNodeMessage, +) -> Result<(), std::io::Error> { + let line = serde_json::to_string(message).map_err(std::io::Error::other)?; + let mut stdin = stdin.lock().await; + stdin.write_all(line.as_bytes()).await?; + stdin.write_all(b"\n").await?; + stdin.flush().await?; + Ok(()) +} + +fn message_request_id(message: &NodeToHostMessage) -> &str { match message { - NodeToHostMessage::ToolCall { session_id, .. } - | NodeToHostMessage::Yielded { session_id, .. } - | NodeToHostMessage::Terminated { session_id, .. } - | NodeToHostMessage::Result { session_id, .. } => *session_id, + NodeToHostMessage::ToolCall { tool_call } => &tool_call.request_id, + NodeToHostMessage::Yielded { request_id, .. } + | NodeToHostMessage::Terminated { request_id, .. } + | NodeToHostMessage::Result { request_id, .. } => request_id, } } diff --git a/codex-rs/core/src/tools/code_mode_runner.cjs b/codex-rs/core/src/tools/code_mode_runner.cjs index d64e369f3..02255b917 100644 --- a/codex-rs/core/src/tools/code_mode_runner.cjs +++ b/codex-rs/core/src/tools/code_mode_runner.cjs @@ -265,6 +265,7 @@ function codeModeWorkerMain() { 'set_max_output_tokens_per_exec_call', 'set_yield_time', 'store', + 'background', ], function initCodeModeModule() { this.setExport('load', load); @@ -288,6 +289,9 @@ function codeModeWorkerMain() { return normalized; }); this.setExport('store', store); + this.setExport('background', () => { + parentPort.postMessage({ type: 'yield' }); + }); }, { context } ); @@ -466,11 +470,16 @@ function createProtocol() { if (message.type === 'poll') { const session = sessions.get(message.session_id); if (session) { - schedulePollYield(protocol, session, normalizeYieldTime(message.yield_time_ms ?? 0)); + session.request_id = String(message.request_id); + if (session.pending_result) { + void completeSession(protocol, sessions, session, session.pending_result); + } else { + schedulePollYield(protocol, session, normalizeYieldTime(message.yield_time_ms ?? 0)); + } } else { void protocol.send({ type: 'result', - session_id: message.session_id, + request_id: message.request_id, content_items: [], stored_values: {}, error_text: `exec session ${message.session_id} not found`, @@ -483,11 +492,12 @@ function createProtocol() { if (message.type === 'terminate') { const session = sessions.get(message.session_id); if (session) { + session.request_id = String(message.request_id); void terminateSession(protocol, sessions, session); } else { void protocol.send({ type: 'result', - session_id: message.session_id, + request_id: message.request_id, content_items: [], stored_values: {}, error_text: `exec session ${message.session_id} not found`, @@ -498,11 +508,11 @@ function createProtocol() { } if (message.type === 'response') { - const entry = pending.get(message.session_id + ':' + message.id); + const entry = pending.get(message.request_id + ':' + message.id); if (!entry) { return; } - pending.delete(message.session_id + ':' + message.id); + pending.delete(message.request_id + ':' + message.id); entry.resolve(message.code_mode_result ?? ''); return; } @@ -537,12 +547,13 @@ function createProtocol() { }); } - function request(sessionId, type, payload) { + function request(type, payload) { + const requestId = 'req-' + ++nextId; const id = 'msg-' + ++nextId; - const pendingKey = sessionId + ':' + id; + const pendingKey = requestId + ':' + id; return new Promise((resolve, reject) => { pending.set(pendingKey, { resolve, reject }); - void send({ type, session_id: sessionId, id, ...payload }).catch((error) => { + void send({ type, request_id: requestId, id, ...payload }).catch((error) => { pending.delete(pendingKey); reject(error); }); @@ -565,7 +576,9 @@ function startSession(protocol, sessions, start) { initial_yield_timer: null, initial_yield_triggered: false, max_output_tokens_per_exec_call: DEFAULT_MAX_OUTPUT_TOKENS_PER_EXEC_CALL, + pending_result: null, poll_yield_timer: null, + request_id: String(start.request_id), worker: new Worker(sessionWorkerSource(), { eval: true, workerData: start, @@ -620,18 +633,30 @@ async function handleWorkerMessage(protocol, sessions, session, message) { return; } + if (message.type === 'yield') { + void sendYielded(protocol, session); + return; + } + if (message.type === 'tool_call') { void forwardToolCall(protocol, session, message); return; } if (message.type === 'result') { - await completeSession(protocol, sessions, session, { + const result = { type: 'result', stored_values: cloneJsonValue(message.stored_values ?? {}), error_text: typeof message.error_text === 'string' ? message.error_text : undefined, - }); + }; + if (session.request_id === null) { + session.pending_result = result; + session.initial_yield_timer = clearTimer(session.initial_yield_timer); + session.poll_yield_timer = clearTimer(session.poll_yield_timer); + return; + } + await completeSession(protocol, sessions, session, result); return; } @@ -640,7 +665,7 @@ async function handleWorkerMessage(protocol, sessions, session, message) { async function forwardToolCall(protocol, session, message) { try { - const result = await protocol.request(session.id, 'tool_call', { + const result = await protocol.request('tool_call', { name: String(message.name), input: message.input, }); @@ -669,18 +694,20 @@ async function forwardToolCall(protocol, session, message) { } async function sendYielded(protocol, session) { - if (session.completed) { + if (session.completed || session.request_id === null) { return; } const contentItems = takeContentItems(session); + const requestId = session.request_id; try { session.worker.postMessage({ type: 'clear_content' }); } catch {} await protocol.send({ type: 'yielded', - session_id: session.id, + request_id: requestId, content_items: contentItems, }); + session.request_id = null; } function scheduleInitialYield(protocol, session, yieldTime) { @@ -711,17 +738,25 @@ async function completeSession(protocol, sessions, session, message) { if (session.completed) { return; } + if (session.request_id === null) { + session.pending_result = message; + session.initial_yield_timer = clearTimer(session.initial_yield_timer); + session.poll_yield_timer = clearTimer(session.poll_yield_timer); + return; + } + const requestId = session.request_id; session.completed = true; session.initial_yield_timer = clearTimer(session.initial_yield_timer); session.poll_yield_timer = clearTimer(session.poll_yield_timer); sessions.delete(session.id); const contentItems = takeContentItems(session); + session.pending_result = null; try { session.worker.postMessage({ type: 'clear_content' }); } catch {} await protocol.send({ ...message, - session_id: session.id, + request_id: requestId, content_items: contentItems, max_output_tokens_per_exec_call: session.max_output_tokens_per_exec_call, }); @@ -741,7 +776,7 @@ async function terminateSession(protocol, sessions, session) { } catch {} await protocol.send({ type: 'terminated', - session_id: session.id, + request_id: session.request_id, content_items: contentItems, }); } diff --git a/codex-rs/core/tests/suite/code_mode.rs b/codex-rs/core/tests/suite/code_mode.rs index 23fcd9c08..976c553dc 100644 --- a/codex-rs/core/tests/suite/code_mode.rs +++ b/codex-rs/core/tests/suite/code_mode.rs @@ -845,8 +845,12 @@ async fn code_mode_exec_wait_terminate_returns_completed_session_if_it_finished_ let test = builder.build(&server).await?; let session_a_gate = test.workspace_path("code-mode-session-a-finished.ready"); let session_b_gate = test.workspace_path("code-mode-session-b-blocked.ready"); + let session_a_done_marker = test.workspace_path("code-mode-session-a-done.txt"); let session_a_wait = wait_for_file_source(&session_a_gate)?; let session_b_wait = wait_for_file_source(&session_b_gate)?; + let session_a_done_marker_quoted = + shlex::try_join([session_a_done_marker.to_string_lossy().as_ref()])?; + let session_a_done_command = format!("printf done > {session_a_done_marker_quoted}"); let session_a_code = format!( r#" @@ -857,6 +861,7 @@ output_text("session a start"); set_yield_time(10); {session_a_wait} output_text("session a done"); +await exec_command({{ cmd: {session_a_done_command:?} }}); "# ); let session_b_code = format!( @@ -966,6 +971,14 @@ output_text("session b done"); session_b_id ); + for _ in 0..100 { + if session_a_done_marker.exists() { + break; + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + assert!(session_a_done_marker.exists()); + responses::mount_sse_once( &server, sse(vec![ @@ -995,14 +1008,124 @@ output_text("session b done"); let fourth_request = fourth_completion.single_request(); let fourth_items = function_tool_output_items(&fourth_request, "call-4"); - assert_eq!(fourth_items.len(), 1); + match fourth_items.len() { + 1 => { + assert_regex_match( + concat!( + r"(?s)\A", + r"Script terminated\nWall time \d+\.\d seconds\nOutput:\n\z" + ), + text_item(&fourth_items, 0), + ); + } + 2 => { + assert_regex_match( + concat!( + r"(?s)\A", + r"Script (?:completed|terminated)\nWall time \d+\.\d seconds\nOutput:\n\z" + ), + text_item(&fourth_items, 0), + ); + assert_eq!(text_item(&fourth_items, 1), "session a done"); + } + other => panic!("unexpected number of content items: {other}"), + } + + Ok(()) +} + +#[cfg_attr(windows, ignore = "no exec_command on Windows")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn code_mode_background_keeps_running_on_later_turn_without_exec_wait() -> Result<()> { + skip_if_no_network!(Ok(())); + + let server = responses::start_mock_server().await; + let mut builder = test_codex().with_config(move |config| { + let _ = config.features.enable(Feature::CodeMode); + }); + let test = builder.build(&server).await?; + let resumed_file = test.workspace_path("code-mode-yield-resumed.txt"); + let resumed_file_quoted = shlex::try_join([resumed_file.to_string_lossy().as_ref()])?; + let write_file_command = format!("printf resumed > {resumed_file_quoted}"); + let wait_for_file_command = + format!("while [ ! -f {resumed_file_quoted} ]; do sleep 0.01; done; printf ready"); + let code = format!( + r#" +import {{ background, output_text }} from "@openai/code_mode"; +import {{ exec_command }} from "tools.js"; + +output_text("before yield"); +background(); +await exec_command({{ cmd: {write_file_command:?} }}); +output_text("after yield"); +"# + ); + + responses::mount_sse_once( + &server, + sse(vec![ + ev_response_created("resp-1"), + ev_custom_tool_call("call-1", "exec", &code), + ev_completed("resp-1"), + ]), + ) + .await; + let first_completion = responses::mount_sse_once( + &server, + sse(vec![ + ev_assistant_message("msg-1", "exec yielded"), + ev_completed("resp-2"), + ]), + ) + .await; + + test.submit_turn("start yielded exec").await?; + + let first_request = first_completion.single_request(); + let first_items = custom_tool_output_items(&first_request, "call-1"); + assert_eq!(first_items.len(), 2); assert_regex_match( concat!( r"(?s)\A", - r"Script terminated\nWall time \d+\.\d seconds\nOutput:\n\z" + r"Script running with session ID \d+\nWall time \d+\.\d seconds\nOutput:\n\z" ), - text_item(&fourth_items, 0), + text_item(&first_items, 0), ); + assert_eq!(text_item(&first_items, 1), "before yield"); + + responses::mount_sse_once( + &server, + sse(vec![ + ev_response_created("resp-3"), + responses::ev_function_call( + "call-2", + "exec_command", + &serde_json::to_string(&serde_json::json!({ + "cmd": wait_for_file_command, + }))?, + ), + ev_completed("resp-3"), + ]), + ) + .await; + let second_completion = responses::mount_sse_once( + &server, + sse(vec![ + ev_assistant_message("msg-2", "file appeared"), + ev_completed("resp-4"), + ]), + ) + .await; + + test.submit_turn("wait for resumed file").await?; + + let second_request = second_completion.single_request(); + assert!( + second_request + .function_call_output_text("call-2") + .is_some_and(|output| output.ends_with("ready")) + ); + assert_eq!(fs::read_to_string(&resumed_file)?, "resumed"); Ok(()) }