Fix interrupt prompt history race

This commit is contained in:
Joe Gershenson
2026-04-20 23:49:29 -07:00
parent ab26554a3a
commit 553376ffcb
6 changed files with 102 additions and 81 deletions

View File

@@ -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<Session>,
turn_context: &Arc<TurnContext>,
) -> 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<Session>,
turn_context: &Arc<TurnContext>,

View File

@@ -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.

View File

@@ -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::<Option<String>>().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 `<turn_aborted>` 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)]

View File

@@ -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::<Vec<_>>();
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

View File

@@ -71,6 +71,8 @@ pub(crate) struct TurnContext {
pub(crate) turn_metadata_state: Arc<TurnMetadataState>,
pub(crate) turn_skills: TurnSkillsContext,
pub(crate) turn_timing_state: Arc<TurnTimingState>,
pub(crate) turn_start_transcript_inputs: Arc<Mutex<Vec<Vec<UserInput>>>>,
pub(crate) transcript_serialization_lock: Arc<Mutex<()>>,
}
impl TurnContext {
pub(crate) fn model_context_window(&self) -> Option<i64> {
@@ -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(())),
}
}

View File

@@ -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;