Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions codex-rs/exec-server/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,8 @@ mod server;
mod shell_snapshot;
#[cfg(unix)]
mod shell_snapshot_file;
#[cfg(unix)]
mod shell_snapshot_process;
mod telemetry;
mod trace_context;
mod websocket_pong_watchdog;
Expand Down
9 changes: 5 additions & 4 deletions codex-rs/exec-server/src/shell_snapshot.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ use crate::protocol::ExecParams;
use crate::protocol::ShellSnapshotRequest;
use crate::rpc::internal_error;
use crate::rpc::invalid_params;
use crate::shell_snapshot_process::SnapshotCapture;
use crate::telemetry::ExecServerTelemetry;

const MAX_CACHED_SNAPSHOTS: usize = 16;
Expand Down Expand Up @@ -359,8 +360,7 @@ async fn capture_snapshot(
.envs(&prepared.env)
.stdin(Stdio::null())
.stdout(Stdio::piped())
.stderr(Stdio::null())
.kill_on_drop(true);
.stderr(Stdio::null());
if let Some(arg0) = &prepared.arg0 {
command.arg0(arg0);
}
Expand All @@ -377,7 +377,7 @@ async fn capture_snapshot(
});
}
}
let mut child = command.spawn().map_err(|err| {
let mut child = SnapshotCapture::spawn(&mut command).map_err(|err| {
(
"spawn_failed",
internal_error(format!("cannot capture shell snapshot: {err}")),
Expand Down Expand Up @@ -410,7 +410,7 @@ async fn capture_snapshot(
)),
));
}
let status = child.wait().await.map_err(|err| {
let status = child.wait_for_exit().await.map_err(|err| {
(
"wait_failed",
internal_error(format!("cannot finish shell snapshot: {err}")),
Expand Down Expand Up @@ -452,6 +452,7 @@ async fn capture_snapshot(
})?;
let mut snapshot = parse_snapshot(shell_type, captured, params.env_policy.as_ref())?;
snapshot.file_source = flag == b"1\0";
child.preserve_helpers();
Ok(snapshot)
}

Expand Down
103 changes: 103 additions & 0 deletions codex-rs/exec-server/src/shell_snapshot_process.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
//! Own a capture process group until its output and exit status are validated.
//! Keep the leader unreaped while cleanup is armed, so its group ID cannot be
//! reused. Successful captures release startup helpers; escaped groups and
//! executor death are outside this guard's scope. This includes bubblewrap's
//! separate session when it inherits the host PID namespace.

use std::io;
use std::os::unix::process::ExitStatusExt;
use std::process::ExitStatus;

use codex_utils_pty::process_group::kill_process_group;
use tokio::process::Child;
use tokio::process::ChildStdout;
use tokio::process::Command;
use tokio::signal::unix::Signal;
use tokio::signal::unix::SignalKind;
use tokio::signal::unix::signal;

pub(super) struct SnapshotCapture {
child: Child,
sigchld: Signal,
cleanup: bool,
pub(super) stdout: Option<ChildStdout>,
}

impl SnapshotCapture {
pub(super) fn spawn(command: &mut Command) -> io::Result<Self> {
// Subscribe before spawning so a fast exit cannot lose its notification.
let sigchld = signal(SignalKind::child())?;
// Group cleanup must not signal the executor or other captures.
let mut child = command.process_group(/*pgroup*/ 0).spawn()?;
let stdout = child.stdout.take();
Ok(Self {
child,
sigchld,
cleanup: true,
stdout,
})
}

pub(super) async fn wait_for_exit(&mut self) -> io::Result<ExitStatus> {
let pid = self
.child
.id()
.ok_or_else(|| io::Error::other("missing capture PID"))?;
loop {
{
let mut info = std::mem::MaybeUninit::<libc::siginfo_t>::zeroed();
// SAFETY: we own this child and provide writable storage. WNOWAIT
// reserves its PID until validation either releases or kills the group.
if unsafe {
libc::waitid(
libc::P_PID,
pid,
info.as_mut_ptr(),
libc::WEXITED | libc::WNOHANG | libc::WNOWAIT,
)
} == -1
{
let error = io::Error::last_os_error();
if error.kind() == io::ErrorKind::Interrupted {
continue;
}
if error.raw_os_error() == Some(libc::ECHILD) {
self.cleanup = false;
}
return Err(error);
}
// SAFETY: successful waitid initialized the zeroed signal information.
let info = unsafe { info.assume_init() };
if unsafe { info.si_pid() } != 0 {
let status = unsafe { info.si_status() };
return Ok(ExitStatus::from_raw(if info.si_code == libc::CLD_EXITED {
status << 8
} else if info.si_code == libc::CLD_DUMPED {
status | 0x80
} else {
status
}));
}
}
self.sigchld
.recv()
.await
.ok_or_else(|| io::Error::other("SIGCHLD stream closed"))?;
}
}

pub(super) fn preserve_helpers(&mut self) {
self.cleanup = false;
}
}

impl Drop for SnapshotCapture {
fn drop(&mut self) {
if self.cleanup
&& let Some(pid) = self.child.id()
{
let _ = kill_process_group(pid);
}
// Tokio reaps the child after group cleanup, including on cancellation.
}
}
90 changes: 75 additions & 15 deletions codex-rs/exec-server/tests/exec_process.rs
Original file line number Diff line number Diff line change
Expand Up @@ -458,21 +458,34 @@ async fn shell_snapshot_v2_remote_managed_proxy_uses_prepared_execution_context(
}

#[cfg(unix)]
#[test_case(false, false, "bash", 1; "local_pipe_recovery")]
#[test_case(false, true, "bash", 1; "local_tty_recovery")]
#[test_case(true, false, "bash", 1; "remote_pipe_recovery")]
#[test_case(true, true, "bash", 1; "remote_tty_recovery")]
#[test_case(false, false, "bash", 3; "local_retry_budget_exhausted")]
#[test_case(true, false, "bash", 3; "remote_retry_budget_exhausted")]
#[cfg_attr(target_os = "macos", test_case(false, false, "zsh", 1; "local_zsh_recovery"))]
#[cfg_attr(target_os = "macos", test_case(true, false, "zsh", 1; "remote_zsh_recovery"))]
#[derive(Clone, Copy)]
enum CaptureFailure {
Exit(i32),
Cancel,
Timeout,
}

#[cfg(unix)]
#[test_case(false, false, "bash", 1, CaptureFailure::Exit(7); "local_pipe_recovery")]
#[test_case(false, true, "bash", 1, CaptureFailure::Exit(7); "local_tty_recovery")]
#[test_case(true, false, "bash", 1, CaptureFailure::Exit(7); "remote_pipe_recovery")]
#[test_case(true, true, "bash", 1, CaptureFailure::Exit(7); "remote_tty_recovery")]
#[test_case(false, false, "bash", 3, CaptureFailure::Exit(7); "local_retry_budget_exhausted")]
#[test_case(true, false, "bash", 3, CaptureFailure::Exit(7); "remote_retry_budget_exhausted")]
#[test_case(false, false, "bash", 1, CaptureFailure::Exit(0); "local_invalid_output_recovery")]
#[cfg_attr(target_os = "macos", test_case(false, false, "zsh", 1, CaptureFailure::Exit(0); "local_zsh_invalid_output_recovery"))]
#[cfg_attr(target_os = "macos", test_case(false, false, "zsh", 1, CaptureFailure::Exit(7); "local_zsh_recovery"))]
#[cfg_attr(target_os = "macos", test_case(true, false, "zsh", 1, CaptureFailure::Exit(7); "remote_zsh_recovery"))]
#[test_case(false, false, "bash", 1, CaptureFailure::Cancel; "cancel_capture")]
#[test_case(false, false, "bash", 1, CaptureFailure::Timeout; "timeout_capture")]
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial_test::serial(remote_exec_server)]
async fn shell_snapshot_v2_capture_failure_falls_back_and_retries(
use_remote: bool,
tty: bool,
shell_name: &str,
failures_before_repair: usize,
failure: CaptureFailure,
) -> Result<()> {
if use_remote
&& let Some(warning) =
Expand All @@ -489,9 +502,21 @@ async fn shell_snapshot_v2_capture_failure_falls_back_and_retries(
"zsh" => ("/bin/zsh", ".zshrc"),
name => anyhow::bail!("unsupported test shell {name}"),
};
let child_pids = home.path().join("startup-pids");
let _cleanup = shell_snapshot::StartupProcesses(child_pids.clone());
// Sandboxed remote PIDs may be namespace-local; check descendants on the native path.
let startup_child = if use_remote {
""
} else {
"/bin/sleep 30 >/dev/null 2>&1 &\nprintf '%s\\n' \"$!\" >> \"$HOME/startup-pids\"\n"
};
let finish = match failure {
CaptureFailure::Exit(code) => format!("exit {code}"),
CaptureFailure::Cancel | CaptureFailure::Timeout => "wait".to_string(),
};
std::fs::write(
home.path().join(profile_name),
"printf x >> \"$HOME/captures\"\nexit 7\n",
format!("printf x >> \"$HOME/captures\"\n{startup_child}{finish}\n"),
)?;
let policy = ExecEnvPolicy {
inherit: ShellEnvironmentPolicyInherit::All,
Expand Down Expand Up @@ -537,12 +562,47 @@ async fn shell_snapshot_v2_capture_failure_falls_back_and_retries(

for attempt in 0..failures_before_repair {
params.process_id = ProcessId::from(format!("snapshot-fallback-{attempt}"));
let fallback = context.backend.start(params.clone()).await?;
let fallback_output = collect_process_output_from_events(fallback.process).await?;
assert_eq!(
fallback_output,
("original".to_string(), String::new(), Some(0), true)
);
let mut start = context.backend.start(params.clone());
if matches!(failure, CaptureFailure::Cancel) {
tokio::select! {
result = &mut start => anyhow::bail!("capture finished before cancellation: {}", result.is_ok()),
_ = async {
while std::fs::read_to_string(&child_pids).unwrap_or_default().trim().is_empty() {
sleep(Duration::from_millis(20)).await;
}
} => {}
}
drop(start);
} else {
let fallback = start.await?;
assert_eq!(
collect_process_output_from_events(fallback.process).await?,
("original".to_string(), String::new(), Some(0), true)
);
}
if !use_remote {
let pids = std::fs::read_to_string(&child_pids)?;
let pid = pids
.split_whitespace()
.last()
.context("startup child PID")?;
timeout(Duration::from_secs(5), async {
loop {
let output = std::process::Command::new("/bin/ps")
.args(["-o", "stat=", "-p", pid])
.output()?;
let state = String::from_utf8(output.stdout)?;
if state.trim().is_empty() || state.trim().starts_with('Z') {
// Forget the completed child before retries can reuse its PID.
std::fs::remove_file(&child_pids)?;
return Ok::<_, anyhow::Error>(());
}
sleep(Duration::from_millis(20)).await;
}
})
.await
.context("failed capture left its startup child running")??;
}
// A real remote executor has its own clock; the unit test uses a
// paused clock to check requests made during the one-second backoff.
sleep(Duration::from_millis(1100)).await;
Expand Down
78 changes: 78 additions & 0 deletions codex-rs/exec-server/tests/exec_process/shell_snapshot.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,84 @@ enum SnapshotSandbox {
DenyFdPath,
}

pub(super) struct StartupProcesses(pub(super) std::path::PathBuf);

impl Drop for StartupProcesses {
fn drop(&mut self) {
if let Ok(contents) = std::fs::read_to_string(&self.0) {
for pid in contents
.split_whitespace()
.filter_map(|pid| pid.parse::<i32>().ok())
{
if pid > 0 {
// SAFETY: these PIDs were written by our startup profile.
unsafe { libc::kill(pid, libc::SIGKILL) };
}
}
}
}
}

#[test_case("bash"; "bash")]
#[cfg_attr(target_os = "macos", test_case("zsh"; "zsh"))]
#[tokio::test]
async fn shell_snapshot_preserves_successful_startup_output_and_services(
shell: &str,
) -> Result<()> {
let context = create_process_context(/*use_remote*/ false).await?;
let home = TempDir::new()?;
let _service = StartupProcesses(home.path().join("service-pid"));
std::fs::write(
home.path().join(format!(".{shell}rc")),
"printf x >> \"$HOME/captures\"\n/bin/sleep 30 >/dev/null 2>&1 &\nexport SNAPSHOT_SERVICE_PID=$!\nprintf '%s' \"$SNAPSHOT_SERVICE_PID\" > \"$HOME/service-pid\"\nexec > >(/bin/sleep 0.2; /bin/cat)\nprofile_helper() { printf 'captured:%s' \"$1\"; }\n",
)?;
for index in 0..2 {
let started = context
.backend
.start(ExecParams {
metadata: Default::default(),
process_id: format!("startup-capture-{index}").into(),
argv: vec![
format!("/bin/{shell}"),
"-lc".to_string(),
format!(
"[ \"$SNAPSHOT_SERVICE_PID\" = \"$(/bin/cat \"$HOME/service-pid\")\" ] || exit 41\ncase \"$(/bin/ps -o stat= -p \"$SNAPSHOT_SERVICE_PID\")\" in ''|*Z*) exit 42;; esac\nprofile_helper {index}"
),
],
cwd: PathUri::from_host_native_path(home.path())?,
env: HashMap::from([
(
"HOME".to_string(),
home.path().to_string_lossy().into_owned(),
),
("PATH".to_string(), "/usr/bin:/bin".to_string()),
]),
env_policy: None,
shell_snapshot: Some(ShellSnapshotRequest {
scope_id: "startup-capture".to_string(),
shell: ShellInfo {
name: shell.to_string(),
path: format!("/bin/{shell}"),
},
}),
tty: false,
pipe_stdin: false,
arg0: Some("codex-linux-sandbox".to_string()),
sandbox: None,
enforce_managed_network: false,
managed_network: None,
network_proxy: None,
})
.await?;
assert_eq!(
collect_process_output_from_events(started.process).await?,
(format!("captured:{index}"), String::new(), Some(0), true)
);
}
assert_eq!(std::fs::read_to_string(home.path().join("captures"))?, "x");
Ok(())
}

#[test_case("bash", false, SnapshotSandbox::None; "bash")]
#[test_case("bash", true, SnapshotSandbox::None; "bash_tty")]
#[cfg_attr(target_os = "macos", test_case("zsh", false, SnapshotSandbox::None; "zsh"))]
Expand Down
Loading