From cb98305eab96eefb6a7d2f025a4e42bec971ba20 Mon Sep 17 00:00:00 2001 From: Charles Cunningham Date: Fri, 20 Mar 2026 12:43:10 -0700 Subject: [PATCH] address review --- codex-rs/core/src/thread_manager.rs | 38 +++++----- codex-rs/core/src/thread_manager_tests.rs | 88 +++++++++++++++++++++-- 2 files changed, 102 insertions(+), 24 deletions(-) diff --git a/codex-rs/core/src/thread_manager.rs b/codex-rs/core/src/thread_manager.rs index 07cbea9a54..ea04ab4708 100644 --- a/codex-rs/core/src/thread_manager.rs +++ b/codex-rs/core/src/thread_manager.rs @@ -25,6 +25,7 @@ use crate::rollout::truncation; use crate::shell_snapshot::ShellSnapshot; use crate::skills::SkillsManager; use crate::tasks::interrupted_turn_history_marker; +use codex_app_server_protocol::ThreadHistoryBuilder; use codex_protocol::ThreadId; use codex_protocol::config_types::CollaborationModeMask; #[cfg(test)] @@ -151,10 +152,10 @@ pub enum ForkSnapshot { /// Fork the current persisted history as if the source thread had been /// interrupted now. /// - /// If the source thread is mid-turn, this appends the same - /// `` marker produced by a real interrupt. If the source - /// thread is already at a turn boundary, this returns the current persisted - /// history unchanged. + /// If the persisted snapshot ends mid-turn, this appends the same + /// `` marker produced by a real interrupt. If the snapshot is + /// already at a turn boundary, this returns the current persisted history + /// unchanged. Interrupted, } @@ -586,21 +587,10 @@ impl ThreadManager { parent_trace: Option, ) -> CodexResult { let history = RolloutRecorder::get_rollout_history(&path).await?; - let source_thread = { - let threads = self.state.threads.read().await; - threads - .values() - .find(|thread| thread.rollout_path().as_ref() == Some(&path)) - .cloned() - }; - let source_mid_turn = if let Some(thread) = source_thread { - thread.codex.session.active_turn.lock().await.is_some() - } else { - false - }; + let snapshot_mid_turn = snapshot_ends_mid_turn(&history); let history = match snapshot { ForkSnapshot::TruncateBeforeNthUserMessage(nth_user_message) => { - truncate_before_nth_user_message(history, nth_user_message, source_mid_turn) + truncate_before_nth_user_message(history, nth_user_message, snapshot_mid_turn) } ForkSnapshot::Interrupted => { let history = match history { @@ -608,7 +598,7 @@ impl ThreadManager { InitialHistory::Forked(history) => InitialHistory::Forked(history), InitialHistory::Resumed(resumed) => InitialHistory::Forked(resumed.history), }; - if source_mid_turn { + if snapshot_mid_turn { inject_interrupted_marker(history) } else { history @@ -906,11 +896,11 @@ impl ThreadManagerState { fn truncate_before_nth_user_message( history: InitialHistory, n: usize, - source_mid_turn: bool, + snapshot_mid_turn: bool, ) -> InitialHistory { let items: Vec = history.get_rollout_items(); let user_positions = truncation::user_message_positions_in_rollout(&items); - let rolled = if source_mid_turn && n >= user_positions.len() { + let rolled = if snapshot_mid_turn && n >= user_positions.len() { if let Some(cut_idx) = user_positions.last().copied() { items[..cut_idx].to_vec() } else { @@ -927,6 +917,14 @@ fn truncate_before_nth_user_message( } } +fn snapshot_ends_mid_turn(history: &InitialHistory) -> bool { + let mut builder = ThreadHistoryBuilder::new(); + for item in history.get_rollout_items() { + builder.handle_rollout_item(&item); + } + builder.has_active_turn() +} + /// 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. diff --git a/codex-rs/core/src/thread_manager_tests.rs b/codex-rs/core/src/thread_manager_tests.rs index fb83a1f8dd..52e56e0549 100644 --- a/codex-rs/core/src/thread_manager_tests.rs +++ b/codex-rs/core/src/thread_manager_tests.rs @@ -3,6 +3,7 @@ use crate::codex::make_session_and_context; use crate::config::test_config; use crate::models_manager::collaboration_mode_presets::CollaborationModesConfig; use crate::models_manager::manager::RefreshStrategy; +use crate::rollout::RolloutRecorder; use crate::tasks::interrupted_turn_history_marker; use codex_protocol::models::ContentItem; use codex_protocol::models::ReasoningItemReasoningSummary; @@ -71,7 +72,7 @@ fn truncates_before_requested_user_message() { let truncated = truncate_before_nth_user_message( InitialHistory::Forked(initial), 1, - /*source_mid_turn*/ false, + /*snapshot_mid_turn*/ false, ); let got_items = truncated.get_rollout_items(); let expected_items = vec![ @@ -92,7 +93,7 @@ fn truncates_before_requested_user_message() { let truncated2 = truncate_before_nth_user_message( InitialHistory::Forked(initial2.clone()), 2, - /*source_mid_turn*/ false, + /*snapshot_mid_turn*/ false, ); assert_eq!( serde_json::to_value(truncated2.get_rollout_items()).unwrap(), @@ -112,7 +113,7 @@ fn out_of_range_truncation_drops_only_unfinished_suffix_mid_turn() { let truncated = truncate_before_nth_user_message( InitialHistory::Forked(items.clone()), usize::MAX, - /*source_mid_turn*/ true, + /*snapshot_mid_turn*/ true, ); assert_eq!( @@ -139,7 +140,7 @@ async fn ignores_session_prefix_messages_when_truncating() { let truncated = truncate_before_nth_user_message( InitialHistory::Forked(rollout_items), 1, - /*source_mid_turn*/ false, + /*snapshot_mid_turn*/ false, ); let got_items = truncated.get_rollout_items(); @@ -245,3 +246,82 @@ fn interrupted_fork_snapshot_appends_interrupt_marker() { .expect("serialize expected interrupted empty history"), ); } + +#[tokio::test] +async fn interrupted_fork_snapshot_uses_persisted_mid_turn_history_without_live_source() { + let temp_dir = tempdir().expect("tempdir"); + let mut config = test_config(); + config.codex_home = temp_dir.path().join("codex-home"); + config.cwd = config.codex_home.clone(); + std::fs::create_dir_all(&config.codex_home).expect("create codex home"); + + let auth_manager = + AuthManager::from_auth_for_testing(CodexAuth::create_dummy_chatgpt_auth_for_testing()); + let manager = ThreadManager::new( + &config, + auth_manager.clone(), + SessionSource::Exec, + CollaborationModesConfig::default(), + ); + + let source = manager + .resume_thread_with_history( + config.clone(), + InitialHistory::Forked(vec![ + RolloutItem::ResponseItem(user_msg("hello")), + RolloutItem::ResponseItem(assistant_msg("partial")), + ]), + auth_manager, + /*persist_extended_history*/ false, + /*parent_trace*/ None, + ) + .await + .expect("create source thread from partial history"); + let source_path = source + .thread + .rollout_path() + .expect("source rollout path should exist"); + let source_history = RolloutRecorder::get_rollout_history(&source_path) + .await + .expect("read source rollout history"); + assert!(snapshot_ends_mid_turn(&source_history)); + manager.remove_thread(&source.thread_id).await; + + let forked = manager + .fork_thread( + ForkSnapshot::Interrupted, + config, + source_path, + /*persist_extended_history*/ false, + /*parent_trace*/ None, + ) + .await + .expect("fork interrupted snapshot"); + let forked_path = forked + .thread + .rollout_path() + .expect("forked rollout path should exist"); + let history = RolloutRecorder::get_rollout_history(&forked_path) + .await + .expect("read forked rollout history"); + + let forked_rollout_items: Vec<_> = history + .get_rollout_items() + .into_iter() + .filter(|item| !matches!(item, RolloutItem::SessionMeta(_))) + .collect(); + let interrupted_marker_json = serde_json::to_value(RolloutItem::ResponseItem( + interrupted_turn_history_marker(), + )) + .expect("serialize interrupted marker"); + assert_eq!( + forked_rollout_items + .iter() + .filter(|item| { + serde_json::to_value(item).expect("serialize forked rollout item") + == interrupted_marker_json + }) + .count(), + 1, + ); +}