From 65bd8965c90e1ff701c7897d84024adfbbd151f4 Mon Sep 17 00:00:00 2001 From: jimmyfraiture Date: Thu, 11 Sep 2025 16:29:47 -0700 Subject: [PATCH] Fix the race condition --- codex-rs/core/src/codex.rs | 62 +++++++++++++++++------- codex-rs/core/src/shell.rs | 98 +++++++++++++++++++++++++------------- 2 files changed, 112 insertions(+), 48 deletions(-) diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index 238a72e32f..8a23029d12 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -2306,28 +2306,51 @@ fn should_translate_shell_command( || shell_policy.use_profile || matches!( shell, - crate::shell::Shell::Posix(shell) if shell.shell_snapshot.borrow().is_some() + crate::shell::Shell::Posix(shell) + if !shell.shell_snapshot.borrow().is_unavailable() ) } -fn maybe_translate_shell_command( - params: ExecParams, +async fn maybe_translate_shell_command( + mut params: ExecParams, sess: &Session, turn_context: &TurnContext, ) -> ExecParams { let should_translate = should_translate_shell_command(&sess.user_shell, &turn_context.shell_environment_policy); - if should_translate - && let Some(command) = sess - .user_shell - .format_default_shell_invocation(params.command.clone()) - { - return ExecParams { command, ..params }; + if !should_translate { + return params; } + + if let crate::shell::Shell::Posix(shell) = &sess.user_shell + && shell.shell_snapshot.borrow().is_pending() + { + wait_for_shell_snapshot(shell).await; + } + + let original_command = std::mem::take(&mut params.command); + params.command = sess + .user_shell + .format_default_shell_invocation(&original_command) + .unwrap_or(original_command); + params } +async fn wait_for_shell_snapshot(shell: &crate::shell::PosixShell) { + if !shell.shell_snapshot.borrow().is_pending() { + return; + } + + let mut rx = shell.shell_snapshot.clone(); + while rx.changed().await.is_ok() { + if !rx.borrow().is_pending() { + break; + } + } +} + async fn handle_container_exec_with_params( params: ExecParams, sess: &Session, @@ -2488,7 +2511,7 @@ async fn handle_container_exec_with_params( ), }; - let params = maybe_translate_shell_command(params, sess, turn_context); + let params = maybe_translate_shell_command(params, sess, turn_context).await; let output_result = sess .run_exec_with_events( turn_diff_tracker, @@ -2963,7 +2986,7 @@ mod tests { } } - fn zsh_shell(shell_snapshot: Option>) -> shell::Shell { + fn zsh_shell(shell_snapshot: shell::ShellSnapshotState) -> shell::Shell { let (_tx, rx) = tokio::sync::watch::channel(shell_snapshot); shell::Shell::Posix(shell::PosixShell { shell_path: "/bin/zsh".to_string(), @@ -2975,26 +2998,33 @@ mod tests { #[test] fn translates_commands_when_shell_policy_requests_profile() { let policy = shell_policy_with_profile(true); - let shell = zsh_shell(None); + let shell = zsh_shell(shell::ShellSnapshotState::Unavailable); assert!(should_translate_shell_command(&shell, &policy)); } #[test] fn translates_commands_for_zsh_with_snapshot() { let policy = shell_policy_with_profile(false); - let shell = zsh_shell(Some(Arc::new(ShellSnapshot::new(PathBuf::from( - "/tmp/snapshot", - ))))); + let shell = zsh_shell(shell::ShellSnapshotState::Ready(Arc::new( + ShellSnapshot::new(PathBuf::from("/tmp/snapshot")), + ))); assert!(should_translate_shell_command(&shell, &policy)); } #[test] fn bypasses_translation_for_zsh_without_snapshot_or_profile() { let policy = shell_policy_with_profile(false); - let shell = zsh_shell(None); + let shell = zsh_shell(shell::ShellSnapshotState::Unavailable); assert!(!should_translate_shell_command(&shell, &policy)); } + #[test] + fn translates_commands_for_zsh_with_pending_snapshot() { + let policy = shell_policy_with_profile(false); + let shell = zsh_shell(shell::ShellSnapshotState::Pending); + assert!(should_translate_shell_command(&shell, &policy)); + } + #[test] fn prefers_structured_content_when_present() { let ctr = CallToolResult { diff --git a/codex-rs/core/src/shell.rs b/codex-rs/core/src/shell.rs index d2acf75c20..2e997ef69d 100644 --- a/codex-rs/core/src/shell.rs +++ b/codex-rs/core/src/shell.rs @@ -13,11 +13,42 @@ pub struct ShellSnapshot { pub(crate) path: PathBuf, } +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ShellSnapshotState { + Pending, + Ready(Arc), + Unavailable, +} + +impl ShellSnapshotState { + pub fn is_pending(&self) -> bool { + matches!(self, Self::Pending) + } + + pub fn is_unavailable(&self) -> bool { + matches!(self, Self::Unavailable) + } + + pub fn snapshot(&self) -> Option<&Arc> { + match self { + Self::Ready(snapshot) => Some(snapshot), + Self::Pending | Self::Unavailable => None, + } + } + + pub fn into_snapshot(self) -> Option> { + match self { + Self::Ready(snapshot) => Some(snapshot), + Self::Pending | Self::Unavailable => None, + } + } +} + // Usage of fully qualified names for the Receiver for clarity. -type ShellSnapshotRx = tokio::sync::watch::Receiver>>; +type ShellSnapshotRx = tokio::sync::watch::Receiver; pub fn default_shell_snapshot_rx() -> ShellSnapshotRx { - let (_tx, rx) = tokio::sync::watch::channel(None); + let (_tx, rx) = tokio::sync::watch::channel(ShellSnapshotState::Unavailable); rx } @@ -68,22 +99,19 @@ pub enum Shell { } impl Shell { - pub fn format_default_shell_invocation(&self, command: Vec) -> Option> { + pub fn format_default_shell_invocation(&self, command: &[String]) -> Option> { match self { Shell::Posix(shell) => { - let joined = strip_bash_lc(&command) + let joined = strip_bash_lc(command) .or_else(|| shlex::try_join(command.iter().map(|s| s.as_str())).ok())?; - let mut source_path = Path::new(&shell.rc_path); - - let shell_snapshot = shell.shell_snapshot.borrow().clone(); - let session_cmd = if let Some(shell_snapshot) = shell_snapshot.as_ref() - && shell_snapshot.path.exists() + let snapshot_state = shell.shell_snapshot.borrow(); + let (source_path, session_cmd) = if let Some(snapshot) = snapshot_state.snapshot() + && snapshot.path.exists() { - source_path = shell_snapshot.path.as_path(); - "-c".to_string() + (snapshot.path.clone(), "-c".to_string()) } else { - "-lc".to_string() + (PathBuf::from(&shell.rc_path), "-lc".to_string()) }; let source_path_str = source_path.to_string_lossy().to_string(); @@ -95,7 +123,7 @@ impl Shell { } Shell::PowerShell(ps) => { // If model generated a bash command, prefer a detected bash fallback - if let Some(script) = strip_bash_lc(&command) { + if let Some(script) = strip_bash_lc(command) { return match &ps.bash_exe_fallback { Some(bash) => Some(vec![ bash.to_string_lossy().to_string(), @@ -121,7 +149,7 @@ impl Shell { if first != Some(ps.exe.as_str()) { // TODO (CODEX_2900): Handle escaping newlines. if command.iter().any(|a| a.contains('\n') || a.contains('\r')) { - return Some(command); + return Some(command.to_vec()); } let joined = shlex::try_join(command.iter().map(|s| s.as_str())).ok(); @@ -136,7 +164,7 @@ impl Shell { } // Model generated a PowerShell command. Run it. - Some(command) + Some(command.to_vec()) } Shell::Unknown => None, } @@ -154,14 +182,14 @@ impl Shell { pub fn get_snapshot(&self) -> Option> { match self { - Shell::Posix(shell) => shell.shell_snapshot.borrow().clone(), + Shell::Posix(shell) => shell.shell_snapshot.borrow().snapshot().cloned(), _ => None, } } } -fn strip_bash_lc(command: &Vec) -> Option { - match command.as_slice() { +fn strip_bash_lc(command: &[String]) -> Option { + match command { // exactly three items [first, second, third] // first two must be "bash", "-lc" @@ -197,7 +225,7 @@ async fn detect_default_user_shell(session_id: Uuid, codex_home: &Path) -> Shell return Shell::Unknown; }; - let (tx, rx) = tokio::sync::watch::channel(None); + let (tx, rx) = tokio::sync::watch::channel(ShellSnapshotState::Pending); { let shell_path = shell_path.clone(); @@ -216,9 +244,12 @@ async fn detect_default_user_shell(session_id: Uuid, codex_home: &Path) -> Shell if snapshot_path.is_none() { trace!("failed to prepare posix snapshot; using live profile"); } - let shell_snapshot = - snapshot_path.map(|snapshot| Arc::new(ShellSnapshot::new(snapshot))); - if tx.send(shell_snapshot).is_err() { + let snapshot_state = snapshot_path + .map(|snapshot| { + ShellSnapshotState::Ready(Arc::new(ShellSnapshot::new(snapshot))) + }) + .unwrap_or(ShellSnapshotState::Unavailable); + if tx.send(snapshot_state).is_err() { trace!("failed to send posix snapshot; using live profile"); } }); @@ -433,7 +464,8 @@ pub(crate) mod tests { rc_path: "/does/not/exist/.zshrc".to_string(), shell_snapshot: default_shell_snapshot_rx(), }); - let actual_cmd = shell.format_default_shell_invocation(vec!["myecho".to_string()]); + let command = vec!["myecho".to_string()]; + let actual_cmd = shell.format_default_shell_invocation(&command); assert_eq!( actual_cmd, Some(vec![ @@ -495,8 +527,8 @@ pub(crate) mod tests { shell_snapshot: default_shell_snapshot_rx(), }); - let actual_cmd = shell - .format_default_shell_invocation(input.iter().map(|s| s.to_string()).collect()); + let input = input.iter().map(|s| s.to_string()).collect::>(); + let actual_cmd = shell.format_default_shell_invocation(&input); let expected_cmd = expected_cmd .iter() .map(|s| { @@ -600,8 +632,9 @@ mod macos_tests { let snapshot_path = temp_dir.path().join("snapshot.zsh"); std::fs::write(&snapshot_path, "export SNAPSHOT_READY=1").unwrap(); - let (_tx, rx) = - tokio::sync::watch::channel(Some(Arc::new(ShellSnapshot::new(snapshot_path.clone())))); + let (_tx, rx) = tokio::sync::watch::channel(ShellSnapshotState::Ready(Arc::new( + ShellSnapshot::new(snapshot_path.clone()), + ))); let shell = Shell::Posix(PosixShell { shell_path: "/bin/zsh".to_string(), @@ -613,7 +646,8 @@ mod macos_tests { shell_snapshot: rx, }); - let invocation = shell.format_default_shell_invocation(vec!["echo".to_string()]); + let command = vec!["echo".to_string()]; + let invocation = shell.format_default_shell_invocation(&command); let expected_command = vec!["/bin/zsh".to_string(), "-c".to_string(), { let snapshot_path = snapshot_path.to_string_lossy(); format!("[ -f {snapshot_path} ] && . {snapshot_path}; (echo)") @@ -692,8 +726,8 @@ mod macos_tests { shell_snapshot: default_shell_snapshot_rx(), }); - let actual_cmd = shell - .format_default_shell_invocation(input.iter().map(|s| s.to_string()).collect()); + let input = input.iter().map(|s| s.to_string()).collect::>(); + let actual_cmd = shell.format_default_shell_invocation(&input); let expected_cmd = expected_cmd .iter() .map(|s| { @@ -819,8 +853,8 @@ mod tests_windows { ]; for (shell, input, expected_cmd) in cases { - let actual_cmd = shell - .format_default_shell_invocation(input.iter().map(|s| s.to_string()).collect()); + let input = input.iter().map(|s| s.to_string()).collect::>(); + let actual_cmd = shell.format_default_shell_invocation(&input); assert_eq!( actual_cmd, Some(expected_cmd.iter().map(|s| s.to_string()).collect())