mirror of
https://github.com/openai/codex.git
synced 2026-09-14 11:57:03 +00:00
Fix rollback baseline rehydration
This commit is contained in:
@@ -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")
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user