This commit is contained in:
Ahmed Ibrahim
2025-09-10 09:49:24 -07:00
parent 35107a121a
commit 2f50987567
3 changed files with 69 additions and 55 deletions

View File

@@ -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<RolloutItem> = 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<RolloutItem> = (&state.response_items).into();
let rolled_event_msgs: Vec<RolloutItem> = (&state.event_msgs).into();
warn!("rolled_event_msgs: {:?}", rolled_event_msgs);
[rolled_response_items, rolled_event_msgs].concat()
};
let event = Event {

View File

@@ -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)
}
}

View File

@@ -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<NewConversation> {
// 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<RolloutItem> = history.get_rollout_items();
// Determine the cut point among response items (counting only ResponseItem::Message with role=="user").
let response_items: Vec<ResponseItem> = 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<String> =
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<EventMsg> =
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<ResponseItem> = 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<EventMsg> {
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<usize> {
if n == 0 {
@@ -224,53 +262,30 @@ fn user_message_text_for_response(item: &ResponseItem) -> Option<String> {
})
}
/// 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<usize> {
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<RolloutItem>,
cut_resp_index: usize,
event_cut_index: Option<usize>,
event_msgs: &[EventMsg],
response_items: &[ResponseItem],
) -> Vec<RolloutItem> {
let mut kept_response_seen = 0usize;
let mut rolled: Vec<RolloutItem> = 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<RolloutItem> = 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
}