From dd7aaab82520b7c38dce7ffd6568eb2535fe46d7 Mon Sep 17 00:00:00 2001 From: Joe Gershenson Date: Tue, 21 Apr 2026 00:40:08 -0700 Subject: [PATCH] codex: harden interrupt prompt drain --- codex-rs/core/src/hook_runtime.rs | 81 ++++++++++++++++------- codex-rs/core/src/session/review.rs | 2 +- codex-rs/core/src/session/tests.rs | 52 +++++++++++++++ codex-rs/core/src/session/turn_context.rs | 25 ++++++- codex-rs/core/src/tasks/mod.rs | 21 +++--- 5 files changed, 144 insertions(+), 37 deletions(-) diff --git a/codex-rs/core/src/hook_runtime.rs b/codex-rs/core/src/hook_runtime.rs index 2d3e7087eb..8d73b38275 100644 --- a/codex-rs/core/src/hook_runtime.rs +++ b/codex-rs/core/src/hook_runtime.rs @@ -35,6 +35,7 @@ use serde_json::Value; use crate::event_mapping::parse_turn_item; use crate::session::session::Session; use crate::session::turn_context::TurnContext; +use crate::session::turn_context::TurnStartUserPromptSubmitOutcome; use crate::tools::sandboxing::PermissionRequestPayload; pub(crate) struct HookRuntimeOutcome { @@ -277,31 +278,48 @@ pub(crate) async fn drain_turn_start_transcript_inputs( return false; }; if run_pending_session_start_hooks(sess, turn_context).await { - turn_context - .turn_start_transcript_inputs - .lock() - .await - .clear(); + turn_context.lock_turn_start_transcript_inputs().clear(); return false; } loop { - let input = { - let inputs = turn_context.turn_start_transcript_inputs.lock().await; + let queued_input = { + let inputs = turn_context.lock_turn_start_transcript_inputs(); inputs.first().cloned() }; - let Some(input) = input else { + let Some(queued_input) = queued_input else { break; }; + let input = queued_input.input; 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; + let user_prompt_submit_outcome = match queued_input.user_prompt_submit_outcome { + Some(outcome) => outcome, + None => { + let outcome = run_user_prompt_submit_hooks( + sess, + turn_context, + UserMessageItem::new(&input).message(), + ) + .await; + let outcome = TurnStartUserPromptSubmitOutcome { + should_stop: outcome.should_stop, + additional_contexts: outcome.additional_contexts, + }; + { + let mut inputs = turn_context.lock_turn_start_transcript_inputs(); + if let Some(queued) = inputs.first_mut() + && queued.input == input + && queued.user_prompt_submit_outcome.is_none() + { + queued.user_prompt_submit_outcome = Some(outcome.clone()); + } + } + outcome + } + }; + if user_prompt_submit_outcome.should_stop { record_additional_contexts( sess, @@ -309,27 +327,42 @@ pub(crate) async fn drain_turn_start_transcript_inputs( 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) { + let mut inputs = turn_context.lock_turn_start_transcript_inputs(); + if inputs.first().is_some_and(|queued| queued.input == input) { inputs.remove(0); } return false; } - sess.record_user_prompt_and_emit_turn_item( - turn_context.as_ref(), - input.as_slice(), - response_item, - ) - .await; + if !queued_input.user_prompt_recorded { + sess.record_conversation_items( + turn_context.as_ref(), + std::slice::from_ref(&response_item), + ) + .await; + { + let mut inputs = turn_context.lock_turn_start_transcript_inputs(); + if let Some(queued) = inputs.first_mut() + && queued.input == input + { + queued.user_prompt_recorded = true; + } + } + let turn_item = TurnItem::UserMessage(UserMessageItem::new(&input)); + sess.emit_turn_item_started(turn_context.as_ref(), &turn_item) + .await; + sess.emit_turn_item_completed(turn_context.as_ref(), turn_item) + .await; + sess.ensure_rollout_materialized().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) { + let mut inputs = turn_context.lock_turn_start_transcript_inputs(); + if inputs.first().is_some_and(|queued| queued.input == input) { inputs.remove(0); } } diff --git a/codex-rs/core/src/session/review.rs b/codex-rs/core/src/session/review.rs index b0b48f4ac9..c0a0c51af4 100644 --- a/codex-rs/core/src/session/review.rs +++ b/codex-rs/core/src/session/review.rs @@ -139,7 +139,7 @@ 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())), + turn_start_transcript_inputs: Arc::new(std::sync::Mutex::new(Vec::new())), transcript_serialization_lock: Arc::new(Semaphore::new(1)), }; diff --git a/codex-rs/core/src/session/tests.rs b/codex-rs/core/src/session/tests.rs index bbe21c3bc8..049a7a024d 100644 --- a/codex-rs/core/src/session/tests.rs +++ b/codex-rs/core/src/session/tests.rs @@ -55,6 +55,8 @@ use tracing::Span; use crate::RolloutRecorderParams; use crate::rollout::policy::EventPersistenceMode; use crate::rollout::recorder::RolloutRecorder; +use crate::session::turn_context::TurnStartTranscriptInput; +use crate::session::turn_context::TurnStartUserPromptSubmitOutcome; use crate::state::TaskKind; use crate::tasks::SessionTask; use crate::tasks::SessionTaskContext; @@ -5708,6 +5710,56 @@ async fn abort_regular_task_records_prompt_before_interrupt_marker() { assert_eq!(history.raw_items(), expected.as_slice()); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[test_log::test] +async fn abort_regular_task_replays_context_without_replaying_prompt() { + let (sess, tc, _rx) = make_session_and_context_with_rx().await; + let input = vec![UserInput::Text { + text: "hello".to_string(), + text_elements: Vec::new(), + }]; + let response_item = ResponseItem::Message { + id: None, + role: "user".to_string(), + content: vec![ContentItem::InputText { + text: "hello".to_string(), + }], + end_turn: None, + phase: None, + }; + sess.record_conversation_items(tc.as_ref(), std::slice::from_ref(&response_item)) + .await; + tc.lock_turn_start_transcript_inputs() + .push(TurnStartTranscriptInput { + input, + user_prompt_submit_outcome: Some(TurnStartUserPromptSubmitOutcome { + should_stop: false, + additional_contexts: vec!["hook context".to_string()], + }), + user_prompt_recorded: true, + }); + sess.spawn_task( + Arc::clone(&tc), + Vec::new(), + NeverEndingTask { + kind: TaskKind::Regular, + listen_to_cancellation_token: true, + }, + ) + .await; + + sess.abort_all_tasks(TurnAbortReason::Interrupted).await; + + let history = sess.clone_history().await; + let expected = vec![ + response_item, + DeveloperInstructions::new("hook context".to_string()).into(), + crate::tasks::interrupted_turn_history_marker(), + ]; + assert_eq!(history.raw_items(), expected.as_slice()); + assert!(tc.lock_turn_start_transcript_inputs().is_empty()); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[test_log::test] async fn abort_non_regular_task_keeps_pending_session_start_source() { diff --git a/codex-rs/core/src/session/turn_context.rs b/codex-rs/core/src/session/turn_context.rs index c2cdfa8d43..a6c4c6bc18 100644 --- a/codex-rs/core/src/session/turn_context.rs +++ b/codex-rs/core/src/session/turn_context.rs @@ -24,6 +24,19 @@ impl TurnSkillsContext { } } +#[derive(Clone, Debug, PartialEq)] +pub(crate) struct TurnStartTranscriptInput { + pub(crate) input: Vec, + pub(crate) user_prompt_submit_outcome: Option, + pub(crate) user_prompt_recorded: bool, +} + +#[derive(Clone, Debug, PartialEq)] +pub(crate) struct TurnStartUserPromptSubmitOutcome { + pub(crate) should_stop: bool, + pub(crate) additional_contexts: Vec, +} + /// The context needed for a single turn of the thread. #[derive(Debug)] pub(crate) struct TurnContext { @@ -71,10 +84,18 @@ 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) turn_start_transcript_inputs: Arc>>, pub(crate) transcript_serialization_lock: Arc, } impl TurnContext { + pub(crate) fn lock_turn_start_transcript_inputs( + &self, + ) -> std::sync::MutexGuard<'_, Vec> { + self.turn_start_transcript_inputs + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + } + pub(crate) fn model_context_window(&self) -> Option { let effective_context_window_percent = self.model_info.effective_context_window_percent; self.model_info @@ -447,7 +468,7 @@ 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())), + turn_start_transcript_inputs: Arc::new(std::sync::Mutex::new(Vec::new())), transcript_serialization_lock: Arc::new(Semaphore::new(1)), } } diff --git a/codex-rs/core/src/tasks/mod.rs b/codex-rs/core/src/tasks/mod.rs index b7f0c969ca..957f3521ee 100644 --- a/codex-rs/core/src/tasks/mod.rs +++ b/codex-rs/core/src/tasks/mod.rs @@ -28,6 +28,7 @@ use crate::hook_runtime::record_additional_contexts; use crate::hook_runtime::record_pending_input; use crate::session::session::Session; use crate::session::turn_context::TurnContext; +use crate::session::turn_context::TurnStartTranscriptInput; use crate::state::ActiveTurn; use crate::state::RunningTask; use crate::state::TaskKind; @@ -302,10 +303,12 @@ impl Session { if task_kind == TaskKind::Regular && !input.is_empty() { turn_context - .turn_start_transcript_inputs - .lock() - .await - .push(input.clone()); + .lock_turn_start_transcript_inputs() + .push(TurnStartTranscriptInput { + input: input.clone(), + user_prompt_submit_outcome: None, + user_prompt_recorded: false, + }); } let mut active = self.active_turn.lock().await; let turn = active.get_or_insert_with(ActiveTurn::default); @@ -637,12 +640,10 @@ impl Session { task.handle.abort(); if reason == TurnAbortReason::Interrupted && task.kind == TaskKind::Regular { - let has_turn_start_transcript_input = !task - .turn_context - .turn_start_transcript_inputs - .lock() - .await - .is_empty(); + let has_turn_start_transcript_input = { + let inputs = task.turn_context.lock_turn_start_transcript_inputs(); + !inputs.is_empty() + }; if has_turn_start_transcript_input { let _ = drain_turn_start_transcript_inputs(self, &task.turn_context).await; }