diff --git a/codex-rs/app-server/tests/common/lib.rs b/codex-rs/app-server/tests/common/lib.rs index 4a2a99db23..8982bb11ae 100644 --- a/codex-rs/app-server/tests/common/lib.rs +++ b/codex-rs/app-server/tests/common/lib.rs @@ -28,6 +28,7 @@ pub use models_cache::write_models_cache; pub use models_cache::write_models_cache_with_models; pub use responses::create_apply_patch_sse_response; pub use responses::create_exec_command_sse_response; +pub use responses::create_exec_command_sse_response_for_command; pub use responses::create_final_assistant_message_sse_response; pub use responses::create_request_user_input_sse_response; pub use responses::create_shell_command_sse_response; diff --git a/codex-rs/app-server/tests/common/mcp_process.rs b/codex-rs/app-server/tests/common/mcp_process.rs index 249a280213..2b14a7be25 100644 --- a/codex-rs/app-server/tests/common/mcp_process.rs +++ b/codex-rs/app-server/tests/common/mcp_process.rs @@ -53,7 +53,9 @@ use codex_app_server_protocol::SetDefaultModelParams; use codex_app_server_protocol::SkillsListParams; use codex_app_server_protocol::ThreadArchiveParams; use codex_app_server_protocol::ThreadCompactStartParams; +use codex_app_server_protocol::ThreadDecrementElicitationParams; use codex_app_server_protocol::ThreadForkParams; +use codex_app_server_protocol::ThreadIncrementElicitationParams; use codex_app_server_protocol::ThreadListParams; use codex_app_server_protocol::ThreadLoadedListParams; use codex_app_server_protocol::ThreadReadParams; @@ -472,6 +474,26 @@ impl McpProcess { self.send_request("thread/read", params).await } + /// Send a `thread/increment_elicitation` JSON-RPC request. + pub async fn send_thread_increment_elicitation_request( + &mut self, + params: ThreadIncrementElicitationParams, + ) -> anyhow::Result { + let params = Some(serde_json::to_value(params)?); + self.send_request("thread/increment_elicitation", params) + .await + } + + /// Send a `thread/decrement_elicitation` JSON-RPC request. + pub async fn send_thread_decrement_elicitation_request( + &mut self, + params: ThreadDecrementElicitationParams, + ) -> anyhow::Result { + let params = Some(serde_json::to_value(params)?); + self.send_request("thread/decrement_elicitation", params) + .await + } + /// Send a `model/list` JSON-RPC request. pub async fn send_list_models_request( &mut self, diff --git a/codex-rs/app-server/tests/common/responses.rs b/codex-rs/app-server/tests/common/responses.rs index e15319e02f..6b47f4645e 100644 --- a/codex-rs/app-server/tests/common/responses.rs +++ b/codex-rs/app-server/tests/common/responses.rs @@ -50,9 +50,18 @@ pub fn create_exec_command_sse_response(call_id: &str) -> anyhow::Result let command = std::iter::once(cmd.to_string()) .chain(args.into_iter().map(str::to_string)) .collect::>(); + create_exec_command_sse_response_for_command(command, 500, call_id) +} + +pub fn create_exec_command_sse_response_for_command( + command: Vec, + yield_time_ms: u64, + call_id: &str, +) -> anyhow::Result { + let command_str = shlex::try_join(command.iter().map(String::as_str))?; let tool_call_arguments = serde_json::to_string(&json!({ - "cmd": command.join(" "), - "yield_time_ms": 500 + "cmd": command_str, + "yield_time_ms": yield_time_ms }))?; Ok(responses::sse(vec![ responses::ev_response_created("resp-1"), diff --git a/codex-rs/app-server/tests/fixtures/elicitation_stopwatch/orchestrator.py b/codex-rs/app-server/tests/fixtures/elicitation_stopwatch/orchestrator.py new file mode 100644 index 0000000000..1c95ebe61c --- /dev/null +++ b/codex-rs/app-server/tests/fixtures/elicitation_stopwatch/orchestrator.py @@ -0,0 +1,77 @@ +#!/usr/bin/env python3 + +import argparse +import sys +import time +from pathlib import Path + +REQUESTED_FILENAME = "elicitation_requested" +RELEASE_FILENAME = "elicitation_release" + + +def requested_path(state_dir: Path) -> Path: + return state_dir / REQUESTED_FILENAME + + +def release_path(state_dir: Path) -> Path: + return state_dir / RELEASE_FILENAME + + +def cmd_wait_for_request(state_dir: Path, timeout_seconds: float) -> int: + deadline = time.monotonic() + timeout_seconds + while time.monotonic() < deadline: + if requested_path(state_dir).exists(): + return 0 + time.sleep(0.05) + + print( + f"timed out waiting for {requested_path(state_dir)}", + file=sys.stderr, + ) + return 2 + + +def cmd_release(state_dir: Path) -> int: + release_path(state_dir).write_text("approved\n", encoding="utf-8") + return 0 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument( + "--state-dir", + required=True, + type=Path, + help="Directory shared with the elicitation trigger script.", + ) + subparsers = parser.add_subparsers(dest="command", required=True) + + wait_parser = subparsers.add_parser("wait-for-request") + wait_parser.add_argument( + "--timeout-seconds", + type=float, + default=5.0, + ) + + subparsers.add_parser("release") + + return parser.parse_args() + + +def main() -> int: + args = parse_args() + state_dir: Path = args.state_dir + state_dir.mkdir(parents=True, exist_ok=True) + + if args.command == "wait-for-request": + return cmd_wait_for_request(state_dir, args.timeout_seconds) + + if args.command == "release": + return cmd_release(state_dir) + + print(f"unsupported command: {args.command}", file=sys.stderr) + return 1 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/codex-rs/app-server/tests/fixtures/elicitation_stopwatch/trigger_elicitation.py b/codex-rs/app-server/tests/fixtures/elicitation_stopwatch/trigger_elicitation.py new file mode 100644 index 0000000000..c31edec3b9 --- /dev/null +++ b/codex-rs/app-server/tests/fixtures/elicitation_stopwatch/trigger_elicitation.py @@ -0,0 +1,43 @@ +#!/usr/bin/env python3 + +import argparse +import os +import sys +import time +from pathlib import Path + +REQUESTED_FILENAME = "elicitation_requested" +RELEASE_FILENAME = "elicitation_release" + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument( + "--state-dir", + required=True, + type=Path, + help="Directory shared with the test orchestrator.", + ) + return parser.parse_args() + + +def main() -> int: + args = parse_args() + state_dir: Path = args.state_dir + state_dir.mkdir(parents=True, exist_ok=True) + + requested = state_dir / REQUESTED_FILENAME + release = state_dir / RELEASE_FILENAME + + requested.write_text(f"pid={os.getpid()}\n", encoding="utf-8") + print("waited for a user approval", file=sys.stderr, flush=True) + + while not release.exists(): + time.sleep(0.05) + + print("approval received", flush=True) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) 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 d2d850176d..10ce739a0b 100644 --- a/codex-rs/app-server/tests/suite/v2/turn_start.rs +++ b/codex-rs/app-server/tests/suite/v2/turn_start.rs @@ -2,6 +2,7 @@ use anyhow::Result; use app_test_support::McpProcess; use app_test_support::create_apply_patch_sse_response; use app_test_support::create_exec_command_sse_response; +use app_test_support::create_exec_command_sse_response_for_command; use app_test_support::create_fake_rollout; use app_test_support::create_final_assistant_message_sse_response; use app_test_support::create_mock_responses_server_sequence; @@ -26,6 +27,10 @@ use codex_app_server_protocol::PatchChangeKind; use codex_app_server_protocol::RequestId; use codex_app_server_protocol::ServerRequest; use codex_app_server_protocol::TextElement; +use codex_app_server_protocol::ThreadDecrementElicitationParams; +use codex_app_server_protocol::ThreadDecrementElicitationResponse; +use codex_app_server_protocol::ThreadIncrementElicitationParams; +use codex_app_server_protocol::ThreadIncrementElicitationResponse; use codex_app_server_protocol::ThreadItem; use codex_app_server_protocol::ThreadStartParams; use codex_app_server_protocol::ThreadStartResponse; @@ -45,12 +50,14 @@ use codex_protocol::config_types::ModeKind; use codex_protocol::config_types::Personality; use codex_protocol::config_types::Settings; use codex_protocol::openai_models::ReasoningEffort; +use codex_utils_cargo_bin::find_resource; use core_test_support::responses; use core_test_support::skip_if_no_network; use pretty_assertions::assert_eq; use std::collections::BTreeMap; use std::path::Path; use tempfile::TempDir; +use tokio::process::Command; use tokio::time::timeout; #[cfg(windows)] @@ -1716,6 +1723,207 @@ async fn turn_start_file_change_approval_decline_v2() -> Result<()> { Ok(()) } +#[tokio::test] +#[cfg_attr( + windows, + ignore = "relies on local Python fixture scripts and POSIX unified exec timing" +)] +async fn thread_elicitation_pauses_unified_exec_stopwatch() -> Result<()> { + let tmp = TempDir::new()?; + let codex_home = tmp.path().join("codex_home"); + std::fs::create_dir(&codex_home)?; + let workspace = tmp.path().join("workspace"); + std::fs::create_dir(&workspace)?; + let state_dir = tmp.path().join("elicitation_state"); + std::fs::create_dir(&state_dir)?; + + let orchestrator_script = + find_resource!("tests/fixtures/elicitation_stopwatch/orchestrator.py")?; + let trigger_script = + find_resource!("tests/fixtures/elicitation_stopwatch/trigger_elicitation.py")?; + + let responses = vec![ + create_exec_command_sse_response_for_command( + vec![ + "python3".to_string(), + trigger_script.to_string_lossy().to_string(), + "--state-dir".to_string(), + state_dir.to_string_lossy().to_string(), + ], + 30_000, + "uexec-elicitation", + )?, + create_final_assistant_message_sse_response("done")?, + ]; + let server = create_mock_responses_server_sequence(responses).await; + create_config_toml_with_sandbox( + &codex_home, + &server.uri(), + "never", + &BTreeMap::from([(Feature::UnifiedExec, true)]), + "danger-full-access", + )?; + + let mut mcp = McpProcess::new(&codex_home).await?; + timeout(DEFAULT_READ_TIMEOUT, mcp.initialize()).await??; + + let start_id = mcp + .send_thread_start_request(ThreadStartParams { + model: Some("mock-model".to_string()), + ..Default::default() + }) + .await?; + let start_resp: JSONRPCResponse = timeout( + DEFAULT_READ_TIMEOUT, + mcp.read_stream_until_response_message(RequestId::Integer(start_id)), + ) + .await??; + let ThreadStartResponse { thread, .. } = to_response::(start_resp)?; + + let turn_id = mcp + .send_turn_start_request(TurnStartParams { + thread_id: thread.id.clone(), + input: vec![V2UserInput::Text { + text: "run the local approval fixture".to_string(), + text_elements: Vec::new(), + }], + cwd: Some(workspace.clone()), + sandbox_policy: Some(codex_app_server_protocol::SandboxPolicy::DangerFullAccess), + ..Default::default() + }) + .await?; + timeout( + DEFAULT_READ_TIMEOUT, + mcp.read_stream_until_response_message(RequestId::Integer(turn_id)), + ) + .await??; + + let started_command = 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 let ThreadItem::CommandExecution { .. } = started.item { + return Ok::(started.item); + } + } + }) + .await??; + let ThreadItem::CommandExecution { id, status, .. } = started_command else { + unreachable!("loop ensures we break on command execution items"); + }; + assert_eq!(id, "uexec-elicitation"); + assert_eq!(status, CommandExecutionStatus::InProgress); + + run_elicitation_orchestrator( + &orchestrator_script, + &state_dir, + "wait-for-request", + Some(5), + ) + .await?; + + let increment_request_id = mcp + .send_thread_increment_elicitation_request(ThreadIncrementElicitationParams { + thread_id: thread.id.clone(), + }) + .await?; + let increment_response: JSONRPCResponse = timeout( + DEFAULT_READ_TIMEOUT, + mcp.read_stream_until_response_message(RequestId::Integer(increment_request_id)), + ) + .await??; + let ThreadIncrementElicitationResponse { count, paused } = + to_response::(increment_response)?; + assert_eq!(count, 1); + assert!(paused); + + // Hold longer than the default 10s unified-exec timeout. If the stopwatch is not paused, + // the command exits/times out and the model will issue the second /responses request. + assert!( + timeout( + std::time::Duration::from_secs(11), + read_completed_command_execution_item(&mut mcp), + ) + .await + .is_err(), + "command execution should remain in progress while elicitation is active" + ); + let requests_during_pause = server + .received_requests() + .await + .expect("failed to fetch received requests while paused"); + assert_eq!( + requests_during_pause.len(), + 1, + "unexpected extra inference request while elicitation pause was active" + ); + + let decrement_request_id = mcp + .send_thread_decrement_elicitation_request(ThreadDecrementElicitationParams { + thread_id: thread.id.clone(), + }) + .await?; + let decrement_response: JSONRPCResponse = timeout( + DEFAULT_READ_TIMEOUT, + mcp.read_stream_until_response_message(RequestId::Integer(decrement_request_id)), + ) + .await??; + let ThreadDecrementElicitationResponse { count, paused } = + to_response::(decrement_response)?; + assert_eq!(count, 0); + assert!(!paused); + + run_elicitation_orchestrator(&orchestrator_script, &state_dir, "release", None).await?; + + let completed_command = timeout( + DEFAULT_READ_TIMEOUT, + read_completed_command_execution_item(&mut mcp), + ) + .await??; + let ThreadItem::CommandExecution { + id, + status, + exit_code, + aggregated_output, + .. + } = completed_command + else { + unreachable!("helper only returns command execution items"); + }; + assert_eq!(id, "uexec-elicitation"); + assert_eq!(status, CommandExecutionStatus::Completed); + assert_eq!(exit_code, Some(0)); + let aggregated_output = aggregated_output.expect("expected command output"); + assert!(aggregated_output.contains("waited for a user approval")); + assert!(aggregated_output.contains("approval received")); + + timeout( + DEFAULT_READ_TIMEOUT, + mcp.read_stream_until_notification_message("codex/event/task_complete"), + ) + .await??; + timeout( + DEFAULT_READ_TIMEOUT, + mcp.read_stream_until_notification_message("turn/completed"), + ) + .await??; + + let requests_after_resume = server + .received_requests() + .await + .expect("failed to fetch received requests after resume"); + assert_eq!(requests_after_resume.len(), 2); + + Ok(()) +} + #[tokio::test] #[cfg_attr(windows, ignore = "process id reporting differs on Windows")] async fn command_execution_notifications_include_process_id() -> Result<()> { @@ -1853,6 +2061,55 @@ async fn command_execution_notifications_include_process_id() -> Result<()> { Ok(()) } +async fn run_elicitation_orchestrator( + orchestrator_script: &Path, + state_dir: &Path, + command: &str, + timeout_seconds: Option, +) -> Result<()> { + let mut cmd = Command::new("python3"); + cmd.arg(orchestrator_script) + .arg("--state-dir") + .arg(state_dir) + .arg(command); + + if let Some(timeout_seconds) = timeout_seconds { + cmd.arg("--timeout-seconds") + .arg(timeout_seconds.to_string()); + } + + let output = cmd.output().await?; + if output.status.success() { + return Ok(()); + } + + let stdout = String::from_utf8_lossy(&output.stdout); + let stderr = String::from_utf8_lossy(&output.stderr); + anyhow::bail!( + "orchestrator command `{command}` failed with status {:?}\nstdout:\n{}\nstderr:\n{}", + output.status.code(), + stdout, + stderr, + ); +} + +async fn read_completed_command_execution_item(mcp: &mut McpProcess) -> Result { + loop { + let notif = mcp + .read_stream_until_notification_message("item/completed") + .await?; + let completed: ItemCompletedNotification = serde_json::from_value( + notif + .params + .clone() + .expect("item/completed should include params"), + )?; + if let ThreadItem::CommandExecution { .. } = completed.item { + return Ok(completed.item); + } + } +} + // Helper to create a config.toml pointing at the mock model server. fn create_config_toml( codex_home: &Path,