From e4f8441cab402d71fbe6e5ef800ef7f05b2f4012 Mon Sep 17 00:00:00 2001 From: Sayan Sisodiya Date: Wed, 17 Jun 2026 00:27:23 -0700 Subject: [PATCH] core: track starting environments in snapshots --- codex-rs/core/src/agents_md_tests.rs | 6 +- codex-rs/core/src/environment_selection.rs | 482 +++++++++++++++------ 2 files changed, 358 insertions(+), 130 deletions(-) diff --git a/codex-rs/core/src/agents_md_tests.rs b/codex-rs/core/src/agents_md_tests.rs index 42604629a0..ce0984393b 100644 --- a/codex-rs/core/src/agents_md_tests.rs +++ b/codex-rs/core/src/agents_md_tests.rs @@ -257,8 +257,8 @@ async fn agents_md_paths(config: &TestConfig) -> std::io::Result( environments: [(&str, AbsolutePathBuf); N], ) -> TurnEnvironmentSnapshot { - TurnEnvironmentSnapshot { - turn_environments: environments + TurnEnvironmentSnapshot::from_turn_environments( + environments .into_iter() .map(|(environment_id, cwd)| { TurnEnvironment::new( @@ -272,7 +272,7 @@ fn resolved_local_environments( ) }) .collect(), - } + ) } fn project_provenance(path: AbsolutePathBuf, cwd: AbsolutePathBuf) -> InstructionProvenance { diff --git a/codex-rs/core/src/environment_selection.rs b/codex-rs/core/src/environment_selection.rs index 16b2c47b5d..b92161c732 100644 --- a/codex-rs/core/src/environment_selection.rs +++ b/codex-rs/core/src/environment_selection.rs @@ -2,16 +2,13 @@ use std::collections::HashSet; use std::sync::Arc; use arc_swap::ArcSwap; +use codex_exec_server::Environment; use codex_exec_server::EnvironmentManager; use codex_exec_server::ExecutorFileSystem; -use codex_protocol::error::CodexErr; -use codex_protocol::error::Result as CodexResult; use codex_protocol::protocol::TurnEnvironmentSelection; use codex_utils_absolute_path::AbsolutePathBuf; use codex_utils_path_uri::PathUri; use futures::FutureExt; -use futures::future::BoxFuture; -use futures::future::Shared; use crate::session::turn_context::TurnEnvironment; use crate::shell::Shell; @@ -31,13 +28,17 @@ pub(crate) fn default_thread_environment_selections( .collect() } -type SnapshotTask = Shared>; +#[derive(Clone, Debug)] +pub(crate) struct StartingTurnEnvironment { + pub(crate) selection: TurnEnvironmentSelection, + pub(crate) environment: Arc, +} pub(crate) struct ThreadEnvironments { environment_manager: Arc, local_shell: Shell, shell_snapshot: ShellSnapshot, - snapshot_task: ArcSwap, + snapshot: ArcSwap, } impl ThreadEnvironments { @@ -51,127 +52,180 @@ impl ThreadEnvironments { environment_manager, local_shell, shell_snapshot, - snapshot_task: ArcSwap::from_pointee(futures::future::ready(current).boxed().shared()), + snapshot: ArcSwap::from_pointee(current), } } pub(crate) fn update_selections(&self, environments: &[TurnEnvironmentSelection]) { - let previous = self - .snapshot_task - .load() - .peek() - .cloned() - .unwrap_or_default(); - let environment_manager = Arc::clone(&self.environment_manager); - let local_shell = self.local_shell.clone(); - let shell_snapshot = self.shell_snapshot.clone(); - let environments = environments.to_vec(); - let (snapshot_task, snapshot) = async move { - Self::resolve_snapshot( - environment_manager, - local_shell, - shell_snapshot, - previous, - environments, - ) - .await - } - .remote_handle(); - self.snapshot_task - .store(Arc::new(snapshot.boxed().shared())); - drop(tokio::spawn(snapshot_task)); - } - - async fn resolve_snapshot( - environment_manager: Arc, - local_shell: Shell, - shell_snapshot: ShellSnapshot, - current: TurnEnvironmentSnapshot, - environments: Vec, - ) -> TurnEnvironmentSnapshot { + let previous = self.snapshot.load(); let mut seen_environment_ids = HashSet::with_capacity(environments.len()); let mut turn_environments = Vec::with_capacity(environments.len()); - for selected_environment in &environments { + let mut starting = Vec::with_capacity(environments.len()); + let mut ordered_selections = Vec::with_capacity(environments.len()); + for selected_environment in environments { if !seen_environment_ids.insert(selected_environment.environment_id.as_str()) { continue; } - let turn_environment = match current.turn_environments.iter().find(|environment| { + // Reuse the exact attached or starting environment already selected by this thread. + if let Some(environment) = previous.turn_environments.iter().find(|environment| { environment.environment_id == selected_environment.environment_id && environment.cwd() == &selected_environment.cwd }) { - Some(environment) => environment.clone(), - None => match Self::resolve_selection( - &environment_manager, - &local_shell, - &shell_snapshot, - selected_environment, - ) - .await - { - Ok(environment) => environment, - Err(err) => { - tracing::warn!( - "skipping unresolved turn environment `{}`: {err}", - selected_environment.environment_id - ); - continue; - } - }, + turn_environments.push(environment.clone()); + ordered_selections.push(selected_environment.clone()); + continue; + } + if let Some(environment) = previous + .starting + .iter() + .find(|environment| environment.selection == *selected_environment) + { + starting.push(environment.clone()); + ordered_selections.push(selected_environment.clone()); + continue; + } + + // Only new selections consult the manager; reused selections keep their stable handle. + let environment_id = &selected_environment.environment_id; + let Some(environment) = self.environment_manager.get_environment(environment_id) else { + tracing::warn!("skipping unknown turn environment `{environment_id}`"); + continue; }; - turn_environments.push(turn_environment); + if environment.is_remote() { + // Connect in the background and leave attachment to a later snapshot. + environment.start_connecting(); + starting.push(StartingTurnEnvironment { + selection: selected_environment.clone(), + environment, + }); + } else { + turn_environments.push(self.build_turn_environment( + selected_environment, + environment, + Some(self.local_shell.clone()), + )); + } + ordered_selections.push(selected_environment.clone()); } - TurnEnvironmentSnapshot { turn_environments } + self.snapshot.store(Arc::new(TurnEnvironmentSnapshot { + turn_environments, + starting, + ordered_selections, + })); } - async fn resolve_selection( - environment_manager: &EnvironmentManager, - local_shell: &Shell, - shell_snapshot: &ShellSnapshot, - selected_environment: &TurnEnvironmentSelection, - ) -> CodexResult { - let environment_id = selected_environment.environment_id.clone(); - let environment = environment_manager - .get_environment(&environment_id) - .ok_or_else(|| { - CodexErr::InvalidRequest(format!("unknown turn environment id `{environment_id}`")) - })?; - let shell = if environment.is_remote() { - match environment.info().await { - Ok(info) => match Shell::from_environment_shell_info(info.shell) { - Ok(shell) => Some(shell), - Err(err) => { - tracing::warn!( - "failed to resolve shell for environment `{environment_id}`: {err}" - ); - None - } - }, + async fn resolve_starting_environment( + &self, + starting: &StartingTurnEnvironment, + ) -> TurnEnvironment { + let environment_id = &starting.selection.environment_id; + let shell = match starting.environment.info().boxed().await { + Ok(info) => match Shell::from_environment_shell_info(info.shell) { + Ok(shell) => Some(shell), Err(err) => { - tracing::warn!("failed to get info for environment `{environment_id}`: {err}"); + tracing::warn!( + "failed to resolve shell for environment `{environment_id}`: {err}" + ); None } + }, + Err(err) => { + tracing::warn!("failed to get info for environment `{environment_id}`: {err}"); + None } - } else { - Some(local_shell.clone()) }; + self.build_turn_environment( + &starting.selection, + Arc::clone(&starting.environment), + shell, + ) + } + + fn build_turn_environment( + &self, + selected_environment: &TurnEnvironmentSelection, + environment: Arc, + shell: Option, + ) -> TurnEnvironment { let mut turn_environment = TurnEnvironment::new( - environment_id, + selected_environment.environment_id.clone(), environment, selected_environment.cwd.clone(), shell, ); - let task = shell_snapshot + let task = self + .shell_snapshot .clone() .build(turn_environment.clone()) .boxed() .shared(); drop(tokio::spawn(task.clone())); turn_environment.shell_snapshot = task; - Ok(turn_environment) + turn_environment } pub(crate) async fn snapshot(&self) -> TurnEnvironmentSnapshot { - self.snapshot_task.load_full().as_ref().clone().await + loop { + let current = self.snapshot.load_full(); + if current.starting.is_empty() { + return current.as_ref().clone(); + } + + // Rebuild both lists in configured order while promoting completed startups. + let mut changed = false; + let mut turn_environments = Vec::with_capacity(current.ordered_selections.len()); + let mut starting = Vec::with_capacity(current.starting.len()); + for selection in ¤t.ordered_selections { + if let Some(environment) = current.turn_environments.iter().find(|environment| { + environment.environment_id == selection.environment_id + && environment.cwd() == &selection.cwd + }) { + turn_environments.push(environment.clone()); + continue; + } + let Some(environment) = current + .starting + .iter() + .find(|environment| environment.selection == *selection) + else { + continue; + }; + if !environment.environment.startup_finished() { + // Never wait for an environment whose startup is still running. + starting.push(environment.clone()); + continue; + } + + changed = true; + // Startup finished, so this only reads its saved success or failure. + match environment.environment.wait_until_ready().boxed().await { + Ok(()) => { + turn_environments + .push(self.resolve_starting_environment(environment).await); + } + Err(err) => { + tracing::warn!( + "turn environment `{}` failed to start: {err}", + environment.selection.environment_id + ); + } + } + } + if !changed { + return current.as_ref().clone(); + } + + let next = Arc::new(TurnEnvironmentSnapshot { + turn_environments, + starting, + ordered_selections: current.ordered_selections.clone(), + }); + // Do not overwrite selections changed while shell resolution was in flight. + let previous = self.snapshot.compare_and_swap(¤t, Arc::clone(&next)); + if Arc::ptr_eq(&previous, ¤t) { + return next.as_ref().clone(); + } + } } pub(crate) fn environment_manager(&self) -> Arc { @@ -182,9 +236,25 @@ impl ThreadEnvironments { #[derive(Clone, Debug, Default)] pub(crate) struct TurnEnvironmentSnapshot { pub(crate) turn_environments: Vec, + pub(crate) starting: Vec, + // Attached and starting environments are stored separately, so retain their configured order. + ordered_selections: Vec, } impl TurnEnvironmentSnapshot { + #[cfg(test)] + pub(crate) fn from_turn_environments(turn_environments: Vec) -> Self { + let ordered_selections = turn_environments + .iter() + .map(TurnEnvironment::selection) + .collect(); + Self { + turn_environments, + starting: Vec::new(), + ordered_selections, + } + } + pub(crate) fn primary(&self) -> Option<&TurnEnvironment> { self.turn_environments.first() } @@ -202,10 +272,7 @@ impl TurnEnvironmentSnapshot { } pub(crate) fn to_selections(&self) -> Vec { - self.turn_environments - .iter() - .map(TurnEnvironment::selection) - .collect() + self.ordered_selections.clone() } pub(crate) fn primary_filesystem(&self) -> Option> { @@ -237,7 +304,16 @@ mod tests { use codex_protocol::protocol::TurnEnvironmentSelection; use codex_utils_absolute_path::AbsolutePathBuf; use codex_utils_path_uri::PathUri; + use futures::SinkExt; + use futures::StreamExt; use pretty_assertions::assert_eq; + use serde_json::Value; + use tokio::net::TcpListener; + use tokio::net::TcpStream; + use tokio::time::timeout; + use tokio_tungstenite::WebSocketStream; + use tokio_tungstenite::accept_async; + use tokio_tungstenite::tungstenite::Message; use super::*; @@ -264,6 +340,61 @@ mod tests { .expect("runtime paths") } + async fn read_websocket_json(websocket: &mut WebSocketStream) -> Value { + loop { + match timeout(std::time::Duration::from_secs(5), websocket.next()) + .await + .expect("websocket read should not time out") + .expect("websocket should stay open") + .expect("websocket frame should read") + { + Message::Text(text) => { + return serde_json::from_str(text.as_ref()).expect("valid JSON-RPC message"); + } + Message::Binary(bytes) => { + return serde_json::from_slice(bytes.as_ref()).expect("valid JSON-RPC message"); + } + Message::Ping(_) | Message::Pong(_) => {} + other => panic!("expected JSON-RPC message, got {other:?}"), + } + } + } + + async fn serve_environment_info(listener: TcpListener) { + let (stream, _) = listener.accept().await.expect("connection"); + let mut websocket = accept_async(stream).await.expect("websocket handshake"); + + let initialize = read_websocket_json(&mut websocket).await; + assert_eq!(initialize["method"], "initialize"); + websocket + .send(Message::Text( + serde_json::json!({ + "id": initialize["id"], + "result": { "sessionId": "test-session" } + }) + .to_string() + .into(), + )) + .await + .expect("initialize response"); + let initialized = read_websocket_json(&mut websocket).await; + assert_eq!(initialized["method"], "initialized"); + + let info = read_websocket_json(&mut websocket).await; + assert_eq!(info["method"], "environment/info"); + websocket + .send(Message::Text( + serde_json::json!({ + "id": info["id"], + "result": { "shell": { "name": "zsh", "path": "/bin/zsh" } } + }) + .to_string() + .into(), + )) + .await + .expect("environment info response"); + } + #[tokio::test] async fn default_thread_environment_selections_use_manager_default_id() { let cwd = AbsolutePathBuf::current_dir().expect("cwd"); @@ -447,6 +578,84 @@ url = "ws://127.0.0.1:8765" assert_eq!(resolved.snapshot().await.to_selections(), vec![local]); } + #[tokio::test] + async fn snapshot_keeps_starting_environment_until_it_can_be_attached() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind websocket listener"); + let manager = Arc::new( + EnvironmentManager::create_for_tests_with_local( + Some(format!( + "ws://{}", + listener.local_addr().expect("listener address") + )), + test_runtime_paths(), + ) + .await, + ); + let cwd = AbsolutePathBuf::current_dir().expect("cwd"); + let cwd = PathUri::from_abs_path(&cwd); + let remote = TurnEnvironmentSelection { + environment_id: REMOTE_ENVIRONMENT_ID.to_string(), + cwd: cwd.clone(), + }; + let local = TurnEnvironmentSelection { + environment_id: LOCAL_ENVIRONMENT_ID.to_string(), + cwd, + }; + let turn_environments = ThreadEnvironments::new( + manager, + crate::shell::default_user_shell(), + ShellSnapshot::disabled(), + TurnEnvironmentSnapshot::default(), + ); + turn_environments.update_selections(&[remote.clone(), local.clone()]); + + let starting = turn_environments.snapshot().await; + assert_eq!( + starting + .turn_environments + .iter() + .map(TurnEnvironment::selection) + .collect::>(), + vec![local.clone()] + ); + assert_eq!( + starting + .starting + .iter() + .map(|environment| environment.selection.clone()) + .collect::>(), + vec![remote.clone()] + ); + assert_eq!( + starting.to_selections(), + vec![remote.clone(), local.clone()] + ); + + let server = tokio::spawn(serve_environment_info(listener)); + timeout( + std::time::Duration::from_secs(5), + starting.starting[0].environment.wait_until_ready(), + ) + .await + .expect("environment startup should finish") + .expect("environment startup should succeed"); + let attached = turn_environments.snapshot().await; + + assert!(attached.starting.is_empty()); + assert_eq!( + attached + .turn_environments + .iter() + .map(TurnEnvironment::selection) + .collect::>(), + vec![remote.clone(), local.clone()] + ); + assert_eq!(attached.to_selections(), vec![remote, local]); + server.await.expect("server task"); + } + #[tokio::test] async fn latest_environment_update_wins_while_previous_resolution_is_pending() { let listener = tokio::net::TcpListener::bind("127.0.0.1:0") @@ -495,11 +704,17 @@ url = "ws://127.0.0.1:8765" } #[tokio::test] - async fn matching_environment_id_and_cwd_reuse_resolved_environment() { + async fn matching_environment_id_and_cwd_reuse_starting_environment() { let cwd = AbsolutePathBuf::current_dir().expect("cwd"); + let first_listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind first listener"); let manager = Arc::new( EnvironmentManager::create_for_tests( - Some("ws://127.0.0.1:8765".to_string()), + Some(format!( + "ws://{}", + first_listener.local_addr().expect("first listener address") + )), Some(test_runtime_paths()), ) .await, @@ -508,41 +723,58 @@ url = "ws://127.0.0.1:8765" environment_id: REMOTE_ENVIRONMENT_ID.to_string(), cwd: PathUri::from_abs_path(&cwd), }; - let initial = - resolve_turn_environments(Arc::clone(&manager), std::slice::from_ref(&selection)).await; + let environments = ThreadEnvironments::new( + Arc::clone(&manager), + crate::shell::default_user_shell(), + ShellSnapshot::disabled(), + TurnEnvironmentSnapshot::default(), + ); + environments.update_selections(std::slice::from_ref(&selection)); + let initial_snapshot = environments.snapshot().await; + let second_listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind second listener"); manager .upsert_environment( REMOTE_ENVIRONMENT_ID.to_string(), - "ws://127.0.0.1:9876".to_string(), + format!( + "ws://{}", + second_listener + .local_addr() + .expect("second listener address") + ), ) .expect("replace environment"); - let initial_snapshot = initial.snapshot().await; - initial.update_selections(std::slice::from_ref(&selection)); - let reused_snapshot = initial.snapshot().await; - initial.update_selections(&[TurnEnvironmentSelection { + environments.update_selections(std::slice::from_ref(&selection)); + let reused_snapshot = environments.snapshot().await; + environments.update_selections(&[TurnEnvironmentSelection { cwd: PathUri::from_abs_path(&cwd.join("changed")), ..selection }]); - let changed_snapshot = initial.snapshot().await; + let changed_snapshot = environments.snapshot().await; assert!(Arc::ptr_eq( &initial_snapshot - .primary() + .starting + .first() .expect("initial environment") .environment, &reused_snapshot - .primary() + .starting + .first() .expect("reused environment") .environment, )); assert!(!Arc::ptr_eq( &reused_snapshot - .primary() + .starting + .first() .expect("reused environment") .environment, &changed_snapshot - .primary() + .starting + .first() .expect("changed environment") .environment, )); @@ -566,25 +798,21 @@ url = "ws://127.0.0.1:8765" Environment::create_for_tests(Some("ws://127.0.0.1:8765".to_string())) .expect("remote environment"), ); - let remote = TurnEnvironmentSnapshot { - turn_environments: vec![TurnEnvironment::new( + let remote = TurnEnvironmentSnapshot::from_turn_environments(vec![TurnEnvironment::new( + REMOTE_ENVIRONMENT_ID.to_string(), + remote_environment.clone(), + cwd_uri.clone(), + /*shell*/ None, + )]); + let multiple = TurnEnvironmentSnapshot::from_turn_environments(vec![ + local.primary().expect("local environment").clone(), + TurnEnvironment::new( REMOTE_ENVIRONMENT_ID.to_string(), - remote_environment.clone(), - cwd_uri.clone(), + remote_environment, + cwd_uri, /*shell*/ None, - )], - }; - let multiple = TurnEnvironmentSnapshot { - turn_environments: vec![ - local.primary().expect("local environment").clone(), - TurnEnvironment::new( - REMOTE_ENVIRONMENT_ID.to_string(), - remote_environment, - cwd_uri, - /*shell*/ None, - ), - ], - }; + ), + ]); assert_eq!(local.single_local_environment_cwd(), Some(cwd)); assert_eq!(remote.single_local_environment_cwd(), None);