From 2f50987567d3db5222827de5ea73754477c6dfeb Mon Sep 17 00:00:00 2001 From: Ahmed Ibrahim Date: Wed, 10 Sep 2025 09:49:24 -0700 Subject: [PATCH] forking --- codex-rs/core/src/codex.rs | 3 + codex-rs/core/src/conversation_history.rs | 10 +- codex-rs/core/src/conversation_manager.rs | 111 ++++++++++++---------- 3 files changed, 69 insertions(+), 55 deletions(-) diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index 50af0b1c7d..58142421b4 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -724,6 +724,7 @@ impl Session { let msgs = map_response_item_to_event_messages(&response_item, self.show_raw_agent_reasoning); let user_msgs: Vec = msgs + .clone() .into_iter() .filter_map(|m| match m { EventMsg::UserMessage(ev) => Some(RolloutItem::EventMsg(EventMsg::UserMessage(ev))), @@ -733,6 +734,7 @@ impl Session { if !user_msgs.is_empty() { self.persist_rollout_items(&user_msgs).await; } + self.state.lock_unchecked().event_msgs.record_items(&msgs); } async fn on_exec_command_begin( @@ -1391,6 +1393,7 @@ async fn submission_loop( let state = sess.state.lock_unchecked(); let rolled_response_items: Vec = (&state.response_items).into(); let rolled_event_msgs: Vec = (&state.event_msgs).into(); + warn!("rolled_event_msgs: {:?}", rolled_event_msgs); [rolled_response_items, rolled_event_msgs].concat() }; let event = Event { diff --git a/codex-rs/core/src/conversation_history.rs b/codex-rs/core/src/conversation_history.rs index e66e81d277..e718526748 100644 --- a/codex-rs/core/src/conversation_history.rs +++ b/codex-rs/core/src/conversation_history.rs @@ -1,3 +1,4 @@ +use crate::rollout::policy::should_persist_event_msg; use codex_protocol::models::ResponseItem; use codex_protocol::protocol::EventMsg; use codex_protocol::protocol::RolloutItem; @@ -80,14 +81,9 @@ impl EventMsgsHistory { } } } + fn should_record_item(&self, item: &EventMsg) -> bool { - !matches!( - item, - EventMsg::AgentMessageDelta(_) - | EventMsg::AgentReasoningDelta(_) - | EventMsg::AgentReasoningRawContentDelta(_) - | EventMsg::ExecCommandOutputDelta(_) - ) + should_persist_event_msg(item) } } diff --git a/codex-rs/core/src/conversation_manager.rs b/codex-rs/core/src/conversation_manager.rs index 0afafe8662..0d49b37f63 100644 --- a/codex-rs/core/src/conversation_manager.rs +++ b/codex-rs/core/src/conversation_manager.rs @@ -20,6 +20,7 @@ use std::collections::HashMap; use std::path::PathBuf; use std::sync::Arc; use tokio::sync::RwLock; +use tracing::warn; /// Represents a newly created Codex conversation, including the first event /// (which is [`EventMsg::SessionConfigured`]). @@ -156,6 +157,10 @@ impl ConversationManager { config: Config, ) -> CodexResult { // Compute the prefix up to the cut point. + warn!( + "eventmsgs in fork_conversation: {:?}", + conversation_history.get_event_msgs() + ); let history = truncate_after_dropping_last_messages(conversation_history, num_messages_to_drop); @@ -173,10 +178,8 @@ impl ConversationManager { /// Return a prefix of `items` obtained by dropping the last `n` user messages /// and all items that follow them. fn truncate_after_dropping_last_messages(history: InitialHistory, n: usize) -> InitialHistory { - // Work from response items for cut logic; preserve any existing rollout items when possible. - let rollout_items: Vec = history.get_rollout_items(); + // Determine the cut point among response items (counting only ResponseItem::Message with role=="user"). let response_items: Vec = history.get_response_items(); - if n == 0 { return history; } @@ -189,13 +192,48 @@ fn truncate_after_dropping_last_messages(history: InitialHistory, n: usize) -> I return InitialHistory::New; } - let cut_events_index = - find_matching_user_event_index_in_rollout(&rollout_items, &response_items, cut_resp_index); + // Identify the specific user message text at the cut response index. + let target_message: Option = + user_message_text_for_response(&response_items[cut_resp_index]); - let rolled = build_truncated_rollout(rollout_items, cut_resp_index, cut_events_index); + // Compute event prefix by cutting at the matching user event (if present). + let event_msgs_prefix: Vec = + event_msgs_prefix_until_target(&history, target_message.as_deref()); + warn!("event_msgs_prefix: {:?}", event_msgs_prefix); + + // Keep only response items strictly before the cut response index. + let response_prefix: Vec = response_items[..cut_resp_index].to_vec(); + + let rolled = build_truncated_rollout(&event_msgs_prefix, &response_prefix); InitialHistory::Forked(rolled) } +/// Build the event messages prefix from `history` by cutting at the first +/// `EventMsg::UserMessage` whose message equals `target_message`. +fn event_msgs_prefix_until_target( + history: &InitialHistory, + target_message: Option<&str>, +) -> Vec { + warn!("target_message: {:?}", target_message); + match history.get_event_msgs() { + Some(all_events) => { + if let Some(target) = target_message { + if let Some(idx) = find_matching_user_event_index_in_event_msgs(&all_events, target) + { + warn!("found matching user event index: {}", idx); + all_events[..idx].to_vec() + } else { + warn!("no matching user event index found"); + Vec::new() + } + } else { + Vec::new() + } + } + None => Vec::new(), + } +} + /// Find the index (into response items) of the Nth user message from the end. fn find_cut_response_index(response_items: &[ResponseItem], n: usize) -> Option { if n == 0 { @@ -224,53 +262,30 @@ fn user_message_text_for_response(item: &ResponseItem) -> Option { }) } -/// Given rollout items and the response-item cut index, locate the matching user EventMsg index. -fn find_matching_user_event_index_in_rollout( - rollout_items: &[RolloutItem], - response_items: &[ResponseItem], - cut_resp_index: usize, +/// Locate the matching user EventMsg index in the list of event messages. +fn find_matching_user_event_index_in_event_msgs( + event_msgs: &[EventMsg], + target_message: &str, ) -> Option { - let target_message = user_message_text_for_response(&response_items[cut_resp_index])?; - rollout_items - .iter() - .enumerate() - .find_map(|(i, it)| match it { - RolloutItem::EventMsg(EventMsg::UserMessage(u)) if u.message == target_message => { - Some(i) - } - _ => None, - }) + event_msgs.iter().enumerate().find_map(|(i, it)| match it { + EventMsg::UserMessage(u) if u.message == target_message => Some(i), + _ => None, + }) } -/// Build a truncated rollout keeping response items strictly before `cut_resp_index` and -/// event messages strictly before `event_cut_index` (when provided). Always keeps session meta. +/// Build a truncated rollout by concatenating the (already-sliced) event messages and response items. fn build_truncated_rollout( - rollout_items: Vec, - cut_resp_index: usize, - event_cut_index: Option, + event_msgs: &[EventMsg], + response_items: &[ResponseItem], ) -> Vec { - let mut kept_response_seen = 0usize; - let mut rolled: Vec = Vec::new(); - for (abs_idx, it) in rollout_items.into_iter().enumerate() { - match &it { - RolloutItem::ResponseItem(_) => { - if kept_response_seen < cut_resp_index { - rolled.push(it); - } - kept_response_seen += 1; - } - RolloutItem::EventMsg(_) => { - if let Some(evt_cut) = event_cut_index - && abs_idx < evt_cut - { - rolled.push(it); - } - } - RolloutItem::SessionMeta(_) => { - rolled.push(it); - } - } - } + let mut rolled: Vec = Vec::with_capacity(event_msgs.len() + response_items.len()); + rolled.extend(event_msgs.iter().cloned().map(RolloutItem::EventMsg)); + rolled.extend( + response_items + .iter() + .cloned() + .map(RolloutItem::ResponseItem), + ); rolled }