Bound rollout replay and rollback history

This commit is contained in:
Friel
2026-06-13 00:49:50 +00:00
parent 2852e3b9d1
commit 4f4e4b8e76
4 changed files with 166 additions and 2 deletions

View File

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

View File

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

View File

@@ -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<RolloutItem>,
reference_depth: usize,
remaining_segment_depth: Option<usize>,
},
Item {
item: Box<RolloutItem>,
reference_depth: usize,
remaining_segment_depth: Option<usize>,
},
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,
});
}

View File

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