diff --git a/codex-rs/app-server/src/codex_message_processor.rs b/codex-rs/app-server/src/codex_message_processor.rs index 88d9a885a9..818a034daa 100644 --- a/codex-rs/app-server/src/codex_message_processor.rs +++ b/codex-rs/app-server/src/codex_message_processor.rs @@ -3439,8 +3439,7 @@ impl CodexMessageProcessor { }; let fallback_model_provider = config.model_provider_id.clone(); - let fork_cutoff_nth_user_message = if let Some(fork_after_turn_id) = - fork_after_turn_id.as_deref() + let anchored_fork_history = if let Some(fork_after_turn_id) = fork_after_turn_id.as_deref() { let source_items = match read_rollout_items_from_rollout(rollout_path.as_path()).await { Ok(items) => items, @@ -3468,34 +3467,43 @@ impl CodexMessageProcessor { } }; - match resolve_thread_fork_cutoff_nth_user_message( + match resolve_thread_fork_history_from_anchor( source_items.as_slice(), fork_after_turn_id, ) { - Ok(cutoff) => cutoff, + Ok(history_items) => Some(InitialHistory::Forked(history_items)), Err(message) => { self.send_invalid_request_error(request_id, message).await; return; } } } else { - usize::MAX + None }; let NewThread { thread_id, session_configured, .. - } = match self - .thread_manager - .fork_thread( - fork_cutoff_nth_user_message, - config, - rollout_path.clone(), - persist_extended_history, - ) - .await - { + } = match if let Some(initial_history) = anchored_fork_history { + self.thread_manager + .resume_thread_with_history( + config, + initial_history, + self.auth_manager.clone(), + persist_extended_history, + ) + .await + } else { + self.thread_manager + .fork_thread( + usize::MAX, + config, + rollout_path.clone(), + persist_extended_history, + ) + .await + } { Ok(thread) => thread, Err(err) => { let (code, message) = match err { @@ -6913,10 +6921,10 @@ async fn sync_default_client_residency_requirement( } } -fn resolve_thread_fork_cutoff_nth_user_message( +fn resolve_thread_fork_history_from_anchor( rollout_items: &[RolloutItem], fork_after_turn_id: &str, -) -> Result { +) -> Result, String> { let history = build_thread_history_from_rollout_items(rollout_items); if history.has_synthetic_turn_ids { return Err( @@ -6925,55 +6933,85 @@ fn resolve_thread_fork_cutoff_nth_user_message( ); } - let total_user_turns = history + let Some(target_turn_idx) = history .turns .iter() - .filter(|turn| { - turn.items - .iter() - .any(|item| matches!(item, ThreadItem::UserMessage { .. })) - }) - .count(); - let mut user_turn_index = 0usize; + .position(|turn| turn.id == fork_after_turn_id) + else { + return Err(format!( + "fork turn not found in source thread history: {fork_after_turn_id}" + )); + }; + let target_turn = &history.turns[target_turn_idx]; + let has_user_message = target_turn + .items + .iter() + .any(|item| matches!(item, ThreadItem::UserMessage { .. })); + let has_agent_message = target_turn + .items + .iter() + .any(|item| matches!(item, ThreadItem::AgentMessage { .. })); - for turn in &history.turns { - let has_user_message = turn - .items - .iter() - .any(|item| matches!(item, ThreadItem::UserMessage { .. })); - let has_agent_message = turn - .items - .iter() - .any(|item| matches!(item, ThreadItem::AgentMessage { .. })); - - if turn.id == fork_after_turn_id { - if matches!(turn.status, TurnStatus::InProgress) { - return Err( - "fork turn must be completed/interrupted/failed, not in progress".to_string(), - ); - } - if !has_user_message { - return Err("fork turn must contain a user message".to_string()); - } - if !has_agent_message { - return Err("fork turn must contain an agent message".to_string()); - } - - return if user_turn_index.saturating_add(1) >= total_user_turns { - Ok(usize::MAX) - } else { - Ok(user_turn_index.saturating_add(1)) - }; - } - - if has_user_message { - user_turn_index = user_turn_index.saturating_add(1); - } + if matches!(target_turn.status, TurnStatus::InProgress) { + return Err("fork turn must be completed/interrupted/failed, not in progress".to_string()); + } + if !has_user_message { + return Err("fork turn must contain a user message".to_string()); + } + if !has_agent_message { + return Err("fork turn must contain an agent message".to_string()); } - Err(format!( - "fork turn not found in source thread history: {fork_after_turn_id}" - )) + let Some(next_turn) = history.turns.get(target_turn_idx.saturating_add(1)) else { + return Ok(rollout_items.to_vec()); + }; + let next_turn_id = next_turn.id.as_str(); + + let next_turn_started_idx = rollout_items + .iter() + .position(|item| { + matches!( + item, + RolloutItem::EventMsg(EventMsg::TurnStarted(payload)) + if payload.turn_id == next_turn_id + ) + }) + .ok_or_else(|| { + format!("failed to locate boundary after fork turn: missing next turn `{next_turn_id}`") + })?; + + let target_terminal_idx = rollout_items.iter().rposition(|item| { + matches!( + item, + RolloutItem::EventMsg(EventMsg::TurnComplete(payload)) + if payload.turn_id == fork_after_turn_id + ) || matches!( + item, + RolloutItem::EventMsg(EventMsg::TurnAborted(payload)) + if payload.turn_id.as_deref() == Some(fork_after_turn_id) + ) + }); + + let next_user_boundary_idx = target_terminal_idx + .and_then(|terminal_idx| { + rollout_items + .iter() + .enumerate() + .skip(terminal_idx.saturating_add(1)) + .find_map(|(idx, item)| is_user_turn_rollout_boundary(item).then_some(idx)) + }) + .unwrap_or(usize::MAX); + let cut_idx = next_turn_started_idx.min(next_user_boundary_idx); + + Ok(rollout_items[..cut_idx].to_vec()) +} + +fn is_user_turn_rollout_boundary(item: &RolloutItem) -> bool { + match item { + RolloutItem::ResponseItem(ResponseItem::Message { role, .. }) => role == "user", + RolloutItem::EventMsg(EventMsg::UserMessage(_)) => true, + _ => false, + } } /// Derive the effective [`Config`] by layering three override sources. diff --git a/codex-rs/app-server/tests/suite/v2/thread_fork.rs b/codex-rs/app-server/tests/suite/v2/thread_fork.rs index f96908073c..5f6bb0698f 100644 --- a/codex-rs/app-server/tests/suite/v2/thread_fork.rs +++ b/codex-rs/app-server/tests/suite/v2/thread_fork.rs @@ -220,20 +220,21 @@ async fn thread_fork_can_fork_after_selected_turn() -> Result<()> { &[ ExplicitTurnFixture { turn_id: "turn-1", - user_text: "u1", + user_text: Some("u1"), agent_text: Some("a1"), state: FixtureTurnState::Completed, }, ExplicitTurnFixture { turn_id: "turn-2", - user_text: "u2", + user_text: Some("u2"), agent_text: Some("a2"), state: FixtureTurnState::Completed, }, ExplicitTurnFixture { - turn_id: "turn-3", - user_text: "u3", - agent_text: Some("a3"), + // Explicit turn with no user message exercises the exact cut-after-turn logic. + turn_id: "turn-3-empty", + user_text: None, + agent_text: None, state: FixtureTurnState::Completed, }, ], @@ -282,7 +283,7 @@ async fn thread_fork_rejects_unknown_turn_anchor() -> Result<()> { "2025-01-05T12:00:02Z", &[ExplicitTurnFixture { turn_id: "turn-1", - user_text: "u1", + user_text: Some("u1"), agent_text: Some("a1"), state: FixtureTurnState::Completed, }], @@ -363,7 +364,7 @@ async fn thread_fork_rejects_in_progress_turn_anchor() -> Result<()> { "2025-01-05T12:00:04Z", &[ExplicitTurnFixture { turn_id: "turn-in-progress", - user_text: "u1", + user_text: Some("u1"), agent_text: Some("a1"), state: FixtureTurnState::InProgress, }], @@ -405,7 +406,7 @@ async fn thread_fork_rejects_turn_anchor_without_agent_message() -> Result<()> { "2025-01-05T12:00:05Z", &[ExplicitTurnFixture { turn_id: "turn-no-agent", - user_text: "u1", + user_text: Some("u1"), agent_text: None, state: FixtureTurnState::Completed, }], @@ -468,7 +469,7 @@ enum FixtureTurnState { #[derive(Clone, Copy)] struct ExplicitTurnFixture<'a> { turn_id: &'a str, - user_text: &'a str, + user_text: Option<&'a str>, agent_text: Option<&'a str>, state: FixtureTurnState, } @@ -521,18 +522,20 @@ fn create_fake_rollout_with_explicit_turns( )?]; for (idx, turn) in turns.iter().enumerate() { - lines.push( - json!({ - "timestamp": meta_rfc3339, - "type":"response_item", - "payload": { - "type":"message", - "role":"user", - "content":[{"type":"input_text","text": turn.user_text}] - } - }) - .to_string(), - ); + if let Some(user_text) = turn.user_text { + lines.push( + json!({ + "timestamp": meta_rfc3339, + "type":"response_item", + "payload": { + "type":"message", + "role":"user", + "content":[{"type":"input_text","text": user_text}] + } + }) + .to_string(), + ); + } lines.push(rollout_line( meta_rfc3339, @@ -542,15 +545,17 @@ fn create_fake_rollout_with_explicit_turns( collaboration_mode_kind: Default::default(), })), )?); - lines.push(rollout_line( - meta_rfc3339, - RolloutItem::EventMsg(EventMsg::UserMessage(UserMessageEvent { - message: turn.user_text.to_string(), - images: None, - local_images: Vec::new(), - text_elements: Vec::new(), - })), - )?); + if let Some(user_text) = turn.user_text { + lines.push(rollout_line( + meta_rfc3339, + RolloutItem::EventMsg(EventMsg::UserMessage(UserMessageEvent { + message: user_text.to_string(), + images: None, + local_images: Vec::new(), + text_elements: Vec::new(), + })), + )?); + } if let Some(agent_text) = turn.agent_text { lines.push(rollout_line( meta_rfc3339, @@ -571,7 +576,7 @@ fn create_fake_rollout_with_explicit_turns( )?); } - if idx == 0 { + if idx == 0 && turn.agent_text.is_some() { lines.push( json!({ "timestamp": meta_rfc3339,