diff --git a/codex-rs/core/config.schema.json b/codex-rs/core/config.schema.json index 4f07e475b210..979090546819 100644 --- a/codex-rs/core/config.schema.json +++ b/codex-rs/core/config.schema.json @@ -486,6 +486,9 @@ "code_mode_host": { "$ref": "#/definitions/FeatureToml_for_CodeModeHostConfigToml" }, + "code_mode_interrupt": { + "type": "boolean" + }, "code_mode_only": { "type": "boolean" }, @@ -5148,6 +5151,9 @@ "code_mode_host": { "$ref": "#/definitions/FeatureToml_for_CodeModeHostConfigToml" }, + "code_mode_interrupt": { + "type": "boolean" + }, "code_mode_only": { "type": "boolean" }, diff --git a/codex-rs/core/src/tasks/mod.rs b/codex-rs/core/src/tasks/mod.rs index 75506c4e8585..0a2b39c49c87 100644 --- a/codex-rs/core/src/tasks/mod.rs +++ b/codex-rs/core/src/tasks/mod.rs @@ -876,6 +876,18 @@ impl Session { .complete_task_end(&sub_id); trace!(task_kind = ?task.kind, sub_id, "aborting running task"); task.cancellation_token.cancel(); + if reason == TurnAbortReason::Interrupted + && task + .turn_context + .config + .features + .enabled(Feature::CodeModeInterrupt) + { + self.services + .code_mode_service + .interrupt_active_cells() + .await; + } task.turn_context .turn_metadata_state .cancel_git_enrichment_task(); diff --git a/codex-rs/core/src/tools/code_mode/delegate.rs b/codex-rs/core/src/tools/code_mode/delegate.rs index 70defb24215c..bb46a06e740d 100644 --- a/codex-rs/core/src/tools/code_mode/delegate.rs +++ b/codex-rs/core/src/tools/code_mode/delegate.rs @@ -45,6 +45,15 @@ impl CodeModeDispatchBroker { remove_dispatch_gate(&self.dispatch_gates, cell_id); } + pub(super) fn active_cell_ids(&self) -> Vec { + self.dispatch_gates + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .keys() + .cloned() + .collect() + } + pub(super) fn start_turn_worker( &self, exec: ExecContext, @@ -105,12 +114,13 @@ impl CodeModeDispatchBroker { } let host = Arc::clone(&host); tokio::spawn(async move { + let invocation = + host.invoke_tool(invocation, cancellation_token.clone()); + tokio::pin!(invocation); let response = tokio::select! { - response = host.invoke_tool( - invocation, - cancellation_token.clone(), - ) => response, - _ = cancellation_token.cancelled() => return, + biased; + _ = cancellation_token.cancelled() => invocation.await, + response = &mut invocation => response, }; let _ = response_tx.send(response); }); diff --git a/codex-rs/core/src/tools/code_mode/mod.rs b/codex-rs/core/src/tools/code_mode/mod.rs index f6375e75c00b..0ce1e66b9cf0 100644 --- a/codex-rs/core/src/tools/code_mode/mod.rs +++ b/codex-rs/core/src/tools/code_mode/mod.rs @@ -18,6 +18,7 @@ use codex_code_mode::CodeModeSessionProvider; use codex_code_mode::CodeModeToolKind; use codex_code_mode::RuntimeResponse; use codex_protocol::models::FunctionCallOutputContentItem; +use futures::future::join_all; use serde_json::Value as JsonValue; use tokio::sync::OnceCell; use tokio_util::sync::CancellationToken; @@ -141,6 +142,23 @@ impl CodeModeService { self.session().await?.terminate(cell_id).await } + pub(crate) async fn interrupt_active_cells(&self) { + let Some(session) = self.session.get() else { + return; + }; + join_all( + self.dispatch_broker + .active_cell_ids() + .into_iter() + .map(|cell_id| async move { + if let Err(error) = session.terminate(cell_id.clone()).await { + tracing::warn!(%cell_id, %error, "failed to terminate interrupted code-mode cell"); + } + }), + ) + .await; + } + pub(crate) async fn shutdown(&self) -> Result<(), String> { self.shutting_down.store(true, Ordering::Release); // Join any initialization already in progress without initializing an unused service. diff --git a/codex-rs/core/tests/suite/code_mode.rs b/codex-rs/core/tests/suite/code_mode.rs index d8b6fabbb320..0fefaca196fe 100644 --- a/codex-rs/core/tests/suite/code_mode.rs +++ b/codex-rs/core/tests/suite/code_mode.rs @@ -10,7 +10,12 @@ use codex_core::config::Config; use codex_core::config::CurrentTimeReminderConfig; use codex_extension_api::ExtensionData; use codex_extension_api::ExtensionRegistryBuilder; +use codex_extension_api::ToolCallOutcome; use codex_extension_api::ToolContributor; +use codex_extension_api::ToolFinishInput; +use codex_extension_api::ToolLifecycleContributor; +use codex_extension_api::ToolLifecycleFuture; +use codex_extension_api::ToolStartInput; use codex_features::CurrentTimeSource; use codex_features::Feature; use codex_login::CodexAuth; @@ -80,11 +85,13 @@ use std::fs; use std::io::Cursor; use std::path::Path; use std::sync::Arc; +use std::sync::Mutex; use std::sync::atomic::AtomicUsize; use std::sync::atomic::Ordering; use std::time::Duration; use std::time::Instant; use test_case::test_case; +use tokio::sync::oneshot; use wiremock::Mock; use wiremock::MockServer; use wiremock::ResponseTemplate; @@ -3246,6 +3253,208 @@ text("after yield"); Ok(()) } +struct InterruptedNestedToolObserver { + started: Mutex>>, + finished: Mutex>>, +} + +impl ToolLifecycleContributor for InterruptedNestedToolObserver { + fn on_tool_start<'a>(&'a self, input: ToolStartInput<'a>) -> ToolLifecycleFuture<'a> { + Box::pin(async move { + let codex_extension_api::ToolCallSource::CodeMode { cell_id, .. } = input.source else { + return; + }; + if input.tool_name.name != "test_sync_tool" { + return; + } + if let Some(started) = self + .started + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .take() + { + let _ = started.send(cell_id); + } + }) + } + + fn on_tool_finish<'a>(&'a self, input: ToolFinishInput<'a>) -> ToolLifecycleFuture<'a> { + Box::pin(async move { + if input.tool_name.name != "test_sync_tool" + || !matches!( + input.source, + codex_extension_api::ToolCallSource::CodeMode { .. } + ) + { + return; + } + if let Some(finished) = self + .finished + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .take() + { + let _ = finished.send(input.outcome); + } + }) + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn code_mode_interrupt_terminates_active_cells_and_nested_tools() -> Result<()> { + skip_if_no_network!(Ok(())); + + let (started_tx, started_rx) = oneshot::channel(); + let (finished_tx, finished_rx) = oneshot::channel(); + let mut extensions = ExtensionRegistryBuilder::::new(); + extensions.tool_lifecycle_contributor(Arc::new(InterruptedNestedToolObserver { + started: Mutex::new(Some(started_tx)), + finished: Mutex::new(Some(finished_tx)), + })); + + let server = responses::start_mock_server().await; + let mut builder = test_codex() + .with_model("test-gpt-5.1-codex") + .with_extensions(Arc::new(extensions.build())) + .with_config(|config| { + let _ = config.features.enable(Feature::CodeMode); + let _ = config.features.enable(Feature::CodeModeInterrupt); + }); + let test = builder.build_with_auto_env(&server).await?; + + let setup = responses::mount_sse_sequence( + &server, + vec![ + sse(vec![ + ev_response_created("resp-store"), + ev_custom_tool_call("call-store", "exec", r#"store("persisted", "preserved");"#), + ev_completed("resp-store"), + ]), + sse(vec![ + ev_assistant_message("msg-store", "stored"), + ev_completed("resp-store-complete"), + ]), + sse(vec![ + ev_response_created("resp-background"), + ev_custom_tool_call( + "call-background", + "exec", + "yield_control(); await new Promise(() => {});", + ), + ev_completed("resp-background"), + ]), + sse(vec![ + ev_assistant_message("msg-background", "running"), + ev_completed("resp-background-complete"), + ]), + ], + ) + .await; + + test.submit_turn("store a value in the reusable code-mode session") + .await?; + test.submit_turn("start a background code-mode cell") + .await?; + let background_response = setup + .last_request() + .expect("background cell should be returned to the model"); + let background_items = custom_tool_output_items(&background_response, "call-background"); + assert!( + text_item(&background_items, /*index*/ 0).starts_with("Script running with cell ID "), + "background cell should remain active: {background_items:?}" + ); + let background_cell_id = + extract_running_cell_id(text_item(&background_items, /*index*/ 0)); + + responses::mount_sse_once( + &server, + sse(vec![ + ev_response_created("resp-interrupted"), + ev_custom_tool_call( + "call-interrupted", + "exec", + "await tools.test_sync_tool({ sleep_after_ms: 60_000 });", + ), + ev_completed("resp-interrupted"), + ]), + ) + .await; + + test.codex + .submit(Op::UserInput { + items: vec![UserInput::Text { + text: "start a long-running nested tool".to_string(), + text_elements: Vec::new(), + }], + final_output_json_schema: None, + responsesapi_client_metadata: None, + additional_context: Default::default(), + thread_settings: Default::default(), + }) + .await?; + let active_cell_id = tokio::time::timeout(Duration::from_secs(10), started_rx).await??; + + test.codex.submit(Op::Interrupt).await?; + wait_for_event(&test.codex, |event| { + matches!(event, EventMsg::TurnAborted(_)) + }) + .await; + let nested_outcome = tokio::time::timeout(Duration::from_secs(10), finished_rx).await??; + assert_eq!(nested_outcome, ToolCallOutcome::Aborted); + + let recovery = responses::mount_sse_sequence( + &server, + vec![ + sse(vec![ + ev_response_created("resp-wait-background"), + responses::ev_function_call( + "call-wait-background", + "wait", + &serde_json::to_string(&serde_json::json!({ + "cell_id": background_cell_id, + "yield_time_ms": 1, + }))?, + ), + ev_completed("resp-wait-background"), + ]), + sse(vec![ + ev_response_created("resp-wait-active"), + responses::ev_function_call( + "call-wait-active", + "wait", + &serde_json::to_string(&serde_json::json!({ + "cell_id": active_cell_id, + "yield_time_ms": 1, + }))?, + ), + ev_completed("resp-wait-active"), + ]), + sse(vec![ + ev_response_created("resp-recovery"), + ev_custom_tool_call("call-recovery", "exec", r#"text(load("persisted"));"#), + ev_completed("resp-recovery"), + ]), + sse(vec![ + ev_assistant_message("msg-recovery", "recovered"), + ev_completed("resp-recovery-complete"), + ]), + ], + ) + .await; + + test.submit_turn("verify interrupted cells and reuse their session") + .await?; + let requests = recovery.requests(); + let background_output = function_tool_output_items(&requests[1], "call-wait-background"); + assert!(text_item(&background_output, /*index*/ 1).contains("not found")); + let active_output = function_tool_output_items(&requests[2], "call-wait-active"); + assert!(text_item(&active_output, /*index*/ 1).contains("not found")); + let recovery_items = custom_tool_output_items(&requests[3], "call-recovery"); + assert_eq!(text_item(&recovery_items, /*index*/ 1), "preserved"); + + Ok(()) +} + #[cfg_attr(windows, ignore = "no exec_command on Windows")] #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn code_mode_wait_uses_its_own_max_tokens_budget() -> Result<()> { diff --git a/codex-rs/features/src/lib.rs b/codex-rs/features/src/lib.rs index a861529b39b7..d0b1847977a9 100644 --- a/codex-rs/features/src/lib.rs +++ b/codex-rs/features/src/lib.rs @@ -103,6 +103,8 @@ pub enum Feature { CodeModeBufferedExec, /// Run JavaScript code mode in the standalone host process. CodeModeHost, + /// Terminate active code mode cells when their turn is interrupted. + CodeModeInterrupt, /// Restrict model-visible tools to code mode entrypoints (`exec`, `wait`). CodeModeOnly, /// Use the single unified PTY-backed exec tool. @@ -933,6 +935,12 @@ pub const FEATURES: &[FeatureSpec] = &[ stage: Stage::Stable, default_enabled: true, }, + FeatureSpec { + id: Feature::CodeModeInterrupt, + key: "code_mode_interrupt", + stage: Stage::UnderDevelopment, + default_enabled: false, + }, FeatureSpec { id: Feature::CodeModeOnly, key: "code_mode_only",