diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index e71fb7cd1c..2db2bec875 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -2659,6 +2659,7 @@ mod handlers { use codex_protocol::protocol::ThreadNameUpdatedEvent; use codex_protocol::protocol::ThreadRolledBackEvent; use codex_protocol::protocol::TurnAbortReason; + use codex_protocol::protocol::TurnContextItem; use codex_protocol::protocol::WarningEvent; use codex_protocol::request_user_input::RequestUserInputResponse; @@ -2676,6 +2677,8 @@ mod handlers { use tracing::info; use tracing::warn; + use crate::state::session::SessionState; + pub async fn interrupt(sess: &Arc) { sess.interrupt_task().await; } @@ -3109,27 +3112,12 @@ mod handlers { sess.replace_history(history.raw_items().to_vec()).await; { let mut state = sess.state.lock().await; - let mut updated_turn_context_history = existing_turn_context_history; - let truncated_len = updated_turn_context_history - .len() - .saturating_sub(num_turns as usize); - updated_turn_context_history.truncate(truncated_len); - if updated_turn_context_history.len() < user_turns { - updated_turn_context_history.resize_with(user_turns, || None); - } - state.set_turn_context_history(updated_turn_context_history); - let mut applied = false; - if state.turn_context_history.len() == user_turns - && let Some(collaboration_mode) = state - .turn_context_history - .last() - .and_then(|turn_context| turn_context.as_ref()) - .and_then(|turn_context| turn_context.collaboration_mode.clone()) - { - state.session_configuration.collaboration_mode = collaboration_mode; - applied = true; - } - state.force_inject_collaboration_instructions = !applied; + apply_rollback_turn_context_history( + &mut state, + existing_turn_context_history, + num_turns, + user_turns, + ); } sess.recompute_token_usage(turn_context.as_ref()).await; @@ -3206,6 +3194,35 @@ mod handlers { .await; } + fn apply_rollback_turn_context_history( + state: &mut SessionState, + existing_turn_context_history: Vec>, + num_turns: u32, + user_turns: usize, + ) { + let mut updated_turn_context_history = existing_turn_context_history; + let truncated_len = updated_turn_context_history + .len() + .saturating_sub(num_turns as usize); + updated_turn_context_history.truncate(truncated_len); + if updated_turn_context_history.len() < user_turns { + updated_turn_context_history.resize_with(user_turns, || None); + } + state.set_turn_context_history(updated_turn_context_history); + if state.turn_context_history.len() == user_turns + && let Some(collaboration_mode) = state + .turn_context_history + .last() + .and_then(|turn_context| turn_context.as_ref()) + .and_then(|turn_context| turn_context.collaboration_mode.clone()) + { + state.session_configuration.collaboration_mode = collaboration_mode; + state.force_inject_collaboration_instructions = false; + } else { + state.force_inject_collaboration_instructions = true; + } + } + pub async fn shutdown(sess: &Arc, sub_id: String) -> bool { sess.abort_all_tasks(TurnAbortReason::Interrupted).await; sess.services