codex: harden interrupt prompt drain

This commit is contained in:
Joe Gershenson
2026-04-21 00:40:08 -07:00
parent 051accdb6a
commit dd7aaab825
5 changed files with 144 additions and 37 deletions

View File

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

View File

@@ -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)),
};

View File

@@ -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() {

View File

@@ -24,6 +24,19 @@ impl TurnSkillsContext {
}
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct TurnStartTranscriptInput {
pub(crate) input: Vec<UserInput>,
pub(crate) user_prompt_submit_outcome: Option<TurnStartUserPromptSubmitOutcome>,
pub(crate) user_prompt_recorded: bool,
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct TurnStartUserPromptSubmitOutcome {
pub(crate) should_stop: bool,
pub(crate) additional_contexts: Vec<String>,
}
/// 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<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) turn_start_transcript_inputs: Arc<std::sync::Mutex<Vec<TurnStartTranscriptInput>>>,
pub(crate) transcript_serialization_lock: Arc<Semaphore>,
}
impl TurnContext {
pub(crate) fn lock_turn_start_transcript_inputs(
&self,
) -> std::sync::MutexGuard<'_, Vec<TurnStartTranscriptInput>> {
self.turn_start_transcript_inputs
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
pub(crate) fn model_context_window(&self) -> Option<i64> {
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)),
}
}

View File

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