From 96fe0ed92ddd19bb7e4e09c3315ea10a0183f302 Mon Sep 17 00:00:00 2001 From: Charles Cunningham Date: Mon, 23 Mar 2026 10:04:27 -0700 Subject: [PATCH] Persist abort boundary in interrupted fork snapshots Co-authored-by: Codex --- codex-rs/core/src/thread_manager.rs | 38 +++++++++--- codex-rs/core/src/thread_manager_tests.rs | 73 ++++++++++++++++++++--- 2 files changed, 96 insertions(+), 15 deletions(-) diff --git a/codex-rs/core/src/thread_manager.rs b/codex-rs/core/src/thread_manager.rs index ecf8bd7c13..fcea98468a 100644 --- a/codex-rs/core/src/thread_manager.rs +++ b/codex-rs/core/src/thread_manager.rs @@ -36,6 +36,8 @@ use codex_protocol::protocol::McpServerRefreshConfig; use codex_protocol::protocol::Op; use codex_protocol::protocol::RolloutItem; use codex_protocol::protocol::SessionSource; +use codex_protocol::protocol::TurnAbortReason; +use codex_protocol::protocol::TurnAbortedEvent; use codex_protocol::protocol::W3cTraceContext; use futures::StreamExt; use futures::stream::FuturesUnordered; @@ -599,7 +601,7 @@ impl ThreadManager { InitialHistory::Resumed(resumed) => InitialHistory::Forked(resumed.history), }; if snapshot_mid_turn { - inject_interrupted_marker(history) + append_interrupted_boundary(history) } else { history } @@ -945,22 +947,42 @@ fn snapshot_ends_mid_turn(history: &InitialHistory) -> bool { }) } -/// Append the same model-visible interruption marker used by the live interrupt -/// path to an existing fork snapshot after the source thread has been confirmed -/// to be mid-turn. -fn inject_interrupted_marker(history: InitialHistory) -> InitialHistory { +/// Append the same persisted interrupt boundary used by the live interrupt path +/// to an existing fork snapshot after the source thread has been confirmed to +/// be mid-turn. +fn append_interrupted_boundary(history: InitialHistory) -> InitialHistory { + let turn_id = { + let rollout_items = history.get_rollout_items(); + let mut builder = ThreadHistoryBuilder::new(); + for item in &rollout_items { + builder.handle_rollout_item(item); + } + if builder.has_active_turn() { + builder.active_turn_snapshot().map(|turn| turn.id) + } else { + None + } + }; + let aborted_event = RolloutItem::EventMsg(EventMsg::TurnAborted(TurnAbortedEvent { + turn_id, + reason: TurnAbortReason::Interrupted, + })); + match history { - InitialHistory::New => InitialHistory::Forked(vec![RolloutItem::ResponseItem( - interrupted_turn_history_marker(), - )]), + InitialHistory::New => InitialHistory::Forked(vec![ + RolloutItem::ResponseItem(interrupted_turn_history_marker()), + aborted_event, + ]), InitialHistory::Forked(mut history) => { history.push(RolloutItem::ResponseItem(interrupted_turn_history_marker())); + history.push(aborted_event); InitialHistory::Forked(history) } InitialHistory::Resumed(mut resumed) => { resumed .history .push(RolloutItem::ResponseItem(interrupted_turn_history_marker())); + resumed.history.push(aborted_event); InitialHistory::Forked(resumed.history) } } diff --git a/codex-rs/core/src/thread_manager_tests.rs b/codex-rs/core/src/thread_manager_tests.rs index 2bc572ab6a..668fd525e5 100644 --- a/codex-rs/core/src/thread_manager_tests.rs +++ b/codex-rs/core/src/thread_manager_tests.rs @@ -224,25 +224,33 @@ async fn new_uses_configured_openai_provider_for_model_refresh() { } #[test] -fn interrupted_fork_snapshot_appends_interrupt_marker() { +fn interrupted_fork_snapshot_appends_interrupt_boundary() { let committed_history = InitialHistory::Forked(vec![RolloutItem::ResponseItem(user_msg("hello"))]); assert_eq!( - serde_json::to_value(inject_interrupted_marker(committed_history).get_rollout_items()) + serde_json::to_value(append_interrupted_boundary(committed_history).get_rollout_items()) .expect("serialize interrupted fork history"), serde_json::to_value(vec![ RolloutItem::ResponseItem(user_msg("hello")), RolloutItem::ResponseItem(interrupted_turn_history_marker()), + RolloutItem::EventMsg(EventMsg::TurnAborted(TurnAbortedEvent { + turn_id: None, + reason: TurnAbortReason::Interrupted, + })), ]) .expect("serialize expected interrupted fork history"), ); assert_eq!( - serde_json::to_value(inject_interrupted_marker(InitialHistory::New).get_rollout_items()) + serde_json::to_value(append_interrupted_boundary(InitialHistory::New).get_rollout_items()) .expect("serialize interrupted empty fork history"), - serde_json::to_value(vec![RolloutItem::ResponseItem( - interrupted_turn_history_marker() - )]) + serde_json::to_value(vec![ + RolloutItem::ResponseItem(interrupted_turn_history_marker()), + RolloutItem::EventMsg(EventMsg::TurnAborted(TurnAbortedEvent { + turn_id: None, + reason: TurnAbortReason::Interrupted, + })), + ]) .expect("serialize expected interrupted empty history"), ); } @@ -290,7 +298,7 @@ async fn interrupted_fork_snapshot_uses_persisted_mid_turn_history_without_live_ let forked = manager .fork_thread( ForkSnapshot::Interrupted, - config, + config.clone(), source_path, /*persist_extended_history*/ false, /*parent_trace*/ None, @@ -304,6 +312,7 @@ async fn interrupted_fork_snapshot_uses_persisted_mid_turn_history_without_live_ let history = RolloutRecorder::get_rollout_history(&forked_path) .await .expect("read forked rollout history"); + assert!(!snapshot_ends_mid_turn(&history)); let forked_rollout_items: Vec<_> = history .get_rollout_items() @@ -323,4 +332,54 @@ async fn interrupted_fork_snapshot_uses_persisted_mid_turn_history_without_live_ .count(), 1, ); + + manager.remove_thread(&forked.thread_id).await; + let reforked = manager + .fork_thread( + ForkSnapshot::Interrupted, + config, + forked_path, + /*persist_extended_history*/ false, + /*parent_trace*/ None, + ) + .await + .expect("re-fork interrupted snapshot"); + let reforked_path = reforked + .thread + .rollout_path() + .expect("re-forked rollout path should exist"); + let reforked_history = RolloutRecorder::get_rollout_history(&reforked_path) + .await + .expect("read re-forked rollout history"); + let reforked_rollout_items: Vec<_> = reforked_history + .get_rollout_items() + .into_iter() + .filter(|item| !matches!(item, RolloutItem::SessionMeta(_))) + .collect(); + + assert_eq!( + reforked_rollout_items + .iter() + .filter(|item| { + serde_json::to_value(item).expect("serialize re-forked rollout item") + == interrupted_marker_json + }) + .count(), + 1, + ); + assert_eq!( + reforked_rollout_items + .iter() + .filter(|item| { + matches!( + item, + RolloutItem::EventMsg(EventMsg::TurnAborted(TurnAbortedEvent { + reason: TurnAbortReason::Interrupted, + .. + })) + ) + }) + .count(), + 1, + ); }