diff --git a/codex-rs/core/src/conversation_manager.rs b/codex-rs/core/src/conversation_manager.rs index cfd35b6ab4..de47be79b2 100644 --- a/codex-rs/core/src/conversation_manager.rs +++ b/codex-rs/core/src/conversation_manager.rs @@ -7,7 +7,6 @@ use crate::codex_conversation::CodexConversation; use crate::config::Config; use crate::error::CodexErr; use crate::error::Result as CodexResult; -use crate::event_mapping::map_response_item_to_event_messages; use crate::protocol::Event; use crate::protocol::EventMsg; use crate::protocol::SessionConfiguredEvent; @@ -187,13 +186,9 @@ fn truncate_after_dropping_last_messages(history: InitialHistory, n: usize) -> I return InitialHistory::New; } - // 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]); - - // Compute event prefix by cutting at the matching user event (if present). + // Compute event prefix by dropping the last `n` user events (counted from the end). let event_msgs_prefix: Vec = - event_msgs_prefix_until_target(&history, target_message.as_deref()); + event_msgs_prefix_after_dropping_last_user_events(&history, n); // Keep only response items strictly before the cut response index. let response_prefix: Vec = response_items[..cut_resp_index].to_vec(); @@ -202,21 +197,16 @@ fn truncate_after_dropping_last_messages(history: InitialHistory, n: usize) -> I 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( +/// Build the event messages prefix from `history` by dropping the last `n` user +/// events (counted from the end) and taking everything before that cut. +fn event_msgs_prefix_after_dropping_last_user_events( history: &InitialHistory, - target_message: Option<&str>, + n: usize, ) -> Vec { 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) - { - all_events[..idx].to_vec() - } else { - Vec::new() - } + if let Some(idx) = find_cut_event_index(&all_events, n) { + all_events[..idx].to_vec() } else { Vec::new() } @@ -244,24 +234,21 @@ fn find_cut_response_index(response_items: &[ResponseItem], n: usize) -> Option< None } -/// Derive the user message text (if any) associated with a response item using event mapping. -fn user_message_text_for_response(item: &ResponseItem) -> Option { - let mapped = map_response_item_to_event_messages(item, false); - mapped.into_iter().find_map(|ev| match ev { - EventMsg::UserMessage(u) => Some(u.message), - _ => None, - }) -} - -/// 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 { - event_msgs.iter().enumerate().find_map(|(i, it)| match it { - EventMsg::UserMessage(u) if u.message == target_message => Some(i), - _ => None, - }) +/// Find the index (into event messages) of the Nth user event from the end. +fn find_cut_event_index(event_msgs: &[EventMsg], n: usize) -> Option { + if n == 0 { + return None; + } + let mut remaining = n; + for (idx, ev) in event_msgs.iter().enumerate().rev() { + if matches!(ev, EventMsg::UserMessage(_)) { + remaining -= 1; + if remaining == 0 { + return Some(idx); + } + } + } + None } /// Build a truncated rollout by concatenating the (already-sliced) event messages and response items. @@ -283,6 +270,8 @@ fn build_truncated_rollout( #[cfg(test)] mod tests { use super::*; + use crate::event_mapping::map_response_item_to_event_messages; + use crate::protocol::EventMsg; use codex_protocol::models::ContentItem; use codex_protocol::models::ReasoningItemReasoningSummary; use codex_protocol::models::ResponseItem; @@ -296,6 +285,15 @@ mod tests { }], } } + fn user_input(text: &str) -> ResponseItem { + ResponseItem::Message { + id: None, + role: "user".to_string(), + content: vec![ContentItem::InputText { + text: text.to_string(), + }], + } + } fn assistant_msg(text: &str) -> ResponseItem { ResponseItem::Message { id: None, @@ -357,4 +355,54 @@ mod tests { let truncated2 = truncate_after_dropping_last_messages(InitialHistory::Forked(initial2), 2); assert!(matches!(truncated2, InitialHistory::New)); } + + #[test] + fn event_prefix_counts_from_end_with_duplicate_user_prompts() { + // Two identical user prompts with assistant replies between them. + let responses = vec![ + user_input("same"), + assistant_msg("a1"), + user_input("same"), + assistant_msg("a2"), + ]; + + // Derive event messages in order from responses (user → UserMessage, assistant → AgentMessage). + let mut events: Vec = Vec::new(); + for r in &responses { + events.extend(map_response_item_to_event_messages(r, false)); + } + + // Build initial history containing both events and responses. + let mut initial: Vec = Vec::new(); + initial.extend(events.iter().cloned().map(RolloutItem::EventMsg)); + initial.extend(responses.iter().cloned().map(RolloutItem::ResponseItem)); + + // Drop the last user turn. + let truncated = truncate_after_dropping_last_messages(InitialHistory::Forked(initial), 1); + + // Expect the event prefix to include the first user + first assistant only, + // and the response prefix to include the first user + first assistant only. + let got_items = truncated.get_rollout_items(); + + // Compute expected events and responses after cut. + let expected_event_prefix: Vec = events[..2] + .iter() + .cloned() + .map(RolloutItem::EventMsg) + .collect(); + let expected_response_prefix: Vec = responses[..2] + .iter() + .cloned() + .map(RolloutItem::ResponseItem) + .collect(); + + let mut expected: Vec = Vec::new(); + expected.extend(expected_event_prefix); + expected.extend(expected_response_prefix); + + assert_eq!( + serde_json::to_value(&got_items).unwrap(), + serde_json::to_value(&expected).unwrap() + ); + } }