From 553376ffcb5d4848ebb52cb7fb022e6d836d8eb6 Mon Sep 17 00:00:00 2001 From: Joe Gershenson Date: Mon, 20 Apr 2026 23:49:29 -0700 Subject: [PATCH] Fix interrupt prompt history race --- codex-rs/core/src/hook_runtime.rs | 62 +++++++++++++++++++++ codex-rs/core/src/session/review.rs | 2 + codex-rs/core/src/session/tests.rs | 68 ++++++----------------- codex-rs/core/src/session/turn.rs | 30 +--------- codex-rs/core/src/session/turn_context.rs | 6 ++ codex-rs/core/src/tasks/mod.rs | 15 ++++- 6 files changed, 102 insertions(+), 81 deletions(-) diff --git a/codex-rs/core/src/hook_runtime.rs b/codex-rs/core/src/hook_runtime.rs index 342fca91e8..a0292d9956 100644 --- a/codex-rs/core/src/hook_runtime.rs +++ b/codex-rs/core/src/hook_runtime.rs @@ -17,6 +17,7 @@ use codex_hooks::UserPromptSubmitRequest; use codex_otel::HOOK_RUN_DURATION_METRIC; use codex_otel::HOOK_RUN_METRIC; use codex_protocol::items::TurnItem; +use codex_protocol::items::UserMessageItem; use codex_protocol::models::DeveloperInstructions; use codex_protocol::models::ResponseInputItem; use codex_protocol::models::ResponseItem; @@ -268,6 +269,67 @@ pub(crate) async fn inspect_pending_input( } } +pub(crate) async fn drain_turn_start_transcript_inputs( + sess: &Arc, + turn_context: &Arc, +) -> bool { + let _guard = turn_context.transcript_serialization_lock.lock().await; + if run_pending_session_start_hooks(sess, turn_context).await { + return false; + } + + loop { + let input = { + let inputs = turn_context.turn_start_transcript_inputs.lock().await; + inputs.first().cloned() + }; + let Some(input) = input else { + break; + }; + + let initial_input_for_turn: ResponseInputItem = ResponseInputItem::from(input.clone()); + let response_item: ResponseItem = initial_input_for_turn.into(); + let user_prompt_submit_outcome = run_user_prompt_submit_hooks( + sess, + turn_context, + UserMessageItem::new(&input).message(), + ) + .await; + if user_prompt_submit_outcome.should_stop { + record_additional_contexts( + sess, + turn_context, + user_prompt_submit_outcome.additional_contexts, + ) + .await; + let mut inputs = turn_context.turn_start_transcript_inputs.lock().await; + if inputs.first().is_some_and(|queued| queued == &input) { + inputs.remove(0); + } + return false; + } + + sess.record_user_prompt_and_emit_turn_item( + turn_context.as_ref(), + input.as_slice(), + response_item, + ) + .await; + record_additional_contexts( + sess, + turn_context, + user_prompt_submit_outcome.additional_contexts, + ) + .await; + let mut inputs = turn_context.turn_start_transcript_inputs.lock().await; + if inputs.first().is_some_and(|queued| queued == &input) { + inputs.remove(0); + } + } + + true +} + pub(crate) async fn record_pending_input( sess: &Arc, turn_context: &Arc, diff --git a/codex-rs/core/src/session/review.rs b/codex-rs/core/src/session/review.rs index 94de4617d5..cfd58e0631 100644 --- a/codex-rs/core/src/session/review.rs +++ b/codex-rs/core/src/session/review.rs @@ -139,6 +139,8 @@ pub(super) async fn spawn_review_thread( turn_metadata_state, turn_skills: TurnSkillsContext::new(parent_turn_context.turn_skills.outcome.clone()), turn_timing_state: Arc::new(TurnTimingState::default()), + turn_start_transcript_inputs: Arc::new(Mutex::new(Vec::new())), + transcript_serialization_lock: Arc::new(Mutex::new(())), }; // Seed the child task with the review prompt as the initial user message. diff --git a/codex-rs/core/src/session/tests.rs b/codex-rs/core/src/session/tests.rs index f901e5bee4..babf659f6d 100644 --- a/codex-rs/core/src/session/tests.rs +++ b/codex-rs/core/src/session/tests.rs @@ -144,7 +144,6 @@ use sha2::Sha512; use std::path::Path; use std::time::Duration; use tokio::sync::Semaphore; -use tokio::time::sleep; use tokio::time::timeout; use tracing_opentelemetry::OpenTelemetrySpanExt; use wiremock::Mock; @@ -5669,49 +5668,14 @@ impl SessionTask for NeverEndingTask { cancellation_token.cancelled().await; return None; } - loop { - sleep(Duration::from_secs(60)).await; - } + std::future::pending::>().await } } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[test_log::test] -async fn abort_regular_task_emits_turn_aborted_only() { - let (sess, tc, rx) = make_session_and_context_with_rx().await; - let input = vec![UserInput::Text { - text: "hello".to_string(), - text_elements: Vec::new(), - }]; - sess.spawn_task( - Arc::clone(&tc), - input, - NeverEndingTask { - kind: TaskKind::Regular, - listen_to_cancellation_token: false, - }, - ) - .await; - - sess.abort_all_tasks(TurnAbortReason::Interrupted).await; - - // Interrupts persist a model-visible `` marker into history, but there is no - // separate client-visible event for that marker (only `EventMsg::TurnAborted`). - let evt = tokio::time::timeout(std::time::Duration::from_secs(2), rx.recv()) - .await - .expect("timeout waiting for event") - .expect("event"); - match evt.msg { - EventMsg::TurnAborted(e) => assert_eq!(TurnAbortReason::Interrupted, e.reason), - other => panic!("unexpected event: {other:?}"), - } - // No extra events should be emitted after an abort. - assert!(rx.try_recv().is_err()); -} - -#[tokio::test] -async fn abort_gracefully_emits_turn_aborted_only() { - let (sess, tc, rx) = make_session_and_context_with_rx().await; +async fn abort_regular_task_records_prompt_before_interrupt_marker() { + let (sess, tc, _rx) = make_session_and_context_with_rx().await; let input = vec![UserInput::Text { text: "hello".to_string(), text_elements: Vec::new(), @@ -5728,18 +5692,20 @@ async fn abort_gracefully_emits_turn_aborted_only() { sess.abort_all_tasks(TurnAbortReason::Interrupted).await; - // Even if tasks handle cancellation gracefully, interrupts still result in `TurnAborted` - // being the only client-visible signal. - let evt = tokio::time::timeout(std::time::Duration::from_secs(2), rx.recv()) - .await - .expect("timeout waiting for event") - .expect("event"); - match evt.msg { - EventMsg::TurnAborted(e) => assert_eq!(TurnAbortReason::Interrupted, e.reason), - other => panic!("unexpected event: {other:?}"), - } - // No extra events should be emitted after an abort. - assert!(rx.try_recv().is_err()); + let history = sess.clone_history().await; + let expected = vec![ + ResponseItem::Message { + id: None, + role: "user".to_string(), + content: vec![ContentItem::InputText { + text: "hello".to_string(), + }], + end_turn: None, + phase: None, + }, + crate::tasks::interrupted_turn_history_marker(), + ]; + assert_eq!(history.raw_items(), expected.as_slice()); } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] diff --git a/codex-rs/core/src/session/turn.rs b/codex-rs/core/src/session/turn.rs index efbcca6716..a05b0b16d1 100644 --- a/codex-rs/core/src/session/turn.rs +++ b/codex-rs/core/src/session/turn.rs @@ -18,12 +18,12 @@ use crate::compact_remote::run_inline_remote_auto_compact_task; use crate::connectors; use crate::feedback_tags; use crate::hook_runtime::PendingInputHookDisposition; +use crate::hook_runtime::drain_turn_start_transcript_inputs; use crate::hook_runtime::emit_hook_completed_events; use crate::hook_runtime::inspect_pending_input; use crate::hook_runtime::record_additional_contexts; use crate::hook_runtime::record_pending_input; use crate::hook_runtime::run_pending_session_start_hooks; -use crate::hook_runtime::run_user_prompt_submit_hooks; use crate::injection::ToolMentionKind; use crate::injection::app_id_from_path; use crate::injection::tool_kind_for_path; @@ -74,7 +74,6 @@ use codex_protocol::error::CodexErr; use codex_protocol::error::Result as CodexResult; use codex_protocol::items::PlanItem; use codex_protocol::items::TurnItem; -use codex_protocol::items::UserMessageItem; use codex_protocol::items::build_hook_prompt_message; use codex_protocol::models::BaseInstructions; use codex_protocol::models::ContentItem; @@ -285,33 +284,9 @@ pub(crate) async fn run_turn( }) .collect::>(); - if run_pending_session_start_hooks(&sess, &turn_context).await { + if !drain_turn_start_transcript_inputs(&sess, &turn_context).await { return None; } - let additional_contexts = if input.is_empty() { - Vec::new() - } else { - let initial_input_for_turn: ResponseInputItem = ResponseInputItem::from(input.clone()); - let response_item: ResponseItem = initial_input_for_turn.clone().into(); - let user_prompt_submit_outcome = run_user_prompt_submit_hooks( - &sess, - &turn_context, - UserMessageItem::new(&input).message(), - ) - .await; - if user_prompt_submit_outcome.should_stop { - record_additional_contexts( - &sess, - &turn_context, - user_prompt_submit_outcome.additional_contexts, - ) - .await; - return None; - } - sess.record_user_prompt_and_emit_turn_item(turn_context.as_ref(), &input, response_item) - .await; - user_prompt_submit_outcome.additional_contexts - }; sess.services .analytics_events_client .track_app_mentioned(tracking.clone(), mentioned_app_invocations); @@ -322,7 +297,6 @@ pub(crate) async fn run_turn( } sess.merge_connector_selection(explicitly_enabled_connectors.clone()) .await; - record_additional_contexts(&sess, &turn_context, additional_contexts).await; if !input.is_empty() { // Track the previous-turn baseline from the regular user-turn path only so // standalone tasks (compact/shell/review/undo) cannot suppress future diff --git a/codex-rs/core/src/session/turn_context.rs b/codex-rs/core/src/session/turn_context.rs index dd86804ee5..27da2ad5dc 100644 --- a/codex-rs/core/src/session/turn_context.rs +++ b/codex-rs/core/src/session/turn_context.rs @@ -71,6 +71,8 @@ pub(crate) struct TurnContext { pub(crate) turn_metadata_state: Arc, pub(crate) turn_skills: TurnSkillsContext, pub(crate) turn_timing_state: Arc, + pub(crate) turn_start_transcript_inputs: Arc>>>, + pub(crate) transcript_serialization_lock: Arc>, } impl TurnContext { pub(crate) fn model_context_window(&self) -> Option { @@ -197,6 +199,8 @@ impl TurnContext { turn_metadata_state: self.turn_metadata_state.clone(), turn_skills: self.turn_skills.clone(), turn_timing_state: Arc::clone(&self.turn_timing_state), + turn_start_transcript_inputs: Arc::clone(&self.turn_start_transcript_inputs), + transcript_serialization_lock: Arc::clone(&self.transcript_serialization_lock), } } @@ -443,6 +447,8 @@ impl Session { turn_metadata_state, turn_skills: TurnSkillsContext::new(skills_outcome), turn_timing_state: Arc::new(TurnTimingState::default()), + turn_start_transcript_inputs: Arc::new(Mutex::new(Vec::new())), + transcript_serialization_lock: Arc::new(Mutex::new(())), } } diff --git a/codex-rs/core/src/tasks/mod.rs b/codex-rs/core/src/tasks/mod.rs index e7030d1d94..50dc3f9744 100644 --- a/codex-rs/core/src/tasks/mod.rs +++ b/codex-rs/core/src/tasks/mod.rs @@ -22,6 +22,7 @@ use tracing::warn; use crate::contextual_user_message::TURN_ABORTED_CLOSE_TAG; use crate::contextual_user_message::TURN_ABORTED_OPEN_TAG; use crate::hook_runtime::PendingInputHookDisposition; +use crate::hook_runtime::drain_turn_start_transcript_inputs; use crate::hook_runtime::inspect_pending_input; use crate::hook_runtime::record_additional_contexts; use crate::hook_runtime::record_pending_input; @@ -305,6 +306,13 @@ impl Session { let done_clone = Arc::clone(&done); let session_ctx = Arc::new(SessionTaskContext::new(Arc::clone(self))); let ctx = Arc::clone(&turn_context); + if task_kind == TaskKind::Regular && !input.is_empty() { + turn_context + .turn_start_transcript_inputs + .lock() + .await + .push(input.clone()); + } let task_for_run = Arc::clone(&task); let task_cancellation_token = cancellation_token.child_token(); // Task-owned turn spans keep a core-owned span open for the @@ -618,6 +626,11 @@ impl Session { .cancel_git_enrichment_task(); let session_task = task.task; + if reason == TurnAbortReason::Interrupted { + self.cleanup_after_interrupt(&task.turn_context).await; + let _ = drain_turn_start_transcript_inputs(self, &task.turn_context).await; + } + select! { _ = task.done.notified() => { }, @@ -634,8 +647,6 @@ impl Session { .await; if reason == TurnAbortReason::Interrupted { - self.cleanup_after_interrupt(&task.turn_context).await; - let marker = interrupted_turn_history_marker(); self.record_into_history(std::slice::from_ref(&marker), task.turn_context.as_ref()) .await;