core: track starting environments in snapshots

This commit is contained in:
Sayan Sisodiya
2026-06-17 00:27:23 -07:00
parent f4f57bb882
commit e4f8441cab
2 changed files with 358 additions and 130 deletions

View File

@@ -257,8 +257,8 @@ async fn agents_md_paths(config: &TestConfig) -> std::io::Result<Vec<AbsolutePat
fn resolved_local_environments<const N: usize>(
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<const N: usize>(
)
})
.collect(),
}
)
}
fn project_provenance(path: AbsolutePathBuf, cwd: AbsolutePathBuf) -> InstructionProvenance {

View File

@@ -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<BoxFuture<'static, TurnEnvironmentSnapshot>>;
#[derive(Clone, Debug)]
pub(crate) struct StartingTurnEnvironment {
pub(crate) selection: TurnEnvironmentSelection,
pub(crate) environment: Arc<Environment>,
}
pub(crate) struct ThreadEnvironments {
environment_manager: Arc<EnvironmentManager>,
local_shell: Shell,
shell_snapshot: ShellSnapshot,
snapshot_task: ArcSwap<SnapshotTask>,
snapshot: ArcSwap<TurnEnvironmentSnapshot>,
}
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<EnvironmentManager>,
local_shell: Shell,
shell_snapshot: ShellSnapshot,
current: TurnEnvironmentSnapshot,
environments: Vec<TurnEnvironmentSelection>,
) -> 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<TurnEnvironment> {
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<Environment>,
shell: Option<Shell>,
) -> 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 &current.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(&current, Arc::clone(&next));
if Arc::ptr_eq(&previous, &current) {
return next.as_ref().clone();
}
}
}
pub(crate) fn environment_manager(&self) -> Arc<EnvironmentManager> {
@@ -182,9 +236,25 @@ impl ThreadEnvironments {
#[derive(Clone, Debug, Default)]
pub(crate) struct TurnEnvironmentSnapshot {
pub(crate) turn_environments: Vec<TurnEnvironment>,
pub(crate) starting: Vec<StartingTurnEnvironment>,
// Attached and starting environments are stored separately, so retain their configured order.
ordered_selections: Vec<TurnEnvironmentSelection>,
}
impl TurnEnvironmentSnapshot {
#[cfg(test)]
pub(crate) fn from_turn_environments(turn_environments: Vec<TurnEnvironment>) -> 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<TurnEnvironmentSelection> {
self.turn_environments
.iter()
.map(TurnEnvironment::selection)
.collect()
self.ordered_selections.clone()
}
pub(crate) fn primary_filesystem(&self) -> Option<Arc<dyn ExecutorFileSystem>> {
@@ -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<TcpStream>) -> 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<_>>(),
vec![local.clone()]
);
assert_eq!(
starting
.starting
.iter()
.map(|environment| environment.selection.clone())
.collect::<Vec<_>>(),
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<_>>(),
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);