From 4f4e4b8e76a23a88cf5e72a42ec172c5bfbcf167 Mon Sep 17 00:00:00 2001 From: Friel Date: Sat, 13 Jun 2026 00:49:50 +0000 Subject: [PATCH] Bound rollout replay and rollback history --- codex-rs/core/src/session/handlers.rs | 26 +++++- codex-rs/core/src/session/tests.rs | 82 +++++++++++++++++++ .../core/src/thread_rollout_truncation.rs | 15 ++++ .../src/thread_rollout_truncation_tests.rs | 45 ++++++++++ 4 files changed, 166 insertions(+), 2 deletions(-) diff --git a/codex-rs/core/src/session/handlers.rs b/codex-rs/core/src/session/handlers.rs index d4525cf45d..8498d36e23 100644 --- a/codex-rs/core/src/session/handlers.rs +++ b/codex-rs/core/src/session/handlers.rs @@ -21,6 +21,7 @@ use crate::tasks::CompactTask; use crate::tasks::UserShellCommandMode; use crate::tasks::UserShellCommandTask; use crate::tasks::execute_user_shell_command; +use crate::thread_rollout_truncation::materialize_rollout_items_for_complete_history; use codex_protocol::models::ContentItem; use codex_protocol::models::ResponseInputItem; use codex_protocol::models::ResponseItem; @@ -515,10 +516,31 @@ pub async fn thread_rollback(sess: &Arc, sub_id: String, num_turns: u32 } }; + let materialized_history = match materialize_rollout_items_for_complete_history( + turn_context.config.codex_home.as_path(), + &stored_history.items, + ) + .await + { + Ok(history) => history, + Err(err) => { + sess.send_event_raw(Event { + id: turn_context.sub_id.clone(), + msg: EventMsg::Error(ErrorEvent { + message: format!( + "failed to materialize thread history for rollback replay: {err}" + ), + codex_error_info: Some(CodexErrorInfo::ThreadRollbackFailed), + }), + }) + .await; + return; + } + }; + let rollback_event = ThreadRolledBackEvent { num_turns }; let rollback_msg = EventMsg::ThreadRolledBack(rollback_event.clone()); - let replay_items = stored_history - .items + let replay_items = materialized_history .into_iter() .chain(std::iter::once(RolloutItem::EventMsg(rollback_msg.clone()))) .collect::>(); diff --git a/codex-rs/core/src/session/tests.rs b/codex-rs/core/src/session/tests.rs index a04290c5d1..01135e50a1 100644 --- a/codex-rs/core/src/session/tests.rs +++ b/codex-rs/core/src/session/tests.rs @@ -32,6 +32,7 @@ use codex_models_manager::model_info; use codex_models_manager::test_support::construct_model_info_offline_for_tests; use codex_models_manager::test_support::get_model_offline_for_tests; use codex_protocol::AgentPath; +use codex_protocol::SegmentId; use codex_protocol::SessionId; use codex_protocol::ThreadId; use codex_protocol::config_types::SERVICE_TIER_DEFAULT_REQUEST_VALUE; @@ -120,6 +121,8 @@ use codex_protocol::protocol::RealtimeVoice; use codex_protocol::protocol::RealtimeVoicesList; use codex_protocol::protocol::ResumedHistory; use codex_protocol::protocol::RolloutItem; +use codex_protocol::protocol::RolloutLine; +use codex_protocol::protocol::RolloutReferenceItem; use codex_protocol::protocol::SessionMeta; use codex_protocol::protocol::SessionMetaLine; use codex_protocol::protocol::SkillScope; @@ -2799,6 +2802,85 @@ async fn thread_rollback_drops_last_turn_from_history() { })); } +#[tokio::test] +async fn thread_rollback_materializes_referenced_predecessor_history() { + let (mut sess, tc, rx) = make_session_and_context_with_rx().await; + let rollout_path = attach_thread_persistence( + Arc::get_mut(&mut sess).expect("session should not have additional references"), + ) + .await; + let predecessor_segment_id = SegmentId::new(); + let predecessor_path = rollout_path.with_file_name("rollback-predecessor.jsonl"); + let turn_1 = vec![ + user_message("turn 1 user"), + assistant_message("turn 1 assistant"), + ]; + let predecessor_items = [ + RolloutItem::SessionMeta(SessionMetaLine { + meta: SessionMeta { + id: sess.thread_id, + segment_id: Some(predecessor_segment_id), + timestamp: "2026-06-13T00:00:00.000Z".to_string(), + cwd: tc.config.cwd.to_path_buf(), + originator: "test".to_string(), + cli_version: "0.0.0".to_string(), + ..SessionMeta::default() + }, + git: None, + }), + RolloutItem::ResponseItem(turn_1[0].clone()), + RolloutItem::ResponseItem(turn_1[1].clone()), + ]; + let predecessor_jsonl = predecessor_items + .into_iter() + .map(|item| { + serde_json::to_string(&RolloutLine { + timestamp: "2026-06-13T00:00:00.000Z".to_string(), + item, + }) + .expect("serialize predecessor rollout line") + }) + .collect::>() + .join("\n") + + "\n"; + tokio::fs::write(&predecessor_path, predecessor_jsonl) + .await + .expect("write predecessor rollout"); + + let turn_2 = vec![ + user_message("turn 2 user"), + assistant_message("turn 2 assistant"), + ]; + sess.persist_rollout_items(&[ + RolloutItem::RolloutReference(RolloutReferenceItem { + rollout_path: predecessor_path, + thread_id: Some(sess.thread_id), + rollout_timestamp: None, + segment_id: Some(predecessor_segment_id), + max_depth: 2, + nth_user_message: None, + compacted_replacement_history_filter_texts: None, + }), + RolloutItem::ResponseItem(turn_2[0].clone()), + RolloutItem::ResponseItem(turn_2[1].clone()), + ]) + .await; + sess.flush_rollout() + .await + .expect("flush referenced history"); + sess.replace_history( + turn_1.iter().chain(&turn_2).cloned().collect(), + Some(tc.to_turn_context_item()), + ) + .await; + + handlers::thread_rollback(&sess, "sub-1".to_string(), /*num_turns*/ 1).await; + + let rollback_event = wait_for_thread_rolled_back(&rx).await; + assert_eq!(rollback_event.num_turns, 1); + assert_eq!(sess.clone_history().await.raw_items(), turn_1); +} + #[tokio::test] async fn thread_rollback_clears_history_when_num_turns_exceeds_existing_turns() { let (mut sess, tc, rx) = make_session_and_context_with_rx().await; diff --git a/codex-rs/core/src/thread_rollout_truncation.rs b/codex-rs/core/src/thread_rollout_truncation.rs index 8d5e893be4..ec5f26bd5d 100644 --- a/codex-rs/core/src/thread_rollout_truncation.rs +++ b/codex-rs/core/src/thread_rollout_truncation.rs @@ -174,6 +174,8 @@ enum RolloutMaterialization { CompleteHistory, } +const MAX_MODEL_REPLAY_REFERENCE_DEPTH: usize = 8; + #[derive(Clone, Debug, Eq, Hash, PartialEq)] enum RolloutReferenceIdentity { Segment(codex_protocol::SegmentId), @@ -237,10 +239,12 @@ async fn materialize_rollout_items( enum Work { Items { rollout_items: Vec, + reference_depth: usize, remaining_segment_depth: Option, }, Item { item: Box, + reference_depth: usize, remaining_segment_depth: Option, }, TruncateSuffix { @@ -254,23 +258,27 @@ async fn materialize_rollout_items( let mut active_references = HashSet::new(); let mut work = vec![Work::Items { rollout_items: rollout_items.to_vec(), + reference_depth: 0, remaining_segment_depth: None, }]; while let Some(next) = work.pop() { match next { Work::Items { rollout_items, + reference_depth, remaining_segment_depth, } => { for item in rollout_items.into_iter().rev() { work.push(Work::Item { item: Box::new(item), + reference_depth, remaining_segment_depth, }); } } Work::Item { item, + reference_depth, remaining_segment_depth, } => { let reference = match *item { @@ -280,6 +288,12 @@ async fn materialize_rollout_items( continue; } }; + if matches!(materialization, RolloutMaterialization::ModelReplay) + && reference_depth >= MAX_MODEL_REPLAY_REFERENCE_DEPTH + { + warn!("rollout reference materialization reached hard model-replay depth cap"); + continue; + } let has_prefix_truncation = reference.nth_user_message.is_some(); let next_remaining_segment_depth = match materialization { RolloutMaterialization::CompleteHistory => None, @@ -358,6 +372,7 @@ async fn materialize_rollout_items( } work.push(Work::Items { rollout_items: reference_items, + reference_depth: reference_depth.saturating_add(1), remaining_segment_depth: next_remaining_segment_depth, }); } diff --git a/codex-rs/core/src/thread_rollout_truncation_tests.rs b/codex-rs/core/src/thread_rollout_truncation_tests.rs index bc1d5b6b48..e3b572c684 100644 --- a/codex-rs/core/src/thread_rollout_truncation_tests.rs +++ b/codex-rs/core/src/thread_rollout_truncation_tests.rs @@ -593,6 +593,51 @@ async fn outer_reference_depth_bounds_descendant_references() { assert!(serialized.contains("middle request")); } +#[tokio::test] +async fn model_replay_clamps_serialized_reference_depth() { + let temp = tempfile::tempdir().expect("tempdir"); + let thread_id = ThreadId::new(); + let mut previous_segment = None; + let mut newest_items = Vec::new(); + + for index in 0..=MAX_MODEL_REPLAY_REFERENCE_DEPTH + 1 { + let segment_id = SegmentId::new(); + let path = temp.path().join(format!("segment-{index}.jsonl")); + let mut items = vec![session_meta_item(thread_id, segment_id)]; + if let Some((previous_path, previous_segment_id)) = previous_segment { + items.push(RolloutItem::RolloutReference(RolloutReferenceItem { + rollout_path: previous_path, + thread_id: Some(thread_id), + rollout_timestamp: None, + segment_id: Some(previous_segment_id), + max_depth: usize::MAX, + nth_user_message: None, + compacted_replacement_history_filter_texts: None, + })); + } + items.push(RolloutItem::ResponseItem(user_msg(&format!( + "segment {index}" + )))); + write_rollout(&path, &items).await; + previous_segment = Some((path, segment_id)); + newest_items = items; + } + + let model_history = + materialize_rollout_items_for_model_replay(temp.path(), &newest_items).await; + let model_json = serde_json::to_string(&model_history).expect("serialize model history"); + assert!(!model_json.contains("\"text\":\"segment 0\"")); + assert!(model_json.contains("\"text\":\"segment 1\"")); + + let complete_history = + materialize_rollout_items_for_complete_history(temp.path(), &newest_items) + .await + .expect("materialize complete history"); + let complete_json = + serde_json::to_string(&complete_history).expect("serialize complete history"); + assert!(complete_json.contains("\"text\":\"segment 0\"")); +} + #[tokio::test] async fn complete_history_reports_unresolvable_references() { let temp = tempfile::tempdir().expect("tempdir");