diff --git a/codex-rs/app-server/tests/suite/v2/turn_start.rs b/codex-rs/app-server/tests/suite/v2/turn_start.rs index a89ca359d8..ea46d06a5c 100644 --- a/codex-rs/app-server/tests/suite/v2/turn_start.rs +++ b/codex-rs/app-server/tests/suite/v2/turn_start.rs @@ -321,6 +321,169 @@ async fn turn_start_forwards_ephemeral_context_to_model_input() -> Result<()> { Ok(()) } +#[tokio::test] +async fn turn_start_steering_active_turn_preserves_ephemeral_context() -> Result<()> { + #[cfg(target_os = "windows")] + let shell_command = vec![ + "powershell".to_string(), + "-Command".to_string(), + "Start-Sleep -Seconds 2".to_string(), + ]; + #[cfg(not(target_os = "windows"))] + let shell_command = vec!["sleep".to_string(), "2".to_string()]; + + let responses = vec![ + create_shell_command_sse_response(shell_command, None, Some(5_000), "call-sleep")?, + create_final_assistant_message_sse_response("Done")?, + ]; + let server = create_mock_responses_server_sequence_unchecked(responses).await; + + let codex_home = TempDir::new()?; + create_config_toml( + codex_home.path(), + &server.uri(), + "never", + &BTreeMap::from([(Feature::Personality, true)]), + )?; + + let mut mcp = McpProcess::new(codex_home.path()).await?; + timeout(DEFAULT_READ_TIMEOUT, mcp.initialize()).await??; + + let thread_req = mcp + .send_thread_start_request(ThreadStartParams { + model: Some("mock-model".to_string()), + ..Default::default() + }) + .await?; + let thread_resp: JSONRPCResponse = timeout( + DEFAULT_READ_TIMEOUT, + mcp.read_stream_until_response_message(RequestId::Integer(thread_req)), + ) + .await??; + let ThreadStartResponse { thread, .. } = to_response::(thread_resp)?; + + let first_turn_req = mcp + .send_turn_start_request(TurnStartParams { + thread_id: thread.id.clone(), + input: vec![V2UserInput::Text { + text: "run sleep".to_string(), + text_elements: Vec::new(), + }], + ..Default::default() + }) + .await?; + let first_turn_resp: JSONRPCResponse = timeout( + DEFAULT_READ_TIMEOUT, + mcp.read_stream_until_response_message(RequestId::Integer(first_turn_req)), + ) + .await??; + let TurnStartResponse { turn: _first_turn } = + to_response::(first_turn_resp)?; + + timeout(DEFAULT_READ_TIMEOUT, async { + loop { + let notif = mcp + .read_stream_until_notification_message("item/started") + .await?; + let started: ItemStartedNotification = serde_json::from_value( + notif + .params + .clone() + .expect("item/started should include params"), + )?; + if matches!(started.item, ThreadItem::CommandExecution { .. }) { + return Ok::<(), anyhow::Error>(()); + } + } + }) + .await??; + + let steering_turn_req = mcp + .send_turn_start_request(TurnStartParams { + thread_id: thread.id, + input: vec![V2UserInput::Text { + text: "Follow-up".to_string(), + text_elements: Vec::new(), + }], + ephemeral_context: Some(vec![EphemeralContext { + title: "Context from my editor".to_string(), + text: "## Active file: src/main.rs".to_string(), + }]), + ..Default::default() + }) + .await?; + let steering_turn_resp: JSONRPCResponse = timeout( + DEFAULT_READ_TIMEOUT, + mcp.read_stream_until_response_message(RequestId::Integer(steering_turn_req)), + ) + .await??; + let TurnStartResponse { + turn: _steering_turn, + } = to_response::(steering_turn_resp)?; + + timeout( + DEFAULT_READ_TIMEOUT, + mcp.read_stream_until_notification_message("turn/completed"), + ) + .await??; + + timeout(DEFAULT_READ_TIMEOUT, async { + loop { + let requests = server + .received_requests() + .await + .expect("failed to fetch received requests"); + let responses_request_count = requests + .iter() + .filter(|request| { + request.method == "POST" && request.url.path().ends_with("/responses") + }) + .count(); + if responses_request_count == 2 { + return Ok::<(), anyhow::Error>(()); + } + if responses_request_count > 2 { + anyhow::bail!( + "expected exactly 2 /responses requests, got {responses_request_count}" + ); + } + tokio::time::sleep(std::time::Duration::from_millis(10)).await; + } + }) + .await??; + + let requests = server + .received_requests() + .await + .expect("failed to fetch received requests"); + let responses_requests = requests + .iter() + .filter(|request| request.method == "POST" && request.url.path().ends_with("/responses")) + .collect::>(); + assert_eq!(responses_requests.len(), 2); + assert!( + body_contains(responses_requests[1], ""), + "expected steered follow-up request to contain additional_context_for_this_turn wrapper" + ); + assert!( + body_contains( + responses_requests[1], + "Context from my editor" + ), + "expected steered follow-up request to contain ephemeral context title" + ); + assert!( + body_contains(responses_requests[1], "## Active file: src/main.rs"), + "expected steered follow-up request to contain ephemeral context body" + ); + assert!( + body_contains(responses_requests[1], "Follow-up"), + "expected steered follow-up request to contain the steered user input" + ); + + Ok(()) +} + #[tokio::test] async fn turn_start_accepts_text_at_limit_with_mention_item() -> Result<()> { let responses = vec![create_final_assistant_message_sse_response("Done")?]; diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index 15409efb83..5509df58a1 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -3747,6 +3747,16 @@ impl Session { &self, input: Vec, expected_turn_id: Option<&str>, + ) -> Result { + self.steer_input_with_ephemeral_context(input, &[], expected_turn_id) + .await + } + + async fn steer_input_with_ephemeral_context( + &self, + input: Vec, + ephemeral_context: &[EphemeralContext], + expected_turn_id: Option<&str>, ) -> Result { if input.is_empty() { return Err(SteerInputError::EmptyInput); @@ -3770,8 +3780,24 @@ impl Session { }); } + let response_input = if ephemeral_context.is_empty() { + input.into() + } else { + let ResponseInputItem::Message { role, content } = ResponseInputItem::from(input) + else { + unreachable!("user input should always convert into a message"); + }; + let mut prefixed_content = + crate::model_visible_fragments::ephemeral_context_content_items(ephemeral_context); + prefixed_content.extend(content); + ResponseInputItem::Message { + role, + content: prefixed_content, + } + }; + let mut turn_state = active_turn.turn_state.lock().await; - turn_state.push_pending_input(input.into()); + turn_state.push_pending_input(response_input); Ok(active_turn_id.clone()) } @@ -4442,7 +4468,10 @@ mod handlers { current_context.session_telemetry.user_prompt(&items); // Attempt to inject input into current task. - if let Err(SteerInputError::NoActiveTurn(items)) = sess.steer_input(items, None).await { + if let Err(SteerInputError::NoActiveTurn(items)) = sess + .steer_input_with_ephemeral_context(items, ¤t_context.ephemeral_context, None) + .await + { sess.refresh_mcp_servers_if_requested(¤t_context) .await; let regular_task = sess.take_startup_regular_task().await.unwrap_or_default(); diff --git a/codex-rs/core/src/codex_tests.rs b/codex-rs/core/src/codex_tests.rs index 0a95bd2907..8bc9998d0e 100644 --- a/codex-rs/core/src/codex_tests.rs +++ b/codex-rs/core/src/codex_tests.rs @@ -66,9 +66,12 @@ use codex_protocol::models::ResponseItem; use codex_protocol::models::developer_personality_spec_text; use codex_protocol::openai_models::ModelsResponse; use codex_protocol::protocol::ConversationAudioParams; +use codex_protocol::protocol::EPHEMERAL_CONTEXT_CLOSE_TAG; +use codex_protocol::protocol::EPHEMERAL_CONTEXT_OPEN_TAG; use codex_protocol::protocol::RealtimeAudioFrame; use codex_protocol::protocol::Submission; use codex_protocol::protocol::W3cTraceContext; +use codex_protocol::user_input::EphemeralContext; use opentelemetry::trace::TraceContextExt; use opentelemetry::trace::TraceId; use opentelemetry::trace::TracerProvider as _; @@ -3973,6 +3976,58 @@ async fn steer_input_returns_active_turn_id() { assert!(sess.has_pending_input().await); } +#[tokio::test] +async fn user_input_or_turn_steering_preserves_ephemeral_context() { + let (sess, tc, _rx) = make_session_and_context_with_rx().await; + let input = vec![UserInput::Text { + text: "hello".to_string(), + text_elements: Vec::new(), + }]; + sess.spawn_task( + Arc::clone(&tc), + input, + NeverEndingTask { + kind: TaskKind::Regular, + listen_to_cancellation_token: false, + }, + ) + .await; + + handlers::user_input_or_turn( + &sess, + "steer-submission".to_string(), + Op::UserInput { + items: vec![UserInput::Text { + text: "steer".to_string(), + text_elements: Vec::new(), + }], + ephemeral_context: vec![EphemeralContext { + title: "Context from my editor".to_string(), + text: "## Active file: src/main.rs".to_string(), + }], + final_output_json_schema: None, + }, + ) + .await; + + assert_eq!( + sess.get_pending_input().await, + vec![ResponseInputItem::Message { + role: "user".to_string(), + content: vec![ + ContentItem::InputText { + text: format!( + "{EPHEMERAL_CONTEXT_OPEN_TAG}\n Context from my editor\n \n## Active file: src/main.rs\n \n{EPHEMERAL_CONTEXT_CLOSE_TAG}" + ), + }, + ContentItem::InputText { + text: "steer".to_string(), + }, + ], + }] + ); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn abort_review_task_emits_exited_then_aborted_and_records_history() { let (sess, tc, rx) = make_session_and_context_with_rx().await; diff --git a/codex-rs/core/src/model_visible_fragments.rs b/codex-rs/core/src/model_visible_fragments.rs index 9cd40f1210..a8e98d45cb 100644 --- a/codex-rs/core/src/model_visible_fragments.rs +++ b/codex-rs/core/src/model_visible_fragments.rs @@ -80,6 +80,7 @@ use codex_protocol::protocol::TurnContextItem; use codex_protocol::protocol::TurnContextNetworkItem; use codex_protocol::protocol::USER_INSTRUCTIONS_CLOSE_TAG; use codex_protocol::protocol::USER_INSTRUCTIONS_OPEN_TAG; +use codex_protocol::user_input::EphemeralContext as ProtocolEphemeralContext; use serde::Deserialize; use serde::Serialize; use std::path::PathBuf; @@ -1068,6 +1069,19 @@ pub(crate) fn is_ephemeral_context_fragment(content_item: &ContentItem) -> bool EphemeralContextFragment::matches_contextual_user_text(text) } +pub(crate) fn ephemeral_context_content_items( + ephemeral_context: &[ProtocolEphemeralContext], +) -> Vec { + ephemeral_context + .iter() + .map(|context| EphemeralContextFragment { + title: context.title.clone(), + text: context.text.clone(), + }) + .map(ModelVisibleContextFragment::into_content_item) + .collect() +} + pub(crate) fn build_turn_state_fragments( reference_context_item: Option<&TurnContextItem>, turn_context: &TurnContext,