Fix rollback baseline rehydration

This commit is contained in:
Charles Cunningham
2026-02-20 23:28:48 -08:00
parent bb0ac5be70
commit 534680fd05

View File

@@ -3511,6 +3511,7 @@ mod handlers {
use crate::mcp::collect_mcp_snapshot_from_manager;
use crate::mcp::effective_mcp_servers;
use crate::review_prompts::resolve_review_request;
use crate::rollout::RolloutRecorder;
use crate::rollout::session_index;
use crate::tasks::CompactTask;
use crate::tasks::UndoTask;
@@ -4104,18 +4105,13 @@ mod handlers {
let turn_context = sess.new_default_turn_with_sub_id(sub_id).await;
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_model`
// so the next regular turn replays any dropped model-switch
// instructions.
history.drop_last_n_user_turns(num_turns);
// 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)
.await;
sess.set_previous_model(None).await;
sess.recompute_token_usage(turn_context.as_ref()).await;
sess.send_event_raw_flushed(Event {
@@ -4123,6 +4119,30 @@ mod handlers {
msg: EventMsg::ThreadRolledBack(ThreadRolledBackEvent { num_turns }),
})
.await;
// Rehydrate the rollback-adjusted baseline from the persisted rollout view so the next
// regular turn gets correct model-switch and full-context-reinjection behavior.
let recorder = { sess.services.rollout.lock().await.clone() };
if let Some(recorder) = recorder {
match RolloutRecorder::load_rollout_items(recorder.rollout_path()).await {
Ok((rollout_items, _, _)) => {
let (previous_turn_context_item, crossed_compaction_after_turn) =
Session::last_rollout_regular_turn_context_lookup(&rollout_items);
let previous_model = previous_turn_context_item.map(|ctx| ctx.model.clone());
let reference_context_item = if crossed_compaction_after_turn {
None
} else {
previous_turn_context_item.cloned()
};
let mut state = sess.state.lock().await;
state.set_previous_model(previous_model);
state.set_reference_context_item(reference_context_item);
}
Err(err) => {
warn!("failed to reload rollout after rollback: {err}");
}
}
}
}
/// Persists the thread name in the session index, updates in-memory state, and emits
@@ -7141,9 +7161,126 @@ mod tests {
let history = sess.clone_history().await;
assert_eq!(expected, history.raw_items());
assert_eq!(sess.previous_model().await, None);
}
#[tokio::test]
async fn thread_rollback_rehydrates_baseline_from_rollout() {
let (sess, tc, rx) = make_session_and_context_with_rx().await;
let (config, session_source, base_instructions, dynamic_tools, persist_extended_history) = {
let state = sess.state.lock().await;
(
Arc::clone(&state.session_configuration.original_config_do_not_use),
state.session_configuration.session_source.clone(),
state.session_configuration.base_instructions.clone(),
state.session_configuration.dynamic_tools.clone(),
state.session_configuration.persist_extended_history,
)
};
let recorder = RolloutRecorder::new(
config.as_ref(),
RolloutRecorderParams::new(
sess.conversation_id,
None,
session_source,
BaseInstructions {
text: base_instructions,
},
dynamic_tools,
if persist_extended_history {
EventPersistenceMode::Extended
} else {
EventPersistenceMode::Limited
},
),
None,
None,
)
.await
.expect("create rollout recorder");
{
let mut guard = sess.services.rollout.lock().await;
*guard = Some(recorder);
}
sess.ensure_rollout_materialized().await;
let turn_1 = vec![user_message("turn 1 user")];
let turn_2 = vec![user_message("turn 2 user")];
sess.record_into_history(&turn_1, tc.as_ref()).await;
sess.record_into_history(&turn_2, tc.as_ref()).await;
let mut turn_1_context = tc.to_turn_context_item();
turn_1_context.turn_id = Some("turn-1".to_string());
turn_1_context.model = "model-a".to_string();
let mut turn_2_context = tc.to_turn_context_item();
turn_2_context.turn_id = Some("turn-2".to_string());
turn_2_context.model = "model-b".to_string();
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(
codex_protocol::protocol::UserMessageEvent {
message: "turn 1 user".to_string(),
images: None,
local_images: Vec::new(),
text_elements: Vec::new(),
},
)),
RolloutItem::TurnContext(turn_1_context.clone()),
RolloutItem::EventMsg(EventMsg::TurnComplete(
codex_protocol::protocol::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(
codex_protocol::protocol::UserMessageEvent {
message: "turn 2 user".to_string(),
images: None,
local_images: Vec::new(),
text_elements: Vec::new(),
},
)),
RolloutItem::TurnContext(turn_2_context),
RolloutItem::EventMsg(EventMsg::TurnComplete(
codex_protocol::protocol::TurnCompleteEvent {
turn_id: "turn-2".to_string(),
last_agent_message: None,
},
)),
])
.await;
sess.flush_rollout().await;
sess.set_previous_model(Some("model-b".to_string())).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;
let rollback_event = wait_for_thread_rolled_back(&rx).await;
assert_eq!(rollback_event.num_turns, 1);
assert_eq!(sess.previous_model().await, Some("model-a".to_string()));
assert_eq!(
sess.previous_model().await,
Some("previous-regular-model".to_string())
serde_json::to_value(sess.reference_context_item().await)
.expect("serialize post-rollback reference context item"),
serde_json::to_value(Some(turn_1_context))
.expect("serialize expected post-rollback reference context item")
);
}