mirror of
https://github.com/pchuan98/codex.git
synced 2026-07-01 00:31:56 +08:00
Replay thread rollback from rollout history (#13615)
- Replay thread rollback from the persisted rollout history instead of truncating in-memory state.\n- Add rollback coverage, including rollback-behind-compaction snapshot coverage.
This commit is contained in:
committed by
GitHub
Unverified
parent
6cf0ed4e79
commit
629cb15bc6
+386
-74
@@ -3942,6 +3942,7 @@ mod handlers {
|
||||
use crate::mcp::auth::compute_auth_statuses;
|
||||
use crate::mcp::collect_mcp_snapshot_from_manager;
|
||||
use crate::review_prompts::resolve_review_request;
|
||||
use crate::rollout::RolloutRecorder;
|
||||
use crate::rollout::session_index;
|
||||
use crate::tasks::CompactTask;
|
||||
use crate::tasks::UndoTask;
|
||||
@@ -3964,6 +3965,7 @@ mod handlers {
|
||||
use codex_protocol::protocol::RemoteSkillSummary;
|
||||
use codex_protocol::protocol::ReviewDecision;
|
||||
use codex_protocol::protocol::ReviewRequest;
|
||||
use codex_protocol::protocol::RolloutItem;
|
||||
use codex_protocol::protocol::SkillsListEntry;
|
||||
use codex_protocol::protocol::ThreadNameUpdatedEvent;
|
||||
use codex_protocol::protocol::ThreadRolledBackEvent;
|
||||
@@ -4538,25 +4540,86 @@ mod handlers {
|
||||
}
|
||||
|
||||
let turn_context = sess.new_default_turn_with_sub_id(sub_id).await;
|
||||
let rollout_path = {
|
||||
let recorder = {
|
||||
let guard = sess.services.rollout.lock().await;
|
||||
guard.clone()
|
||||
};
|
||||
let Some(recorder) = recorder else {
|
||||
sess.send_event_raw(Event {
|
||||
id: turn_context.sub_id.clone(),
|
||||
msg: EventMsg::Error(ErrorEvent {
|
||||
message: "thread rollback requires a persisted rollout path".to_string(),
|
||||
codex_error_info: Some(CodexErrorInfo::ThreadRollbackFailed),
|
||||
}),
|
||||
})
|
||||
.await;
|
||||
return;
|
||||
};
|
||||
recorder.rollout_path().to_path_buf()
|
||||
};
|
||||
if let Some(recorder) = {
|
||||
let guard = sess.services.rollout.lock().await;
|
||||
guard.clone()
|
||||
} && let Err(err) = recorder.flush().await
|
||||
{
|
||||
sess.send_event_raw(Event {
|
||||
id: turn_context.sub_id.clone(),
|
||||
msg: EventMsg::Error(ErrorEvent {
|
||||
message: format!(
|
||||
"failed to flush rollout `{}` for rollback replay: {err}",
|
||||
rollout_path.display()
|
||||
),
|
||||
codex_error_info: Some(CodexErrorInfo::ThreadRollbackFailed),
|
||||
}),
|
||||
})
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
|
||||
let mut history = sess.clone_history().await;
|
||||
// TODO(ccunningham): Fix rollback/backtracking baseline handling.
|
||||
// We clear `reference_context_item` here, but should restore the
|
||||
// post-rollback baseline from the surviving history/rollout instead.
|
||||
// Truncating history should also invalidate/recompute `previous_turn_settings`
|
||||
// so the next regular turn replays any dropped model/realtime
|
||||
// instructions.
|
||||
history.drop_last_n_user_turns(num_turns);
|
||||
let initial_history =
|
||||
match RolloutRecorder::get_rollout_history(rollout_path.as_path()).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 load rollout `{}` for rollback replay: {err}",
|
||||
rollout_path.display()
|
||||
),
|
||||
codex_error_info: Some(CodexErrorInfo::ThreadRollbackFailed),
|
||||
}),
|
||||
})
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
// Replace with the raw items. We don't want to replace with a normalized
|
||||
// version of the history.
|
||||
sess.replace_history(history.raw_items().to_vec(), None)
|
||||
let rollback_event = ThreadRolledBackEvent { num_turns };
|
||||
let replay_items = initial_history
|
||||
.get_rollout_items()
|
||||
.into_iter()
|
||||
.chain(std::iter::once(RolloutItem::EventMsg(
|
||||
EventMsg::ThreadRolledBack(rollback_event.clone()),
|
||||
)))
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
let reconstructed = sess
|
||||
.reconstruct_history_from_rollout(turn_context.as_ref(), replay_items.as_slice())
|
||||
.await;
|
||||
sess.replace_history(
|
||||
reconstructed.history,
|
||||
reconstructed.reference_context_item.clone(),
|
||||
)
|
||||
.await;
|
||||
sess.set_previous_turn_settings(reconstructed.previous_turn_settings)
|
||||
.await;
|
||||
sess.recompute_token_usage(turn_context.as_ref()).await;
|
||||
|
||||
sess.send_event_raw_flushed(Event {
|
||||
id: turn_context.sub_id.clone(),
|
||||
msg: EventMsg::ThreadRolledBack(ThreadRolledBackEvent { num_turns }),
|
||||
msg: EventMsg::ThreadRolledBack(rollback_event),
|
||||
})
|
||||
.await;
|
||||
}
|
||||
@@ -6781,6 +6844,18 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_message(text: &str) -> ResponseItem {
|
||||
ResponseItem::Message {
|
||||
id: None,
|
||||
role: "assistant".to_string(),
|
||||
content: vec![ContentItem::OutputText {
|
||||
text: text.to_string(),
|
||||
}],
|
||||
end_turn: None,
|
||||
phase: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn skill_message(text: &str) -> ResponseItem {
|
||||
ResponseItem::Message {
|
||||
id: None,
|
||||
@@ -7706,59 +7781,37 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn thread_rollback_drops_last_turn_from_history() {
|
||||
let (sess, tc, rx) = make_session_and_context_with_rx().await;
|
||||
let rollout_path = attach_rollout_recorder(&sess).await;
|
||||
|
||||
let initial_context = sess.build_initial_context(tc.as_ref()).await;
|
||||
sess.record_into_history(&initial_context, tc.as_ref())
|
||||
.await;
|
||||
|
||||
let turn_1 = vec![
|
||||
ResponseItem::Message {
|
||||
id: None,
|
||||
role: "user".to_string(),
|
||||
content: vec![ContentItem::InputText {
|
||||
text: "turn 1 user".to_string(),
|
||||
}],
|
||||
end_turn: None,
|
||||
phase: None,
|
||||
},
|
||||
ResponseItem::Message {
|
||||
id: None,
|
||||
role: "assistant".to_string(),
|
||||
content: vec![ContentItem::OutputText {
|
||||
text: "turn 1 assistant".to_string(),
|
||||
}],
|
||||
end_turn: None,
|
||||
phase: None,
|
||||
},
|
||||
user_message("turn 1 user"),
|
||||
assistant_message("turn 1 assistant"),
|
||||
];
|
||||
sess.record_into_history(&turn_1, tc.as_ref()).await;
|
||||
|
||||
let turn_2 = vec![
|
||||
ResponseItem::Message {
|
||||
id: None,
|
||||
role: "user".to_string(),
|
||||
content: vec![ContentItem::InputText {
|
||||
text: "turn 2 user".to_string(),
|
||||
}],
|
||||
end_turn: None,
|
||||
phase: None,
|
||||
},
|
||||
ResponseItem::Message {
|
||||
id: None,
|
||||
role: "assistant".to_string(),
|
||||
content: vec![ContentItem::OutputText {
|
||||
text: "turn 2 assistant".to_string(),
|
||||
}],
|
||||
end_turn: None,
|
||||
phase: None,
|
||||
},
|
||||
user_message("turn 2 user"),
|
||||
assistant_message("turn 2 assistant"),
|
||||
];
|
||||
sess.record_into_history(&turn_2, tc.as_ref()).await;
|
||||
let mut full_history = Vec::new();
|
||||
full_history.extend(initial_context.clone());
|
||||
full_history.extend(turn_1.clone());
|
||||
full_history.extend(turn_2);
|
||||
sess.replace_history(full_history.clone(), Some(tc.to_turn_context_item()))
|
||||
.await;
|
||||
let rollout_items: Vec<RolloutItem> = full_history
|
||||
.into_iter()
|
||||
.map(RolloutItem::ResponseItem)
|
||||
.collect();
|
||||
sess.persist_rollout_items(&rollout_items).await;
|
||||
sess.set_previous_turn_settings(Some(PreviousTurnSettings {
|
||||
model: "previous-regular-model".to_string(),
|
||||
model: "stale-model".to_string(),
|
||||
realtime_active: Some(tc.realtime_active),
|
||||
}))
|
||||
.await;
|
||||
{
|
||||
let mut state = sess.state.lock().await;
|
||||
state.set_reference_context_item(Some(tc.to_turn_context_item()));
|
||||
}
|
||||
|
||||
handlers::thread_rollback(&sess, "sub-1".to_string(), 1).await;
|
||||
|
||||
@@ -7771,33 +7824,41 @@ mod tests {
|
||||
|
||||
let history = sess.clone_history().await;
|
||||
assert_eq!(expected, history.raw_items());
|
||||
assert_eq!(
|
||||
sess.previous_turn_settings().await,
|
||||
Some(PreviousTurnSettings {
|
||||
model: "previous-regular-model".to_string(),
|
||||
realtime_active: Some(tc.realtime_active),
|
||||
})
|
||||
);
|
||||
assert_eq!(sess.previous_turn_settings().await, None);
|
||||
assert!(sess.reference_context_item().await.is_none());
|
||||
|
||||
let InitialHistory::Resumed(resumed) = RolloutRecorder::get_rollout_history(&rollout_path)
|
||||
.await
|
||||
.expect("read rollout history")
|
||||
else {
|
||||
panic!("expected resumed rollout history");
|
||||
};
|
||||
assert!(resumed.history.iter().any(|item| {
|
||||
matches!(
|
||||
item,
|
||||
RolloutItem::EventMsg(EventMsg::ThreadRolledBack(rollback))
|
||||
if rollback.num_turns == 1
|
||||
)
|
||||
}));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn thread_rollback_clears_history_when_num_turns_exceeds_existing_turns() {
|
||||
let (sess, tc, rx) = make_session_and_context_with_rx().await;
|
||||
attach_rollout_recorder(&sess).await;
|
||||
|
||||
let initial_context = sess.build_initial_context(tc.as_ref()).await;
|
||||
sess.record_into_history(&initial_context, tc.as_ref())
|
||||
let turn_1 = vec![user_message("turn 1 user")];
|
||||
let mut full_history = Vec::new();
|
||||
full_history.extend(initial_context.clone());
|
||||
full_history.extend(turn_1);
|
||||
sess.replace_history(full_history.clone(), Some(tc.to_turn_context_item()))
|
||||
.await;
|
||||
|
||||
let turn_1 = vec![ResponseItem::Message {
|
||||
id: None,
|
||||
role: "user".to_string(),
|
||||
content: vec![ContentItem::InputText {
|
||||
text: "turn 1 user".to_string(),
|
||||
}],
|
||||
end_turn: None,
|
||||
phase: None,
|
||||
}];
|
||||
sess.record_into_history(&turn_1, tc.as_ref()).await;
|
||||
let rollout_items: Vec<RolloutItem> = full_history
|
||||
.into_iter()
|
||||
.map(RolloutItem::ResponseItem)
|
||||
.collect();
|
||||
sess.persist_rollout_items(&rollout_items).await;
|
||||
|
||||
handlers::thread_rollback(&sess, "sub-1".to_string(), 99).await;
|
||||
|
||||
@@ -7808,6 +7869,230 @@ mod tests {
|
||||
assert_eq!(initial_context, history.raw_items());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn thread_rollback_fails_without_persisted_rollout_path() {
|
||||
let (sess, tc, rx) = make_session_and_context_with_rx().await;
|
||||
|
||||
let initial_context = sess.build_initial_context(tc.as_ref()).await;
|
||||
sess.record_into_history(&initial_context, tc.as_ref())
|
||||
.await;
|
||||
|
||||
handlers::thread_rollback(&sess, "sub-1".to_string(), 1).await;
|
||||
|
||||
let error_event = wait_for_thread_rollback_failed(&rx).await;
|
||||
assert_eq!(
|
||||
error_event.message,
|
||||
"thread rollback requires a persisted rollout path"
|
||||
);
|
||||
assert_eq!(
|
||||
error_event.codex_error_info,
|
||||
Some(CodexErrorInfo::ThreadRollbackFailed)
|
||||
);
|
||||
assert_eq!(sess.clone_history().await.raw_items(), initial_context);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn thread_rollback_recomputes_previous_turn_settings_and_reference_context_from_replay() {
|
||||
let (sess, tc, rx) = make_session_and_context_with_rx().await;
|
||||
attach_rollout_recorder(&sess).await;
|
||||
|
||||
let first_context_item = tc.to_turn_context_item();
|
||||
let first_turn_id = first_context_item
|
||||
.turn_id
|
||||
.clone()
|
||||
.expect("turn context should have turn_id");
|
||||
let mut rolled_back_context_item = first_context_item.clone();
|
||||
rolled_back_context_item.turn_id = Some("rolled-back-turn".to_string());
|
||||
rolled_back_context_item.model = "rolled-back-model".to_string();
|
||||
let rolled_back_turn_id = rolled_back_context_item
|
||||
.turn_id
|
||||
.clone()
|
||||
.expect("turn context should have turn_id");
|
||||
let turn_one_user = user_message("turn 1 user");
|
||||
let turn_one_assistant = assistant_message("turn 1 assistant");
|
||||
let turn_two_user = user_message("turn 2 user");
|
||||
let turn_two_assistant = assistant_message("turn 2 assistant");
|
||||
|
||||
sess.persist_rollout_items(&[
|
||||
RolloutItem::EventMsg(EventMsg::TurnStarted(
|
||||
codex_protocol::protocol::TurnStartedEvent {
|
||||
turn_id: first_turn_id.clone(),
|
||||
model_context_window: Some(128_000),
|
||||
collaboration_mode_kind: ModeKind::Default,
|
||||
},
|
||||
)),
|
||||
RolloutItem::EventMsg(EventMsg::UserMessage(
|
||||
codex_protocol::protocol::UserMessageEvent {
|
||||
message: "turn 1 user".to_string(),
|
||||
images: None,
|
||||
local_images: Vec::new(),
|
||||
text_elements: Vec::new(),
|
||||
},
|
||||
)),
|
||||
RolloutItem::TurnContext(first_context_item.clone()),
|
||||
RolloutItem::ResponseItem(turn_one_user.clone()),
|
||||
RolloutItem::ResponseItem(turn_one_assistant.clone()),
|
||||
RolloutItem::EventMsg(EventMsg::TurnComplete(TurnCompleteEvent {
|
||||
turn_id: first_turn_id,
|
||||
last_agent_message: None,
|
||||
})),
|
||||
RolloutItem::EventMsg(EventMsg::TurnStarted(
|
||||
codex_protocol::protocol::TurnStartedEvent {
|
||||
turn_id: rolled_back_turn_id.clone(),
|
||||
model_context_window: Some(128_000),
|
||||
collaboration_mode_kind: ModeKind::Default,
|
||||
},
|
||||
)),
|
||||
RolloutItem::EventMsg(EventMsg::UserMessage(
|
||||
codex_protocol::protocol::UserMessageEvent {
|
||||
message: "turn 2 user".to_string(),
|
||||
images: None,
|
||||
local_images: Vec::new(),
|
||||
text_elements: Vec::new(),
|
||||
},
|
||||
)),
|
||||
RolloutItem::TurnContext(rolled_back_context_item),
|
||||
RolloutItem::ResponseItem(turn_two_user),
|
||||
RolloutItem::ResponseItem(turn_two_assistant),
|
||||
RolloutItem::EventMsg(EventMsg::TurnComplete(TurnCompleteEvent {
|
||||
turn_id: rolled_back_turn_id,
|
||||
last_agent_message: None,
|
||||
})),
|
||||
])
|
||||
.await;
|
||||
sess.replace_history(
|
||||
vec![assistant_message("stale history")],
|
||||
Some(first_context_item.clone()),
|
||||
)
|
||||
.await;
|
||||
sess.set_previous_turn_settings(Some(PreviousTurnSettings {
|
||||
model: "stale-model".to_string(),
|
||||
realtime_active: None,
|
||||
}))
|
||||
.await;
|
||||
|
||||
handlers::thread_rollback(&sess, "sub-1".to_string(), 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(),
|
||||
vec![turn_one_user, turn_one_assistant]
|
||||
);
|
||||
assert_eq!(
|
||||
sess.previous_turn_settings().await,
|
||||
Some(PreviousTurnSettings {
|
||||
model: tc.model_info.slug.clone(),
|
||||
realtime_active: Some(tc.realtime_active),
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::to_value(sess.reference_context_item().await)
|
||||
.expect("serialize replay reference context item"),
|
||||
serde_json::to_value(Some(first_context_item))
|
||||
.expect("serialize expected reference context item")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn thread_rollback_persists_marker_and_replays_cumulatively() {
|
||||
let (sess, tc, rx) = make_session_and_context_with_rx().await;
|
||||
let rollout_path = attach_rollout_recorder(&sess).await;
|
||||
let turn_context_item = tc.to_turn_context_item();
|
||||
|
||||
sess.persist_rollout_items(&[
|
||||
RolloutItem::EventMsg(EventMsg::TurnStarted(
|
||||
codex_protocol::protocol::TurnStartedEvent {
|
||||
turn_id: "turn-1".to_string(),
|
||||
model_context_window: Some(128_000),
|
||||
collaboration_mode_kind: ModeKind::Default,
|
||||
},
|
||||
)),
|
||||
RolloutItem::EventMsg(EventMsg::UserMessage(UserMessageEvent {
|
||||
message: "turn 1 user".to_string(),
|
||||
images: None,
|
||||
local_images: Vec::new(),
|
||||
text_elements: Vec::new(),
|
||||
})),
|
||||
RolloutItem::TurnContext(turn_context_item.clone()),
|
||||
RolloutItem::ResponseItem(user_message("turn 1 user")),
|
||||
RolloutItem::ResponseItem(assistant_message("turn 1 assistant")),
|
||||
RolloutItem::EventMsg(EventMsg::TurnComplete(TurnCompleteEvent {
|
||||
turn_id: "turn-1".to_string(),
|
||||
last_agent_message: None,
|
||||
})),
|
||||
RolloutItem::EventMsg(EventMsg::TurnStarted(
|
||||
codex_protocol::protocol::TurnStartedEvent {
|
||||
turn_id: "turn-2".to_string(),
|
||||
model_context_window: Some(128_000),
|
||||
collaboration_mode_kind: ModeKind::Default,
|
||||
},
|
||||
)),
|
||||
RolloutItem::EventMsg(EventMsg::UserMessage(UserMessageEvent {
|
||||
message: "turn 2 user".to_string(),
|
||||
images: None,
|
||||
local_images: Vec::new(),
|
||||
text_elements: Vec::new(),
|
||||
})),
|
||||
RolloutItem::TurnContext(turn_context_item.clone()),
|
||||
RolloutItem::ResponseItem(user_message("turn 2 user")),
|
||||
RolloutItem::ResponseItem(assistant_message("turn 2 assistant")),
|
||||
RolloutItem::EventMsg(EventMsg::TurnComplete(TurnCompleteEvent {
|
||||
turn_id: "turn-2".to_string(),
|
||||
last_agent_message: None,
|
||||
})),
|
||||
RolloutItem::EventMsg(EventMsg::TurnStarted(
|
||||
codex_protocol::protocol::TurnStartedEvent {
|
||||
turn_id: "turn-3".to_string(),
|
||||
model_context_window: Some(128_000),
|
||||
collaboration_mode_kind: ModeKind::Default,
|
||||
},
|
||||
)),
|
||||
RolloutItem::EventMsg(EventMsg::UserMessage(UserMessageEvent {
|
||||
message: "turn 3 user".to_string(),
|
||||
images: None,
|
||||
local_images: Vec::new(),
|
||||
text_elements: Vec::new(),
|
||||
})),
|
||||
RolloutItem::TurnContext(turn_context_item),
|
||||
RolloutItem::ResponseItem(user_message("turn 3 user")),
|
||||
RolloutItem::ResponseItem(assistant_message("turn 3 assistant")),
|
||||
RolloutItem::EventMsg(EventMsg::TurnComplete(TurnCompleteEvent {
|
||||
turn_id: "turn-3".to_string(),
|
||||
last_agent_message: None,
|
||||
})),
|
||||
])
|
||||
.await;
|
||||
|
||||
handlers::thread_rollback(&sess, "sub-1".to_string(), 1).await;
|
||||
let first_rollback = wait_for_thread_rolled_back(&rx).await;
|
||||
assert_eq!(first_rollback.num_turns, 1);
|
||||
handlers::thread_rollback(&sess, "sub-1".to_string(), 1).await;
|
||||
let second_rollback = wait_for_thread_rolled_back(&rx).await;
|
||||
assert_eq!(second_rollback.num_turns, 1);
|
||||
|
||||
assert_eq!(
|
||||
sess.clone_history().await.raw_items(),
|
||||
vec![
|
||||
user_message("turn 1 user"),
|
||||
assistant_message("turn 1 assistant")
|
||||
]
|
||||
);
|
||||
|
||||
let InitialHistory::Resumed(resumed) = RolloutRecorder::get_rollout_history(&rollout_path)
|
||||
.await
|
||||
.expect("read rollout history")
|
||||
else {
|
||||
panic!("expected resumed rollout history");
|
||||
};
|
||||
let rollback_markers = resumed
|
||||
.history
|
||||
.iter()
|
||||
.filter(|item| matches!(item, RolloutItem::EventMsg(EventMsg::ThreadRolledBack(_))))
|
||||
.count();
|
||||
assert_eq!(rollback_markers, 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn thread_rollback_fails_when_turn_in_progress() {
|
||||
let (sess, tc, rx) = make_session_and_context_with_rx().await;
|
||||
@@ -8224,6 +8509,33 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
async fn attach_rollout_recorder(session: &Arc<Session>) -> PathBuf {
|
||||
let config = session.get_config().await;
|
||||
let recorder = RolloutRecorder::new(
|
||||
config.as_ref(),
|
||||
RolloutRecorderParams::new(
|
||||
ThreadId::default(),
|
||||
None,
|
||||
SessionSource::Exec,
|
||||
BaseInstructions::default(),
|
||||
Vec::new(),
|
||||
EventPersistenceMode::Limited,
|
||||
),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("create rollout recorder");
|
||||
let rollout_path = recorder.rollout_path().to_path_buf();
|
||||
{
|
||||
let mut rollout = session.services.rollout.lock().await;
|
||||
*rollout = Some(recorder);
|
||||
}
|
||||
session.ensure_rollout_materialized().await;
|
||||
session.flush_rollout().await;
|
||||
rollout_path
|
||||
}
|
||||
|
||||
fn text_block(s: &str) -> serde_json::Value {
|
||||
json!({
|
||||
"type": "text",
|
||||
|
||||
@@ -20,10 +20,15 @@ use codex_protocol::protocol::EventMsg;
|
||||
use codex_protocol::protocol::Op;
|
||||
use codex_protocol::protocol::WarningEvent;
|
||||
use codex_protocol::user_input::UserInput;
|
||||
use core_test_support::context_snapshot;
|
||||
use core_test_support::context_snapshot::ContextSnapshotOptions;
|
||||
use core_test_support::context_snapshot::ContextSnapshotRenderMode;
|
||||
use core_test_support::responses::ResponseMock;
|
||||
use core_test_support::responses::ResponsesRequest;
|
||||
use core_test_support::responses::ev_assistant_message;
|
||||
use core_test_support::responses::ev_completed;
|
||||
use core_test_support::responses::mount_sse_once_match;
|
||||
use core_test_support::responses::mount_sse_sequence;
|
||||
use core_test_support::responses::sse;
|
||||
use core_test_support::test_codex::test_codex;
|
||||
use core_test_support::wait_for_event;
|
||||
@@ -35,6 +40,7 @@ use tempfile::TempDir;
|
||||
use wiremock::MockServer;
|
||||
|
||||
const AFTER_SECOND_RESUME: &str = "AFTER_SECOND_RESUME";
|
||||
const AFTER_ROLLBACK: &str = "AFTER_ROLLBACK";
|
||||
|
||||
fn network_disabled() -> bool {
|
||||
std::env::var(CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR).is_ok()
|
||||
@@ -406,6 +412,95 @@ async fn compact_resume_after_second_compaction_preserves_history() -> Result<()
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
/// Scenario: rolling back behind a pre-turn compaction should replay
|
||||
/// append-only history from the rollout file and keep earlier compacted
|
||||
/// history visible.
|
||||
async fn snapshot_rollback_past_compaction_replays_append_only_history() -> Result<()> {
|
||||
if network_disabled() {
|
||||
println!("Skipping test because network is disabled in this sandbox");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
const EDITED_AFTER_COMPACT: &str = "EDITED_AFTER_COMPACT";
|
||||
const SECOND_REPLY: &str = "SECOND_REPLY";
|
||||
|
||||
let server = MockServer::start().await;
|
||||
let sse1 = sse(vec![
|
||||
ev_assistant_message("m1", FIRST_REPLY),
|
||||
ev_completed("r1"),
|
||||
]);
|
||||
let sse2 = sse(vec![
|
||||
ev_assistant_message("m2", SUMMARY_TEXT),
|
||||
ev_completed("r2"),
|
||||
]);
|
||||
let sse3 = sse(vec![
|
||||
ev_assistant_message("m3", SECOND_REPLY),
|
||||
ev_completed("r3"),
|
||||
]);
|
||||
let sse4 = sse(vec![ev_completed("r4")]);
|
||||
|
||||
let request_log = mount_sse_sequence(&server, vec![sse1, sse2, sse3, sse4]).await;
|
||||
|
||||
let (_home, _config, _manager, base) = start_test_conversation(&server, None).await;
|
||||
|
||||
user_turn(&base, "hello world").await;
|
||||
compact_conversation(&base).await;
|
||||
user_turn(&base, EDITED_AFTER_COMPACT).await;
|
||||
|
||||
base.submit(Op::ThreadRollback { num_turns: 1 })
|
||||
.await
|
||||
.expect("submit thread rollback");
|
||||
let rollback_event =
|
||||
wait_for_event(&base, |ev| matches!(ev, EventMsg::ThreadRolledBack(_))).await;
|
||||
let EventMsg::ThreadRolledBack(rollback_event) = rollback_event else {
|
||||
panic!("expected thread rolled back event");
|
||||
};
|
||||
assert_eq!(rollback_event.num_turns, 1);
|
||||
|
||||
user_turn(&base, AFTER_ROLLBACK).await;
|
||||
|
||||
let requests = request_log.requests();
|
||||
assert_eq!(requests.len(), 4);
|
||||
assert!(requests[1].body_contains_text(SUMMARIZATION_PROMPT));
|
||||
assert!(requests[2].body_contains_text("hello world"));
|
||||
assert!(requests[2].body_contains_text(SUMMARY_TEXT));
|
||||
assert!(requests[2].body_contains_text(EDITED_AFTER_COMPACT));
|
||||
let after_rollback_user_texts = requests[3].message_input_texts("user");
|
||||
let after_rollback_last = after_rollback_user_texts
|
||||
.last()
|
||||
.unwrap_or_else(|| panic!("post-rollback request missing user messages"));
|
||||
assert_eq!(after_rollback_last, AFTER_ROLLBACK);
|
||||
assert!(
|
||||
requests[3].body_contains_text("hello world"),
|
||||
"the first turn should remain visible after rollback behind compaction",
|
||||
);
|
||||
assert!(
|
||||
!requests[3].body_contains_text(EDITED_AFTER_COMPACT),
|
||||
"the edited post-compaction turn should be removed by rollback",
|
||||
);
|
||||
assert!(
|
||||
requests[3].body_contains_text(SUMMARY_TEXT),
|
||||
"compaction summary should remain for the preserved first turn",
|
||||
);
|
||||
|
||||
insta::assert_snapshot!(
|
||||
"rollback_past_compaction_shapes",
|
||||
context_snapshot::format_labeled_requests_snapshot(
|
||||
"rollback past compaction replay after rollback",
|
||||
&[
|
||||
("compaction request", &requests[1]),
|
||||
("before rollback", &requests[2]),
|
||||
("after rollback", &requests[3]),
|
||||
],
|
||||
&ContextSnapshotOptions::default()
|
||||
.render_mode(ContextSnapshotRenderMode::KindWithTextPrefix { max_chars: 64 }),
|
||||
)
|
||||
);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn normalize_line_endings(value: &mut Value) {
|
||||
match value {
|
||||
Value::String(text) => {
|
||||
@@ -427,10 +522,16 @@ fn normalize_line_endings(value: &mut Value) {
|
||||
}
|
||||
}
|
||||
|
||||
fn gather_request_bodies(request_log: &[ResponseMock]) -> Vec<Value> {
|
||||
let mut bodies = request_log
|
||||
fn gather_requests(request_log: &[ResponseMock]) -> Vec<ResponsesRequest> {
|
||||
request_log
|
||||
.iter()
|
||||
.flat_map(ResponseMock::requests)
|
||||
.collect::<Vec<_>>()
|
||||
}
|
||||
|
||||
fn gather_request_bodies(request_log: &[ResponseMock]) -> Vec<Value> {
|
||||
let mut bodies = gather_requests(request_log)
|
||||
.into_iter()
|
||||
.map(|request| request.body_json())
|
||||
.collect::<Vec<_>>();
|
||||
bodies.iter_mut().for_each(normalize_line_endings);
|
||||
|
||||
+36
@@ -0,0 +1,36 @@
|
||||
---
|
||||
source: core/tests/suite/compact_resume_fork.rs
|
||||
expression: "context_snapshot::format_labeled_requests_snapshot(\"rollback past compaction replay after rollback\",\n&[(\"compaction request\", &requests[1]), (\"before rollback\", &requests[2]),\n(\"after rollback\", &requests[3]),],\n&ContextSnapshotOptions::default().render_mode(ContextSnapshotRenderMode::KindWithTextPrefix\n{ max_chars: 64 }),)"
|
||||
---
|
||||
Scenario: rollback past compaction replay after rollback
|
||||
|
||||
## compaction request
|
||||
00:message/developer:<PERMISSIONS_INSTRUCTIONS>
|
||||
01:message/user[2]:
|
||||
[01] <AGENTS_MD>
|
||||
[02] <ENVIRONMENT_CONTEXT:cwd=<CWD>>
|
||||
02:message/user:hello world
|
||||
03:message/assistant:FIRST_REPLY
|
||||
04:message/user:<SUMMARIZATION_PROMPT>
|
||||
|
||||
## before rollback
|
||||
00:message/user:hello world
|
||||
01:message/user:<COMPACTION_SUMMARY>\nSUMMARY_ONLY_CONTEXT
|
||||
02:message/developer:<PERMISSIONS_INSTRUCTIONS>
|
||||
03:message/user[2]:
|
||||
[01] <AGENTS_MD>
|
||||
[02] <ENVIRONMENT_CONTEXT:cwd=<CWD>>
|
||||
04:message/user:EDITED_AFTER_COMPACT
|
||||
|
||||
## after rollback
|
||||
00:message/user:hello world
|
||||
01:message/user:<COMPACTION_SUMMARY>\nSUMMARY_ONLY_CONTEXT
|
||||
02:message/developer:<PERMISSIONS_INSTRUCTIONS>
|
||||
03:message/user[2]:
|
||||
[01] <AGENTS_MD>
|
||||
[02] <ENVIRONMENT_CONTEXT:cwd=<CWD>>
|
||||
04:message/developer:<PERMISSIONS_INSTRUCTIONS>
|
||||
05:message/user[2]:
|
||||
[01] <AGENTS_MD>
|
||||
[02] <ENVIRONMENT_CONTEXT:cwd=<CWD>>
|
||||
06:message/user:AFTER_ROLLBACK
|
||||
Reference in New Issue
Block a user