mirror of
https://github.com/openai/codex.git
synced 2026-09-04 15:08:45 +00:00
Bound rollout replay and rollback history
This commit is contained in:
@@ -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<_>>();
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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");
|
||||
|
||||
Reference in New Issue
Block a user