diff --git a/codex-rs/agent-roles/src/agent_role_config.rs b/codex-rs/agent-roles/src/agent_role_config.rs new file mode 100644 index 0000000000000000000000000000000000000000..0e53a71957c78a99e6dd2607c4bcd332da353e2e --- /dev/null +++ b/codex-rs/agent-roles/src/agent_role_config.rs @@ -0,0 +1,209 @@ +use codex_config::config_toml::ConfigToml; +use codex_utils_absolute_path::AbsolutePathBufGuard; +use serde::Deserialize; +use std::collections::BTreeSet; +use std::path::Path; +use std::path::PathBuf; +use toml::Value as TomlValue; + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct AgentRoleConfig { + /// Human-facing role documentation used in spawn tool guidance. + /// Required for loaded user-defined roles after deprecated/new metadata precedence resolves. + pub description: Option, + /// Path to a role-specific config layer. + pub config_file: Option, + /// Candidate nicknames for agents spawned with this role. + pub nickname_candidates: Option>, +} + +#[derive(Deserialize, Debug, Clone, Default, PartialEq)] +#[serde(deny_unknown_fields)] +struct RawAgentRoleFileToml { + name: Option, + description: Option, + nickname_candidates: Option>, + #[serde(flatten)] + config: ConfigToml, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct ResolvedAgentRoleFile { + pub role_name: String, + pub description: Option, + pub nickname_candidates: Option>, + pub config: TomlValue, +} + +pub fn parse_agent_role_file_contents( + contents: &str, + role_file_label: &Path, + config_base_dir: &Path, + role_name_hint: Option<&str>, +) -> std::io::Result { + let role_file_toml: TomlValue = toml::from_str(contents).map_err(|err| { + std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!( + "failed to parse agent role file at {}: {err}", + role_file_label.display() + ), + ) + })?; + let _guard = AbsolutePathBufGuard::new(config_base_dir); + let parsed: RawAgentRoleFileToml = role_file_toml.clone().try_into().map_err(|err| { + std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!( + "failed to deserialize agent role file at {}: {err}", + role_file_label.display() + ), + ) + })?; + let description = normalize_agent_role_description( + &format!("agent role file {}.description", role_file_label.display()), + parsed.description.as_deref(), + )?; + validate_agent_role_file_developer_instructions( + role_file_label, + parsed.config.developer_instructions.as_deref(), + role_name_hint.is_none(), + )?; + + let role_name = parsed + .name + .as_deref() + .map(str::trim) + .filter(|name| !name.is_empty()) + .map(ToOwned::to_owned) + .or_else(|| role_name_hint.map(ToOwned::to_owned)) + .ok_or_else(|| { + std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!( + "agent role file at {} must define a non-empty `name`", + role_file_label.display() + ), + ) + })?; + + let nickname_candidates = normalize_agent_role_nickname_candidates( + &format!( + "agent role file {}.nickname_candidates", + role_file_label.display() + ), + parsed.nickname_candidates.as_deref(), + )?; + + let mut config = role_file_toml; + let Some(config_table) = config.as_table_mut() else { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!( + "agent role file at {} must contain a TOML table", + role_file_label.display() + ), + )); + }; + config_table.remove("name"); + config_table.remove("description"); + config_table.remove("nickname_candidates"); + + Ok(ResolvedAgentRoleFile { + role_name, + description, + nickname_candidates, + config, + }) +} + +pub(crate) fn normalize_agent_role_description( + field_label: &str, + description: Option<&str>, +) -> std::io::Result> { + match description.map(str::trim) { + Some("") => Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!("{field_label} cannot be blank"), + )), + Some(description) => Ok(Some(description.to_string())), + None => Ok(None), + } +} + +fn validate_agent_role_file_developer_instructions( + role_file_label: &Path, + developer_instructions: Option<&str>, + require_present: bool, +) -> std::io::Result<()> { + match developer_instructions.map(str::trim) { + Some("") => Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!( + "agent role file at {}.developer_instructions cannot be blank", + role_file_label.display() + ), + )), + Some(_) => Ok(()), + None if require_present => Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!( + "agent role file at {} must define `developer_instructions`", + role_file_label.display() + ), + )), + None => Ok(()), + } +} + +pub(crate) fn normalize_agent_role_nickname_candidates( + field_label: &str, + nickname_candidates: Option<&[String]>, +) -> std::io::Result>> { + let Some(nickname_candidates) = nickname_candidates else { + return Ok(None); + }; + + if nickname_candidates.is_empty() { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!("{field_label} must contain at least one name"), + )); + } + + let mut normalized_candidates = Vec::with_capacity(nickname_candidates.len()); + let mut seen_candidates = BTreeSet::new(); + + for nickname in nickname_candidates { + let normalized_nickname = nickname.trim(); + if normalized_nickname.is_empty() { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!("{field_label} cannot contain blank names"), + )); + } + + if !seen_candidates.insert(normalized_nickname.to_owned()) { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!("{field_label} cannot contain duplicates"), + )); + } + + if !normalized_nickname + .chars() + .all(|c| c.is_ascii_alphanumeric() || matches!(c, ' ' | '-' | '_')) + { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!( + "{field_label} may only contain ASCII letters, digits, spaces, hyphens, and underscores" + ), + )); + } + + normalized_candidates.push(normalized_nickname.to_owned()); + } + + Ok(Some(normalized_candidates)) +} diff --git a/codex-rs/agent-roles/src/discovery.rs b/codex-rs/agent-roles/src/discovery.rs new file mode 100644 index 0000000000000000000000000000000000000000..79d4aae05a34ef76faea89ce759dfdbcb2192482 --- /dev/null +++ b/codex-rs/agent-roles/src/discovery.rs @@ -0,0 +1,40 @@ +use codex_file_system::ExecutorFileSystem; +use codex_utils_absolute_path::AbsolutePathBuf; +use codex_utils_path_uri::PathUri; +use std::io; +use std::io::ErrorKind; + +pub(crate) async fn collect_agent_role_files( + fs: &dyn ExecutorFileSystem, + dir: &AbsolutePathBuf, +) -> io::Result> { + let mut files = Vec::new(); + let mut dirs = vec![dir.clone()]; + while let Some(dir) = dirs.pop() { + let dir_uri = PathUri::from_abs_path(&dir); + let entries = match fs.read_directory(&dir_uri, /*sandbox*/ None).await { + Ok(entries) => entries, + Err(err) if err.kind() == ErrorKind::NotFound => continue, + Err(err) => return Err(err), + }; + + for entry in entries { + let path = dir.join(entry.file_name); + if entry.is_directory { + dirs.push(path); + continue; + } + if entry.is_file + && path + .as_path() + .extension() + .is_some_and(|extension| extension == "toml") + { + files.push(path); + } + } + } + + files.sort(); + Ok(files) +} diff --git a/codex-rs/agent-roles/src/lib.rs b/codex-rs/agent-roles/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..0ad2a7bd9af739e8acdc2703267b3257a9872a14 --- /dev/null +++ b/codex-rs/agent-roles/src/lib.rs @@ -0,0 +1,8 @@ +mod agent_role_config; +mod discovery; +mod loader; + +pub use agent_role_config::AgentRoleConfig; +pub use agent_role_config::ResolvedAgentRoleFile; +pub use agent_role_config::parse_agent_role_file_contents; +pub use loader::load_agent_roles; diff --git a/codex-rs/agent-roles/src/loader.rs b/codex-rs/agent-roles/src/loader.rs new file mode 100644 index 0000000000000000000000000000000000000000..a82649b8e10debe7a2b3c7e21e2e861d557ea3b7 --- /dev/null +++ b/codex-rs/agent-roles/src/loader.rs @@ -0,0 +1,335 @@ +use crate::AgentRoleConfig; +use crate::ResolvedAgentRoleFile; +use crate::agent_role_config::normalize_agent_role_description; +use crate::agent_role_config::normalize_agent_role_nickname_candidates; +use crate::discovery::collect_agent_role_files; +use crate::parse_agent_role_file_contents; +use codex_config::ConfigLayerStack; +use codex_config::config_toml::AgentRoleToml; +use codex_config::config_toml::AgentsToml; +use codex_config::config_toml::ConfigToml; +use codex_file_system::ExecutorFileSystem; +use codex_file_system::GetMetadataOptions; +use codex_file_system::ReadFileOptions; +use codex_utils_absolute_path::AbsolutePathBuf; +use codex_utils_absolute_path::AbsolutePathBufGuard; +use codex_utils_path_uri::PathUri; +use std::collections::BTreeMap; +use std::collections::BTreeSet; +use std::path::Path; +use std::path::PathBuf; +use toml::Value as TomlValue; + +pub async fn load_agent_roles( + fs: &dyn ExecutorFileSystem, + cfg: &ConfigToml, + config_layer_stack: &ConfigLayerStack, + startup_warnings: &mut Vec, +) -> std::io::Result> { + let mut layers = config_layer_stack.layers_low_to_high().peekable(); + if layers.peek().is_none() { + return load_agent_roles_without_layers(fs, cfg).await; + } + + let mut roles: BTreeMap = BTreeMap::new(); + for layer in layers { + let mut layer_roles: BTreeMap = BTreeMap::new(); + let mut declared_role_files = BTreeSet::new(); + let config_folder = layer.config_folder(); + let agents_toml = match agents_toml_from_layer(&layer.config, config_folder.as_deref()) { + Ok(agents_toml) => agents_toml, + Err(err) => { + push_agent_role_warning(startup_warnings, err); + None + } + }; + if let Some(agents_toml) = agents_toml { + for (declared_role_name, role_toml) in &agents_toml.roles { + let (role_name, role) = + match read_declared_role(fs, declared_role_name, role_toml).await { + Ok(role) => role, + Err(err) => { + push_agent_role_warning(startup_warnings, err); + continue; + } + }; + if let Some(config_file) = role.config_file.clone() { + declared_role_files.insert(config_file); + } + if layer_roles.contains_key(&role_name) { + push_agent_role_warning( + startup_warnings, + std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!( + "duplicate agent role name `{role_name}` declared in the same config layer" + ), + ), + ); + continue; + } + layer_roles.insert(role_name, role); + } + } + + if let Some(config_folder) = layer.config_folder() { + for (role_name, role) in discover_agent_roles_in_dir( + fs, + &config_folder.join("agents"), + &declared_role_files, + startup_warnings, + ) + .await? + { + if layer_roles.contains_key(&role_name) { + push_agent_role_warning( + startup_warnings, + std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!( + "duplicate agent role name `{role_name}` declared in the same config layer" + ), + ), + ); + continue; + } + layer_roles.insert(role_name, role); + } + } + + for (role_name, role) in layer_roles { + let mut merged_role = role; + if let Some(existing_role) = roles.get(&role_name) { + merge_missing_role_fields(&mut merged_role, existing_role); + } + if let Err(err) = validate_required_agent_role_description( + &role_name, + merged_role.description.as_deref(), + ) { + push_agent_role_warning(startup_warnings, err); + continue; + } + roles.insert(role_name, merged_role); + } + } + + Ok(roles) +} + +fn push_agent_role_warning(startup_warnings: &mut Vec, err: std::io::Error) { + let message = format!("Ignoring malformed agent role definition: {err}"); + tracing::warn!("{message}"); + startup_warnings.push(message); +} + +async fn load_agent_roles_without_layers( + fs: &dyn ExecutorFileSystem, + cfg: &ConfigToml, +) -> std::io::Result> { + let mut roles = BTreeMap::new(); + if let Some(agents_toml) = cfg.agents.as_ref() { + for (declared_role_name, role_toml) in &agents_toml.roles { + let (role_name, role) = read_declared_role(fs, declared_role_name, role_toml).await?; + validate_required_agent_role_description(&role_name, role.description.as_deref())?; + + if roles.insert(role_name.clone(), role).is_some() { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!("duplicate agent role name `{role_name}` declared in config"), + )); + } + } + } + + Ok(roles) +} + +async fn read_declared_role( + fs: &dyn ExecutorFileSystem, + declared_role_name: &str, + role_toml: &AgentRoleToml, +) -> std::io::Result<(String, AgentRoleConfig)> { + let mut role = agent_role_config_from_toml(fs, declared_role_name, role_toml).await?; + let mut role_name = declared_role_name.to_string(); + if let Some(config_file) = role.config_file.as_deref() { + let config_file = AbsolutePathBuf::from_absolute_path(config_file)?; + let parsed_file = + read_resolved_agent_role_file(fs, &config_file, Some(declared_role_name)).await?; + role_name = parsed_file.role_name; + role.description = parsed_file.description.or(role.description); + role.nickname_candidates = parsed_file.nickname_candidates.or(role.nickname_candidates); + } + + Ok((role_name, role)) +} + +fn merge_missing_role_fields(role: &mut AgentRoleConfig, fallback: &AgentRoleConfig) { + role.description = role.description.clone().or(fallback.description.clone()); + role.config_file = role.config_file.clone().or(fallback.config_file.clone()); + role.nickname_candidates = role + .nickname_candidates + .clone() + .or(fallback.nickname_candidates.clone()); +} + +fn agents_toml_from_layer( + layer_toml: &TomlValue, + config_base_dir: Option<&Path>, +) -> std::io::Result> { + let Some(agents_toml) = layer_toml.get("agents") else { + return Ok(None); + }; + + // AbsolutePathBufGuard resolves relative paths while it remains in scope. + let _guard = config_base_dir.map(AbsolutePathBufGuard::new); + agents_toml + .clone() + .try_into() + .map(Some) + .map_err(|err| std::io::Error::new(std::io::ErrorKind::InvalidData, err)) +} + +async fn agent_role_config_from_toml( + fs: &dyn ExecutorFileSystem, + role_name: &str, + role: &AgentRoleToml, +) -> std::io::Result { + let config_file = role + .config_file + .as_ref() + .map(AbsolutePathBuf::from_absolute_path) + .transpose()?; + validate_agent_role_config_file(fs, role_name, config_file.as_ref()).await?; + let description = normalize_agent_role_description( + &format!("agents.{role_name}.description"), + role.description.as_deref(), + )?; + let nickname_candidates = normalize_agent_role_nickname_candidates( + &format!("agents.{role_name}.nickname_candidates"), + role.nickname_candidates.as_deref(), + )?; + + Ok(AgentRoleConfig { + description, + config_file: config_file.map(AbsolutePathBuf::into_path_buf), + nickname_candidates, + }) +} + +async fn read_resolved_agent_role_file( + fs: &dyn ExecutorFileSystem, + path: &AbsolutePathBuf, + role_name_hint: Option<&str>, +) -> std::io::Result { + let path_uri = PathUri::from_abs_path(path); + let contents = fs + .read_file_text(&path_uri, ReadFileOptions::default(), /*sandbox*/ None) + .await?; + let config_base_dir = path.parent().unwrap_or_else(|| path.clone()); + parse_agent_role_file_contents( + &contents, + path.as_path(), + config_base_dir.as_path(), + role_name_hint, + ) +} + +fn validate_required_agent_role_description( + role_name: &str, + description: Option<&str>, +) -> std::io::Result<()> { + if description.is_some() { + Ok(()) + } else { + Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!("agent role `{role_name}` must define a description"), + )) + } +} + +async fn validate_agent_role_config_file( + fs: &dyn ExecutorFileSystem, + role_name: &str, + config_file: Option<&AbsolutePathBuf>, +) -> std::io::Result<()> { + let Some(config_file) = config_file else { + return Ok(()); + }; + + let config_file_uri = PathUri::from_abs_path(config_file); + let metadata = fs + .get_metadata( + &config_file_uri, + GetMetadataOptions::default(), + /*sandbox*/ None, + ) + .await + .map_err(|e| { + std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!( + "agents.{role_name}.config_file must point to an existing file at {}: {e}", + config_file.as_path().display() + ), + ) + })?; + if metadata.is_file { + Ok(()) + } else { + Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!( + "agents.{role_name}.config_file must point to a file: {}", + config_file.as_path().display() + ), + )) + } +} + +async fn discover_agent_roles_in_dir( + fs: &dyn ExecutorFileSystem, + agents_dir: &AbsolutePathBuf, + declared_role_files: &BTreeSet, + startup_warnings: &mut Vec, +) -> std::io::Result> { + let mut roles = BTreeMap::new(); + + for agent_file in collect_agent_role_files(fs, agents_dir).await? { + if declared_role_files.contains(agent_file.as_path()) { + continue; + } + let parsed_file = + match read_resolved_agent_role_file(fs, &agent_file, /*role_name_hint*/ None).await { + Ok(parsed_file) => parsed_file, + Err(err) => { + push_agent_role_warning(startup_warnings, err); + continue; + } + }; + let role_name = parsed_file.role_name; + if roles.contains_key(&role_name) { + push_agent_role_warning( + startup_warnings, + std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!( + "duplicate agent role name `{role_name}` discovered in {}", + agents_dir.as_path().display() + ), + ), + ); + continue; + } + roles.insert( + role_name, + AgentRoleConfig { + description: parsed_file.description, + config_file: Some(agent_file.to_path_buf()), + nickname_candidates: parsed_file.nickname_candidates, + }, + ); + } + + Ok(roles) +} diff --git a/codex-rs/app-server-protocol-noop-macros/src/lib.rs b/codex-rs/app-server-protocol-noop-macros/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..74af71abc227dffb3d29b5861cb7e37a497c3175 --- /dev/null +++ b/codex-rs/app-server-protocol-noop-macros/src/lib.rs @@ -0,0 +1,20 @@ +//! No-op schema derives for production app-server protocol builds. +//! +//! The real `ts-rs` and `schemars` derives are only needed when regenerating +//! the vendored protocol exports. Normal builds retain the annotations so the +//! protocol definitions stay readable, but use these derives to avoid +//! generating implementations that cannot be reached at runtime. + +use proc_macro::TokenStream; + +/// Accepts `#[schemars(...)]` helper attributes without generating an impl. +#[proc_macro_derive(JsonSchema, attributes(schemars))] +pub fn derive_json_schema(_input: TokenStream) -> TokenStream { + TokenStream::new() +} + +/// Accepts `#[ts(...)]` helper attributes without generating an impl. +#[proc_macro_derive(TS, attributes(ts))] +pub fn derive_ts(_input: TokenStream) -> TokenStream { + TokenStream::new() +} diff --git a/codex-rs/cloud-tasks-mock-client/src/lib.rs b/codex-rs/cloud-tasks-mock-client/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..833ea5e1f8026af59b39d38e4e6f773dd454beb3 --- /dev/null +++ b/codex-rs/cloud-tasks-mock-client/src/lib.rs @@ -0,0 +1,3 @@ +mod mock; + +pub use mock::MockClient; diff --git a/codex-rs/cloud-tasks-mock-client/src/mock.rs b/codex-rs/cloud-tasks-mock-client/src/mock.rs new file mode 100644 index 0000000000000000000000000000000000000000..08fadee371c05c8fa83cf79200300368f0b6964d --- /dev/null +++ b/codex-rs/cloud-tasks-mock-client/src/mock.rs @@ -0,0 +1,267 @@ +use chrono::Utc; +use codex_cloud_tasks_client::ApplyOutcome; +use codex_cloud_tasks_client::ApplyStatus; +use codex_cloud_tasks_client::AttemptStatus; +use codex_cloud_tasks_client::CloudBackend; +use codex_cloud_tasks_client::CloudBackendFuture; +use codex_cloud_tasks_client::CloudTaskError; +use codex_cloud_tasks_client::CreatedTask; +use codex_cloud_tasks_client::DiffSummary; +use codex_cloud_tasks_client::Result; +use codex_cloud_tasks_client::TaskId; +use codex_cloud_tasks_client::TaskListPage; +use codex_cloud_tasks_client::TaskStatus; +use codex_cloud_tasks_client::TaskSummary; +use codex_cloud_tasks_client::TaskText; +use codex_cloud_tasks_client::TurnAttempt; + +#[derive(Clone, Default)] +pub struct MockClient; + +impl MockClient { + async fn list_tasks( + &self, + _env: Option<&str>, + _limit: Option, + _cursor: Option<&str>, + ) -> Result { + // Slightly vary content by env to aid tests that rely on the mock + let rows = match _env { + Some("env-A") => vec![("T-2000", "A: First", TaskStatus::Ready)], + Some("env-B") => vec![ + ("T-3000", "B: One", TaskStatus::Ready), + ("T-3001", "B: Two", TaskStatus::Pending), + ], + _ => vec![ + ("T-1000", "Update README formatting", TaskStatus::Ready), + ("T-1001", "Fix clippy warnings in core", TaskStatus::Pending), + ("T-1002", "Add contributing guide", TaskStatus::Ready), + ], + }; + let environment_id = _env.map(str::to_string); + let environment_label = match _env { + Some("env-A") => Some("Env A".to_string()), + Some("env-B") => Some("Env B".to_string()), + Some(other) => Some(other.to_string()), + None => Some("Global".to_string()), + }; + let mut out = Vec::new(); + for (id_str, title, status) in rows { + let id = TaskId(id_str.to_string()); + let diff = mock_diff_for(&id); + let (a, d) = count_from_unified(&diff); + out.push(TaskSummary { + id, + title: title.to_string(), + status, + updated_at: Utc::now(), + environment_id: environment_id.clone(), + environment_label: environment_label.clone(), + summary: DiffSummary { + files_changed: 1, + lines_added: a, + lines_removed: d, + }, + is_review: false, + attempt_total: Some(if id_str == "T-1000" { 2 } else { 1 }), + }); + } + Ok(TaskListPage { + tasks: out, + cursor: None, + }) + } + + async fn get_task_summary(&self, id: TaskId) -> Result { + let tasks = self + .list_tasks(/*env*/ None, /*limit*/ None, /*cursor*/ None) + .await? + .tasks; + tasks + .into_iter() + .find(|t| t.id == id) + .ok_or_else(|| CloudTaskError::Msg(format!("Task {} not found (mock)", id.0))) + } + + async fn get_task_diff(&self, id: TaskId) -> Result> { + Ok(Some(mock_diff_for(&id))) + } + + async fn get_task_messages(&self, _id: TaskId) -> Result> { + Ok(vec![ + "Mock assistant output: this task contains no diff.".to_string(), + ]) + } + + async fn get_task_text(&self, _id: TaskId) -> Result { + Ok(TaskText { + prompt: Some("Why is there no diff?".to_string()), + messages: vec!["Mock assistant output: this task contains no diff.".to_string()], + turn_id: Some("mock-turn".to_string()), + sibling_turn_ids: Vec::new(), + attempt_placement: Some(0), + attempt_status: AttemptStatus::Completed, + }) + } + + async fn apply_task(&self, id: TaskId, _diff_override: Option) -> Result { + Ok(ApplyOutcome { + applied: true, + status: ApplyStatus::Success, + message: format!("Applied task {} locally (mock)", id.0), + skipped_paths: Vec::new(), + conflict_paths: Vec::new(), + }) + } + + async fn apply_task_preflight( + &self, + id: TaskId, + _diff_override: Option, + ) -> Result { + Ok(ApplyOutcome { + applied: false, + status: ApplyStatus::Success, + message: format!("Preflight passed for task {} (mock)", id.0), + skipped_paths: Vec::new(), + conflict_paths: Vec::new(), + }) + } + + async fn list_sibling_attempts( + &self, + task: TaskId, + _turn_id: String, + ) -> Result> { + if task.0 == "T-1000" { + return Ok(vec![TurnAttempt { + turn_id: "T-1000-attempt-2".to_string(), + attempt_placement: Some(1), + created_at: Some(Utc::now()), + status: AttemptStatus::Completed, + diff: Some(mock_diff_for(&task)), + messages: vec!["Mock alternate attempt".to_string()], + }]); + } + Ok(Vec::new()) + } + + async fn create_task( + &self, + env_id: &str, + prompt: &str, + git_ref: &str, + qa_mode: bool, + best_of_n: usize, + ) -> Result { + let _ = (env_id, prompt, git_ref, qa_mode, best_of_n); + let id = format!("task_local_{}", chrono::Utc::now().timestamp_millis()); + Ok(CreatedTask { id: TaskId(id) }) + } +} + +impl CloudBackend for MockClient { + fn list_tasks<'a>( + &'a self, + env: Option<&'a str>, + limit: Option, + cursor: Option<&'a str>, + ) -> CloudBackendFuture<'a, TaskListPage> { + Box::pin(MockClient::list_tasks(self, env, limit, cursor)) + } + + fn get_task_summary(&self, id: TaskId) -> CloudBackendFuture<'_, TaskSummary> { + Box::pin(MockClient::get_task_summary(self, id)) + } + + fn get_task_diff(&self, id: TaskId) -> CloudBackendFuture<'_, Option> { + Box::pin(MockClient::get_task_diff(self, id)) + } + + fn get_task_messages(&self, id: TaskId) -> CloudBackendFuture<'_, Vec> { + Box::pin(MockClient::get_task_messages(self, id)) + } + + fn get_task_text(&self, id: TaskId) -> CloudBackendFuture<'_, TaskText> { + Box::pin(MockClient::get_task_text(self, id)) + } + + fn apply_task( + &self, + id: TaskId, + diff_override: Option, + ) -> CloudBackendFuture<'_, ApplyOutcome> { + Box::pin(MockClient::apply_task(self, id, diff_override)) + } + + fn apply_task_preflight( + &self, + id: TaskId, + diff_override: Option, + ) -> CloudBackendFuture<'_, ApplyOutcome> { + Box::pin(MockClient::apply_task_preflight(self, id, diff_override)) + } + + fn list_sibling_attempts( + &self, + task: TaskId, + turn_id: String, + ) -> CloudBackendFuture<'_, Vec> { + Box::pin(MockClient::list_sibling_attempts(self, task, turn_id)) + } + + fn create_task<'a>( + &'a self, + env_id: &'a str, + prompt: &'a str, + git_ref: &'a str, + qa_mode: bool, + best_of_n: usize, + ) -> CloudBackendFuture<'a, CreatedTask> { + Box::pin(MockClient::create_task( + self, env_id, prompt, git_ref, qa_mode, best_of_n, + )) + } +} + +fn mock_diff_for(id: &TaskId) -> String { + match id.0.as_str() { + "T-1000" => { + "diff --git a/README.md b/README.md\nindex 000000..111111 100644\n--- a/README.md\n+++ b/README.md\n@@ -1,2 +1,3 @@\n Intro\n-Hello\n+Hello, world!\n+Task: T-1000\n".to_string() + } + "T-1001" => { + "diff --git a/core/src/lib.rs b/core/src/lib.rs\nindex 000000..111111 100644\n--- a/core/src/lib.rs\n+++ b/core/src/lib.rs\n@@ -1,2 +1,1 @@\n-use foo;\n use bar;\n".to_string() + } + _ => { + "diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md\nindex 000000..111111 100644\n--- /dev/null\n+++ b/CONTRIBUTING.md\n@@ -0,0 +1,3 @@\n+## Contributing\n+Please open PRs.\n+Thanks!\n".to_string() + } + } +} + +fn count_from_unified(diff: &str) -> (usize, usize) { + if let Ok(patch) = diffy::Patch::from_str(diff) { + patch + .hunks() + .iter() + .flat_map(diffy::Hunk::lines) + .fold((0, 0), |(a, d), l| match l { + diffy::Line::Insert(_) => (a + 1, d), + diffy::Line::Delete(_) => (a, d + 1), + _ => (a, d), + }) + } else { + let mut a = 0; + let mut d = 0; + for l in diff.lines() { + if l.starts_with("+++") || l.starts_with("---") || l.starts_with("@@") { + continue; + } + match l.as_bytes().first() { + Some(b'+') => a += 1, + Some(b'-') => d += 1, + _ => {} + } + } + (a, d) + } +} diff --git a/codex-rs/code-mode-protocol/src/description.rs b/codex-rs/code-mode-protocol/src/description.rs new file mode 100644 index 0000000000000000000000000000000000000000..860074e769569da5d01cbb357140b230dda1ad2f --- /dev/null +++ b/codex-rs/code-mode-protocol/src/description.rs @@ -0,0 +1,993 @@ +use codex_protocol::ToolName; +use serde::Deserialize; +use serde::Serialize; +use serde_json::Value as JsonValue; +use std::collections::BTreeMap; + +use crate::PUBLIC_TOOL_NAME; +use crate::json_schema_types::render_json_schema_to_typescript; + +const MAX_JS_SAFE_INTEGER: u64 = (1_u64 << 53) - 1; +const DEFERRED_NESTED_TOOLS_GUIDANCE: &str = r#"Some deferred nested tools may be omitted from this description. They are still available on the global `tools` object and listed in `ALL_TOOLS`. +To find one, filter `ALL_TOOLS` by `name` and `description`."#; +const LEGACY_IMAGE_HELPER_DESCRIPTION: &str = r#"`image(imageUrlOrItem: string | { image_url: string; detail?: "auto" | "low" | "high" | "original" | null } | ImageContent, detail?: "auto" | "low" | "high" | "original" | null)`: Appends an image item. `image_url` should be a base64-encoded `data:` URL. To forward an MCP tool image, pass an individual `ImageContent` block from `result.content`, for example `image(result.content[0])`. MCP image blocks may request detail with `_meta: { "codex/imageDetail": "original" }`. When provided, the second `detail` argument overrides any detail embedded in the first argument."#; +const UNIFIED_IMAGE_HELPER_DESCRIPTION: &str = r#"`image(imageUrlOrItem: string | { image_url: string } | ImageContent)`: Appends an image item. `image_url` should be a base64-encoded `data:` URL. To forward an MCP tool image, pass an individual `ImageContent` block from `result.content`, for example `image(result.content[0])`."#; +const EXEC_DESCRIPTION_TEMPLATE: &str = r#"Run JavaScript code to orchestrate/compose tool calls +- Evaluates the provided JavaScript code in a fresh V8 isolate as an async module. +- All nested tools are available on the global `tools` object, for example `await tools.exec_command(...)`. Tool names are exposed as normalized JavaScript identifiers, for example `await tools.mcp__ologs__get_profile(...)`. +- Nested tool methods take either a string or an object as their input argument. +- Nested tools return either an object or a string, based on the description. +- Runs raw JavaScript -- no Node, no file system, no network access, no console. +- Accepts raw JavaScript source text, not JSON, quoted strings, or markdown code fences. +- You may optionally start the tool input with a first-line pragma like `// @exec: {"yield_time_ms": 10000, "max_output_tokens": 1000}`. +- `yield_time_ms` asks `exec` to yield early if the script is still running. Defaults to 10000 ms. +- `max_output_tokens` sets the token budget for direct `exec` results. Defaults to 10000 tokens. +- When the JS code is fully evaluated, the isolate's lifetime ends and unawaited promises are silently discarded. + +- Global helpers: +- `exit()`: Immediately ends the current script successfully (like an early return from the top level). +- `text(value: string | number | boolean | undefined | null)`: Appends a text item. Non-string values are stringified with `JSON.stringify(...)` when possible. +- `image(imageUrlOrItem: string | { image_url: string; detail?: "auto" | "low" | "high" | "original" | null } | ImageContent, detail?: "auto" | "low" | "high" | "original" | null)`: Appends an image item. `image_url` should be a base64-encoded `data:` URL. To forward an MCP tool image, pass an individual `ImageContent` block from `result.content`, for example `image(result.content[0])`. MCP image blocks may request detail with `_meta: { "codex/imageDetail": "original" }`. When provided, the second `detail` argument overrides any detail embedded in the first argument. +- `audio(audioUrlOrItem: string | { audio_url: string } | AudioContent)`: Appends an audio item. `audio_url` should be a base64-encoded `data:` URL. To forward an MCP tool audio block, pass an individual `AudioContent` block from `result.content`, for example `audio(result.content[0])`. +- `generatedImage(result: { image_url: string; output_hint?: string })`: Appends an image-generation result and its optional output hint. HTTP(S) URLs are not supported. +- `store(key: string, value: any)`: stores a serializable value under a string key for later `exec` calls in the same session. +- `load(key: string)`: returns the stored value for a string key, or `undefined` if it is missing. +- `notify(value: string | number | boolean | undefined | null)`: immediately injects an extra `custom_tool_call_output` for the current `exec` call. Values are stringified like `text(...)`. +- `setTimeout(callback: () => void, delayMs?: number)`: schedules a callback to run later and returns a timeout id. Pending timeouts do not keep `exec` alive by themselves; await an explicit promise if you need to wait for one. +- `clearTimeout(timeoutId?: number)`: cancels a timeout created by `setTimeout`. +- `ALL_TOOLS`: metadata for the enabled nested tools as `{ name, description }` entries. +- `yield_control()`: yields the accumulated output to the model immediately while the script keeps running."#; +const WAIT_DESCRIPTION_TEMPLATE: &str = r#"- Use `wait` only after `exec` returns `Script running with cell ID ...`. +- `cell_id` identifies the running `exec` cell to resume. +- `yield_time_ms` controls how long to wait for more output before yielding again. Defaults to 10000 ms. +- `max_tokens` limits how much new output this wait call returns. Defaults to 10000 tokens. +- `terminate: true` stops the running cell; false or omitted waits for output. +- `wait` returns only the new output since the last yield, or the final completion or termination result for that cell. +- If the cell is still running, `wait` may yield again with the same `cell_id`. +- If the cell has already finished, `wait` returns the completed result and closes the cell."#; +// Based off of https://modelcontextprotocol.io/specification/draft/schema#calltoolresult +const MCP_TYPESCRIPT_PREAMBLE: &str = r#"type Role = "user" | "assistant"; +type MetaObject = Record; +type Annotations = { + audience?: Role[]; + priority?: number; + lastModified?: string; +}; +type Icon = { + src: string; + mimeType?: string; + sizes?: string[]; + theme?: "light" | "dark"; +}; +type TextResourceContents = { + uri: string; + mimeType?: string; + _meta?: MetaObject; + text: string; +}; +type BlobResourceContents = { + uri: string; + mimeType?: string; + _meta?: MetaObject; + blob: string; +}; +type TextContent = { + type: "text"; + text: string; + annotations?: Annotations; + _meta?: MetaObject; +}; +type ImageContent = { + type: "image"; + data: string; + mimeType: string; + annotations?: Annotations; + _meta?: MetaObject; +}; +type AudioContent = { + type: "audio"; + data: string; + mimeType: string; + annotations?: Annotations; + _meta?: MetaObject; +}; +type ResourceLink = { + icons?: Icon[]; + name: string; + title?: string; + uri: string; + description?: string; + mimeType?: string; + annotations?: Annotations; + size?: number; + _meta?: MetaObject; + type: "resource_link"; +}; +type EmbeddedResource = { + type: "resource"; + resource: TextResourceContents | BlobResourceContents; + annotations?: Annotations; + _meta?: MetaObject; +}; +type ContentBlock = + | TextContent + | ImageContent + | AudioContent + | ResourceLink + | EmbeddedResource; +type CallToolResult = { + _meta?: MetaObject; + content: ContentBlock[]; + isError?: boolean; + structuredContent?: TStructured; + [key: string]: unknown; +};"#; + +pub const CODE_MODE_PRAGMA_PREFIX: &str = "// @exec:"; + +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum CodeModeToolKind { + Function, + Freeform, +} + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +pub struct ToolDefinition { + pub name: String, + pub tool_name: ToolName, + pub description: String, + pub kind: CodeModeToolKind, + pub input_schema: Option, + pub output_schema: Option, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct ToolNamespaceDescription { + pub name: String, + pub description: String, +} + +#[derive(Debug, Default, Deserialize, PartialEq, Eq)] +#[serde(deny_unknown_fields)] +struct CodeModeExecPragma { + #[serde(default)] + yield_time_ms: Option, + #[serde(default)] + max_output_tokens: Option, +} + +#[derive(Debug, PartialEq, Eq)] +pub struct ParsedExecSource { + pub code: String, + pub yield_time_ms: Option, + pub max_output_tokens: Option, +} + +pub fn parse_exec_source(input: &str) -> Result { + if input.trim().is_empty() { + return Err( + "exec expects raw JavaScript source text (non-empty). Provide JS only, optionally with first-line `// @exec: {\"yield_time_ms\": 10000, \"max_output_tokens\": 1000}`.".to_string(), + ); + } + + let mut args = ParsedExecSource { + code: input.to_string(), + yield_time_ms: None, + max_output_tokens: None, + }; + + let mut lines = input.splitn(2, '\n'); + let first_line = lines.next().unwrap_or_default(); + let rest = lines.next().unwrap_or_default(); + let trimmed = first_line.trim_start(); + let Some(pragma) = trimmed.strip_prefix(CODE_MODE_PRAGMA_PREFIX) else { + return Ok(args); + }; + + if rest.trim().is_empty() { + return Err( + "exec pragma must be followed by JavaScript source on subsequent lines".to_string(), + ); + } + + let directive = pragma.trim(); + if directive.is_empty() { + return Err( + "exec pragma must be a JSON object with supported fields `yield_time_ms` and `max_output_tokens`" + .to_string(), + ); + } + + let value: serde_json::Value = serde_json::from_str(directive).map_err(|err| { + format!( + "exec pragma must be valid JSON with supported fields `yield_time_ms` and `max_output_tokens`: {err}" + ) + })?; + let object = value.as_object().ok_or_else(|| { + "exec pragma must be a JSON object with supported fields `yield_time_ms` and `max_output_tokens`" + .to_string() + })?; + for key in object.keys() { + match key.as_str() { + "yield_time_ms" | "max_output_tokens" => {} + _ => { + return Err(format!( + "exec pragma only supports `yield_time_ms` and `max_output_tokens`; got `{key}`" + )); + } + } + } + + let pragma: CodeModeExecPragma = serde_json::from_value(value).map_err(|err| { + format!( + "exec pragma fields `yield_time_ms` and `max_output_tokens` must be non-negative safe integers: {err}" + ) + })?; + if pragma + .yield_time_ms + .is_some_and(|yield_time_ms| yield_time_ms > MAX_JS_SAFE_INTEGER) + { + return Err( + "exec pragma field `yield_time_ms` must be a non-negative safe integer".to_string(), + ); + } + if pragma.max_output_tokens.is_some_and(|max_output_tokens| { + u64::try_from(max_output_tokens) + .map(|max_output_tokens| max_output_tokens > MAX_JS_SAFE_INTEGER) + .unwrap_or(true) + }) { + return Err( + "exec pragma field `max_output_tokens` must be a non-negative safe integer".to_string(), + ); + } + + args.code = rest.to_string(); + args.yield_time_ms = pragma.yield_time_ms; + args.max_output_tokens = pragma.max_output_tokens; + Ok(args) +} + +pub fn is_code_mode_nested_tool(tool_name: &str) -> bool { + tool_name != crate::PUBLIC_TOOL_NAME && tool_name != crate::WAIT_TOOL_NAME +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum ImageDetailVisibility { + Visible, + Hidden, +} + +pub fn build_exec_tool_description( + enabled_tools: &[ToolDefinition], + deferred_tools: &[ToolDefinition], + namespace_descriptions: &BTreeMap, + default_exec_yield_time_ms: u64, + code_mode_only: bool, + image_detail_visibility: ImageDetailVisibility, +) -> String { + let mut sections = Vec::new(); + sections.push(EXEC_DESCRIPTION_TEMPLATE.replace( + "Defaults to 10000 ms.", + &format!("Defaults to {default_exec_yield_time_ms} ms."), + )); + if image_detail_visibility == ImageDetailVisibility::Hidden { + sections[0] = sections[0].replace( + LEGACY_IMAGE_HELPER_DESCRIPTION, + UNIFIED_IMAGE_HELPER_DESCRIPTION, + ); + } + if !deferred_tools.is_empty() { + sections.push(DEFERRED_NESTED_TOOLS_GUIDANCE.to_string()); + } + if !code_mode_only { + return sections.join("\n\n"); + } + + let has_mcp_tools = enabled_tools + .iter() + .chain(deferred_tools) + .any(|tool| mcp_structured_content_schema(tool.output_schema.as_ref()).is_some()); + if has_mcp_tools { + sections.push(format!( + "Shared MCP Types:\n```ts\n{MCP_TYPESCRIPT_PREAMBLE}\n```" + )); + } + + if !enabled_tools.is_empty() { + let mut current_namespace: Option<&str> = None; + let mut nested_tool_sections = Vec::with_capacity(enabled_tools.len()); + + for tool in enabled_tools { + let name = tool.name.as_str(); + let nested_description = render_code_mode_sample_for_definition(tool); + let namespace_description = tool + .tool_name + .namespace + .as_ref() + .and_then(|namespace| namespace_descriptions.get(namespace)); + let next_namespace = namespace_description + .map(|namespace_description| namespace_description.name.as_str()); + if next_namespace != current_namespace { + if let Some(namespace_description) = namespace_description { + let namespace_description_text = namespace_description.description.trim(); + if !namespace_description_text.is_empty() { + nested_tool_sections.push(format!( + "## {}\n{namespace_description_text}", + namespace_description.name + )); + } + } + current_namespace = next_namespace; + } + + let global_name = normalize_code_mode_identifier(name); + let nested_description = nested_description.trim(); + if nested_description.is_empty() { + nested_tool_sections.push(render_tool_heading(&global_name, name)); + } else { + nested_tool_sections.push(format!( + "{}\n{nested_description}", + render_tool_heading(&global_name, name) + )); + } + } + + sections.push(nested_tool_sections.join("\n\n")); + } + + sections.join("\n\n") +} + +pub fn build_wait_tool_description() -> &'static str { + WAIT_DESCRIPTION_TEMPLATE +} + +pub fn normalize_code_mode_identifier(tool_key: &str) -> String { + let mut identifier = String::new(); + + for (index, ch) in tool_key.chars().enumerate() { + let is_valid = if index == 0 { + ch == '_' || ch == '$' || ch.is_ascii_alphabetic() + } else { + ch == '_' || ch == '$' || ch.is_ascii_alphanumeric() + }; + + if is_valid { + identifier.push(ch); + } else { + identifier.push('_'); + } + } + + if identifier.is_empty() { + "_".to_string() + } else { + identifier + } +} + +pub fn augment_tool_definition(mut definition: ToolDefinition) -> ToolDefinition { + if definition.name != PUBLIC_TOOL_NAME { + definition.description = render_code_mode_sample_for_definition(&definition); + } + definition +} + +pub fn enabled_tool_metadata(definition: &ToolDefinition) -> EnabledToolMetadata { + EnabledToolMetadata { + tool_name: definition.tool_name.clone(), + global_name: normalize_code_mode_identifier(&definition.name), + description: definition.description.clone(), + kind: definition.kind, + } +} + +#[derive(Clone, Debug, Eq, PartialEq, Serialize)] +pub struct EnabledToolMetadata { + pub tool_name: ToolName, + pub global_name: String, + pub description: String, + pub kind: CodeModeToolKind, +} + +pub fn render_code_mode_sample( + description: &str, + tool_name: &str, + input_name: &str, + input_type: String, + output_type: String, +) -> String { + let declaration = format!( + "declare const tools: {{ {} }};", + render_code_mode_tool_declaration(tool_name, input_name, &input_type, &output_type) + ); + format!("{description}\n\nexec tool declaration:\n```ts\n{declaration}\n```") +} + +fn render_code_mode_sample_for_definition(definition: &ToolDefinition) -> String { + let input_name = match definition.kind { + CodeModeToolKind::Function => "args", + CodeModeToolKind::Freeform => "input", + }; + let input_type = match definition.kind { + CodeModeToolKind::Function => definition + .input_schema + .as_ref() + .map(render_json_schema_to_typescript) + .unwrap_or_else(|| "unknown".to_string()), + CodeModeToolKind::Freeform => "string".to_string(), + }; + let output_type = if let Some(structured_content_schema) = + mcp_structured_content_schema(definition.output_schema.as_ref()) + { + let structured_content_type = render_json_schema_to_typescript(structured_content_schema); + if structured_content_type == "unknown" { + "CallToolResult".to_string() + } else { + format!("CallToolResult<{structured_content_type}>") + } + } else { + definition + .output_schema + .as_ref() + .map(render_json_schema_to_typescript) + .unwrap_or_else(|| "unknown".to_string()) + }; + render_code_mode_sample( + &definition.description, + &definition.name, + input_name, + input_type, + output_type, + ) +} + +fn render_code_mode_tool_declaration( + tool_name: &str, + input_name: &str, + input_type: &str, + output_type: &str, +) -> String { + let tool_name = normalize_code_mode_identifier(tool_name); + format!("{tool_name}({input_name}: {input_type}): Promise<{output_type}>;") +} + +fn render_tool_heading(global_name: &str, raw_name: &str) -> String { + if global_name == raw_name { + format!("### `{global_name}`") + } else { + format!("### `{global_name}` (`{raw_name}`)") + } +} + +fn mcp_structured_content_schema(output_schema: Option<&JsonValue>) -> Option<&JsonValue> { + let output_schema = output_schema?; + let properties = output_schema + .get("properties") + .and_then(JsonValue::as_object)?; + let content_schema = properties.get("content").and_then(JsonValue::as_object)?; + if content_schema.get("type").and_then(JsonValue::as_str) != Some("array") { + return None; + } + + if content_schema + .get("items") + .and_then(JsonValue::as_object) + .is_none_or(|items| items.get("type").and_then(JsonValue::as_str) != Some("object")) + { + return None; + } + + if properties + .get("isError") + .and_then(JsonValue::as_object) + .is_none_or(|schema| schema.get("type").and_then(JsonValue::as_str) != Some("boolean")) + { + return None; + } + + if properties + .get("_meta") + .and_then(JsonValue::as_object) + .is_none_or(|schema| schema.get("type").and_then(JsonValue::as_str) != Some("object")) + { + return None; + } + + Some( + properties + .get("structuredContent") + .unwrap_or(&JsonValue::Bool(true)), + ) +} + +#[cfg(test)] +mod tests { + use super::CodeModeToolKind; + use super::ImageDetailVisibility; + use super::ParsedExecSource; + use super::ToolDefinition; + use super::ToolNamespaceDescription; + use super::augment_tool_definition; + use super::build_exec_tool_description; + use super::normalize_code_mode_identifier; + use super::parse_exec_source; + use codex_protocol::ToolName; + use pretty_assertions::assert_eq; + use serde_json::Value as JsonValue; + use serde_json::json; + use std::collections::BTreeMap; + + fn mcp_call_tool_result_schema(structured_content_schema: JsonValue) -> JsonValue { + json!({ + "type": "object", + "properties": { + "content": { + "type": "array", + "items": { + "type": "object" + } + }, + "structuredContent": structured_content_schema, + "isError": { "type": "boolean" }, + "_meta": { "type": "object" } + }, + "required": ["content"], + "additionalProperties": false + }) + } + + #[test] + fn parse_exec_source_without_pragma() { + assert_eq!( + parse_exec_source("text('hi')").unwrap(), + ParsedExecSource { + code: "text('hi')".to_string(), + yield_time_ms: None, + max_output_tokens: None, + } + ); + } + + #[test] + fn parse_exec_source_with_pragma() { + assert_eq!( + parse_exec_source("// @exec: {\"yield_time_ms\": 10}\ntext('hi')").unwrap(), + ParsedExecSource { + code: "text('hi')".to_string(), + yield_time_ms: Some(10), + max_output_tokens: None, + } + ); + } + + #[test] + fn normalize_identifier_rewrites_invalid_characters() { + assert_eq!( + "mcp__ologs__get_profile", + normalize_code_mode_identifier("mcp__ologs__get_profile") + ); + assert_eq!( + "hidden_dynamic_tool", + normalize_code_mode_identifier("hidden-dynamic-tool") + ); + } + + #[test] + fn augment_tool_definition_appends_typed_declaration() { + let definition = ToolDefinition { + name: "hidden_dynamic_tool".to_string(), + tool_name: ToolName::plain("hidden_dynamic_tool"), + description: "Test tool".to_string(), + kind: CodeModeToolKind::Function, + input_schema: Some(json!({ + "type": "object", + "properties": { "city": { "type": "string" } }, + "required": ["city"], + "additionalProperties": false + })), + output_schema: Some(json!({ + "type": "object", + "properties": { "ok": { "type": "boolean" } }, + "required": ["ok"] + })), + }; + + let description = augment_tool_definition(definition).description; + assert!(description.contains("declare const tools")); + assert!( + description.contains( + "hidden_dynamic_tool(args: { city: string; }): Promise<{ ok: boolean; }>;" + ) + ); + } + + #[test] + fn augment_tool_definition_includes_property_descriptions_as_comments() { + let definition = ToolDefinition { + name: "weather_tool".to_string(), + tool_name: ToolName::plain("weather_tool"), + description: "Weather tool".to_string(), + kind: CodeModeToolKind::Function, + input_schema: Some(json!({ + "type": "object", + "properties": { + "weather": { + "type": "array", + "description": "look up weather for a given list of locations", + "items": { + "type": "object", + "properties": { + "location": { "type": "string" } + }, + "required": ["location"] + } + } + }, + "required": ["weather"] + })), + output_schema: Some(json!({ + "type": "object", + "properties": { + "forecast": { + "type": "string", + "description": "human readable weather forecast" + } + }, + "required": ["forecast"] + })), + }; + + let description = augment_tool_definition(definition).description; + assert!(description.contains( + r#"weather_tool(args: { + // look up weather for a given list of locations + weather: Array<{ location: string; }>; +}): Promise<{ + // human readable weather forecast + forecast: string; +}>;"# + )); + } + + #[test] + fn code_mode_types_structured_content_result_refs() { + let definition = ToolDefinition { + name: "mcp__sample__search".to_string(), + tool_name: ToolName::namespaced("mcp__sample__", "search"), + description: "Search".to_string(), + kind: CodeModeToolKind::Function, + input_schema: Some(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })), + output_schema: Some(mcp_call_tool_result_schema(json!({ + "type": "object", + "properties": { + "results": { + "type": "array", + "items": { "$ref": "#/definitions/Result~1item~0v1" } + } + }, + "required": ["results"], + "additionalProperties": false, + "definitions": { + "Result/item~v1": { + "type": "object", + "properties": { + "id": { "type": "string" }, + "score": { "type": "number" } + }, + "required": ["id", "score"], + "additionalProperties": false + } + } + }))), + }; + + let description = augment_tool_definition(definition).description; + assert!(description.contains( + "mcp__sample__search(args: {}): Promise; }>>;" + )); + } + + #[test] + fn code_mode_only_description_includes_nested_tools() { + let description = build_exec_tool_description( + &[ToolDefinition { + name: "foo".to_string(), + tool_name: ToolName::plain("foo"), + description: "bar".to_string(), + kind: CodeModeToolKind::Function, + input_schema: None, + output_schema: None, + }], + &[], + &BTreeMap::new(), + crate::DEFAULT_EXEC_YIELD_TIME_MS, + /*code_mode_only*/ true, + ImageDetailVisibility::Visible, + ); + assert!(description.contains( + "### `foo` +bar" + )); + assert!(!description.contains("do not attempt to use any other tools directly")); + } + + #[test] + fn exec_description_mentions_timeout_helpers() { + let description = build_exec_tool_description( + &[], + &[], + &BTreeMap::new(), + crate::DEFAULT_EXEC_YIELD_TIME_MS, + /*code_mode_only*/ false, + ImageDetailVisibility::Visible, + ); + assert!(description.contains("`audio(audioUrlOrItem:")); + assert!(description.contains("`setTimeout(callback: () => void, delayMs?: number)`")); + assert!(description.contains("`clearTimeout(timeoutId?: number)`")); + } + + #[test] + fn code_mode_only_description_groups_namespace_instructions_once() { + let namespace_descriptions = BTreeMap::from([( + "mcp__sample__".to_string(), + ToolNamespaceDescription { + name: "mcp__sample".to_string(), + description: "Shared namespace guidance.".to_string(), + }, + )]); + let description = build_exec_tool_description( + &[ + ToolDefinition { + name: "mcp__sample__alpha".to_string(), + tool_name: ToolName::namespaced("mcp__sample__", "alpha"), + description: "First tool".to_string(), + kind: CodeModeToolKind::Function, + input_schema: Some(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })), + output_schema: Some(mcp_call_tool_result_schema(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + }))), + }, + ToolDefinition { + name: "mcp__sample__beta".to_string(), + tool_name: ToolName::namespaced("mcp__sample__", "beta"), + description: "Second tool".to_string(), + kind: CodeModeToolKind::Function, + input_schema: Some(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })), + output_schema: Some(mcp_call_tool_result_schema(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + }))), + }, + ], + &[], + &namespace_descriptions, + crate::DEFAULT_EXEC_YIELD_TIME_MS, + /*code_mode_only*/ true, + ImageDetailVisibility::Visible, + ); + assert_eq!(description.matches("## mcp__sample").count(), 1); + assert!(description.contains("## mcp__sample\nShared namespace guidance.")); + assert!(description.contains( + "declare const tools: { mcp__sample__alpha(args: {}): Promise>; };" + )); + assert!(description.contains( + "declare const tools: { mcp__sample__beta(args: {}): Promise>; };" + )); + } + + #[test] + fn code_mode_only_description_omits_empty_namespace_sections() { + let namespace_descriptions = BTreeMap::from([( + "mcp__sample__".to_string(), + ToolNamespaceDescription { + name: "mcp__sample".to_string(), + description: String::new(), + }, + )]); + let description = build_exec_tool_description( + &[ToolDefinition { + name: "mcp__sample__alpha".to_string(), + tool_name: ToolName::namespaced("mcp__sample__", "alpha"), + description: "First tool".to_string(), + kind: CodeModeToolKind::Function, + input_schema: Some(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })), + output_schema: Some(mcp_call_tool_result_schema(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + }))), + }], + &[], + &namespace_descriptions, + crate::DEFAULT_EXEC_YIELD_TIME_MS, + /*code_mode_only*/ true, + ImageDetailVisibility::Visible, + ); + + assert!(!description.contains("## mcp__sample")); + assert!(description.contains("### `mcp__sample__alpha`")); + } + + #[test] + fn code_mode_only_description_renders_shared_mcp_types_once() { + let first_tool = augment_tool_definition(ToolDefinition { + name: "mcp__sample__alpha".to_string(), + tool_name: ToolName::namespaced("mcp__sample__", "alpha"), + description: "First tool".to_string(), + kind: CodeModeToolKind::Function, + input_schema: Some(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })), + output_schema: Some(json!({ + "type": "object", + "properties": { + "content": { + "type": "array", + "items": { + "type": "object" + } + }, + "structuredContent": { + "type": "object", + "properties": { + "echo": { "type": "string" } + }, + "required": ["echo"], + "additionalProperties": false + }, + "isError": { "type": "boolean" }, + "_meta": { "type": "object" } + }, + "required": ["content"], + "additionalProperties": false + })), + }); + let second_tool = augment_tool_definition(ToolDefinition { + name: "mcp__sample__beta".to_string(), + tool_name: ToolName::namespaced("mcp__sample__", "beta"), + description: "Second tool".to_string(), + kind: CodeModeToolKind::Function, + input_schema: Some(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })), + output_schema: Some(json!({ + "type": "object", + "properties": { + "content": { + "type": "array", + "items": { + "type": "object" + } + }, + "structuredContent": { + "type": "object", + "properties": { + "count": { "type": "integer" } + }, + "required": ["count"], + "additionalProperties": false + }, + "isError": { "type": "boolean" }, + "_meta": { "type": "object" } + }, + "required": ["content"], + "additionalProperties": false + })), + }); + + let description = build_exec_tool_description( + &[ + ToolDefinition { + name: first_tool.name, + tool_name: first_tool.tool_name, + description: "First tool".to_string(), + kind: first_tool.kind, + input_schema: first_tool.input_schema, + output_schema: first_tool.output_schema, + }, + ToolDefinition { + name: second_tool.name, + tool_name: second_tool.tool_name, + description: "Second tool".to_string(), + kind: second_tool.kind, + input_schema: second_tool.input_schema, + output_schema: second_tool.output_schema, + }, + ], + &[], + &BTreeMap::new(), + crate::DEFAULT_EXEC_YIELD_TIME_MS, + /*code_mode_only*/ true, + ImageDetailVisibility::Visible, + ); + + assert_eq!( + description + .matches("type CallToolResult") + .count(), + 1 + ); + assert_eq!(description.matches("Shared MCP Types:").count(), 1); + } + + #[test] + fn code_mode_only_description_renders_shared_mcp_types_for_deferred_tools() { + let deferred_tool = ToolDefinition { + name: "mcp__sample__alpha".to_string(), + tool_name: ToolName::namespaced("mcp__sample__", "alpha"), + description: "Deferred tool".to_string(), + kind: CodeModeToolKind::Function, + input_schema: Some(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })), + output_schema: Some(mcp_call_tool_result_schema(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + }))), + }; + + let description = build_exec_tool_description( + &[], + &[deferred_tool], + &BTreeMap::new(), + crate::DEFAULT_EXEC_YIELD_TIME_MS, + /*code_mode_only*/ true, + ImageDetailVisibility::Visible, + ); + + assert!(description.contains("Some deferred nested tools may be omitted")); + assert!(description.contains("Shared MCP Types:")); + assert!(!description.contains("### `mcp__sample__alpha`")); + } + + #[test] + fn exec_description_mentions_deferred_nested_tools_when_available() { + let description = build_exec_tool_description( + &[], + &[ToolDefinition { + name: "deferred_tool".to_string(), + tool_name: ToolName::plain("deferred_tool"), + description: "Deferred tool".to_string(), + kind: CodeModeToolKind::Function, + input_schema: None, + output_schema: None, + }], + &BTreeMap::new(), + crate::DEFAULT_EXEC_YIELD_TIME_MS, + /*code_mode_only*/ false, + ImageDetailVisibility::Visible, + ); + + assert!(description.contains("Some deferred nested tools may be omitted")); + assert!(description.contains("filter `ALL_TOOLS` by `name` and `description`")); + assert!(!description.contains("do not print the full `ALL_TOOLS` array")); + } +} diff --git a/codex-rs/code-mode-protocol/src/grpc/codex.code_mode.v1.proto b/codex-rs/code-mode-protocol/src/grpc/codex.code_mode.v1.proto new file mode 100644 index 0000000000000000000000000000000000000000..7bd175604333ee1f27bdb73f25785aa54a7f8849 --- /dev/null +++ b/codex-rs/code-mode-protocol/src/grpc/codex.code_mode.v1.proto @@ -0,0 +1,269 @@ +syntax = "proto3"; + +package codex.code_mode.v1; + +// Hosts stateful JavaScript execution and delegates nested tool calls to the +// session owner. Large tool inputs and outputs stay off the session event +// stream so independent HTTP/2 streams can make progress concurrently. +service CodeModeHost { + // Opens a session lease. The first event is always SessionOpened; dropping + // this stream closes the session and terminates its active cells. + rpc OpenSession(OpenSessionRequest) returns (stream SessionEvent); + rpc CloseSession(CloseSessionRequest) returns (CloseSessionResponse); + + // Each subscription owns an independent stream of matching invocations. An + // empty tool_names filter matches every tool. Each invocation is routed to + // exactly one matching subscription, even when filters overlap. + rpc SubscribeToToolCalls(SubscribeToToolCallsRequest) + returns (stream ToolCall); + + // Each result receives its own HTTP/2 stream, preventing a large response + // from blocking unrelated tool completions or session control events. + rpc CompleteToolCall(CompleteToolCallRequest) + returns (CompleteToolCallResponse); + rpc AcknowledgeNotification(AcknowledgeNotificationRequest) + returns (AcknowledgeNotificationResponse); + + // Emits ExecutionStarted immediately, followed by one ExecutionOutcome when + // the execution yields, completes, or is terminated. + rpc Execute(ExecuteRequest) returns (stream ExecuteEvent); + rpc Wait(WaitRequest) returns (WaitResponse); + + // Acknowledges that a canceled wait has retired before another wait starts. + rpc CancelWait(CancelWaitRequest) returns (CancelWaitResponse); + rpc Terminate(TerminateRequest) returns (WaitResponse); +} + +message OpenSessionRequest { + optional SessionCellExecutionLimits cell_execution_limits = 1; +} + +message SessionCellExecutionLimits { + optional uint64 max_yield_time_ms = 1; + optional uint64 max_heap_size_bytes = 2; +} + +message SessionEvent { + oneof event { + SessionOpened opened = 1; + ToolCallCancelled tool_call_cancelled = 2; + Notification notification = 3; + NotificationCancelled notification_cancelled = 4; + CellClosed cell_closed = 5; + } +} + +message SessionOpened { + string session_id = 1; +} + +message CloseSessionRequest { + string session_id = 1; +} + +message CloseSessionResponse {} + +message SubscribeToToolCallsRequest { + string session_id = 1; + repeated ToolName tool_names = 2; +} + +message ToolCall { + string session_id = 1; + + // Correlates callbacks with Execute before ExecutionStarted is received. + string execution_id = 2; + string cell_id = 3; + string invocation_id = 4; + string runtime_tool_call_id = 5; + ToolName tool_name = 6; + ToolKind tool_kind = 7; + optional bytes input_json = 8; + + // Starts at one and increases independently for each execution. + uint64 sequence = 9; + + // The host tool invocation span's W3C parent context for this streamed callback. + // gRPC metadata is fixed when the subscription stream opens, so each call + // carries its own context in the message. + optional string traceparent = 10; +} + +message CompleteToolCallRequest { + string session_id = 1; + string invocation_id = 2; + + oneof outcome { + ToolCallSucceeded succeeded = 3; + ToolCallFailed failed = 4; + } +} + +message ToolCallSucceeded { + bytes output_json = 1; +} + +message ToolCallFailed { + string message = 1; +} + +message CompleteToolCallResponse {} + +message ToolCallCancelled { + string invocation_id = 1; + + // Cancellation can arrive before the corresponding ToolCall because session + // control events and tool subscriptions use independent HTTP/2 streams. +} + +message Notification { + string notification_id = 1; + string execution_id = 2; + string cell_id = 3; + string call_id = 4; + string text = 5; +} + +message NotificationCancelled { + string notification_id = 1; +} + +message AcknowledgeNotificationRequest { + string session_id = 1; + string notification_id = 2; +} + +message AcknowledgeNotificationResponse {} + +message CellClosed { + string execution_id = 1; + string cell_id = 2; + + // Last tool-call sequence issued before closure. Clients may retire the cell + // immediately and reject tool calls delivered after its closure. + uint64 final_tool_call_sequence = 3; +} + +message ExecuteRequest { + string session_id = 1; + + // Chosen by the client so callbacks can be correlated before cell admission. + string execution_id = 2; + string tool_call_id = 3; + string source = 4; + repeated ToolDefinition enabled_tools = 5; + optional uint64 yield_time_ms = 6; + optional uint64 max_output_tokens = 7; +} + +message ExecuteEvent { + oneof event { + ExecutionStarted started = 1; + ExecutionOutcome outcome = 2; + } +} + +message ExecutionStarted { + string execution_id = 1; + string cell_id = 2; +} + +message WaitRequest { + string session_id = 1; + string cell_id = 2; + string wait_id = 3; + uint64 yield_time_ms = 4; +} + +message WaitResponse { + oneof state { + ExecutionOutcome live_cell = 1; + ExecutionOutcome missing_cell = 2; + } +} + +message CancelWaitRequest { + string session_id = 1; + string wait_id = 2; +} + +message CancelWaitResponse {} + +message TerminateRequest { + string session_id = 1; + string cell_id = 2; +} + +message ExecutionOutcome { + string cell_id = 1; + repeated ContentItem content_items = 2; + + // Elapsed monotonic time from receipt of this Execute, Wait, or Terminate + // request until its outcome is ready, before serialization and delivery. + // Includes nested tool waits, but not background time between requests. + // Always supplied by the host; zero is a valid measurement. + uint64 code_mode_host_duration_ns = 6; + + oneof outcome { + ExecutionYielded yielded = 3; + ExecutionTerminated terminated = 4; + ExecutionCompleted completed = 5; + } +} + +message ExecutionYielded {} + +message ExecutionTerminated {} + +message ExecutionCompleted { + optional string error_text = 1; +} + +message ToolDefinition { + string name = 1; + ToolName tool_name = 2; + string description = 3; + ToolKind kind = 4; + optional bytes input_schema_json = 5; + optional bytes output_schema_json = 6; +} + +message ToolName { + string name = 1; + optional string namespace = 2; +} + +enum ToolKind { + TOOL_KIND_UNSPECIFIED = 0; + TOOL_KIND_FUNCTION = 1; + TOOL_KIND_FREEFORM = 2; +} + +message ContentItem { + oneof item { + TextContent text = 1; + ImageContent image = 2; + AudioContent audio = 3; + } +} + +message TextContent { + string text = 1; +} + +message ImageContent { + string image_url = 1; + optional ImageDetail detail = 2; +} + +message AudioContent { + string audio_url = 1; +} + +enum ImageDetail { + IMAGE_DETAIL_UNSPECIFIED = 0; + IMAGE_DETAIL_AUTO = 1; + IMAGE_DETAIL_LOW = 2; + IMAGE_DETAIL_HIGH = 3; + IMAGE_DETAIL_ORIGINAL = 4; +} diff --git a/codex-rs/code-mode-protocol/src/grpc/mod.rs b/codex-rs/code-mode-protocol/src/grpc/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..08c88badea9209841881c01d50839de30202a149 --- /dev/null +++ b/codex-rs/code-mode-protocol/src/grpc/mod.rs @@ -0,0 +1,7 @@ +#[cfg(codex_bazel)] +pub use code_mode_proto::codex::code_mode::v1::*; + +#[cfg(not(codex_bazel))] +tonic::include_proto!("codex.code_mode.v1"); + +pub const MAX_IDENTIFIER_BYTES: usize = 256; diff --git a/codex-rs/code-mode-protocol/src/host/codec.rs b/codex-rs/code-mode-protocol/src/host/codec.rs new file mode 100644 index 0000000000000000000000000000000000000000..10d13debe87c1e56de5946e2e7ea2eb3f58ea50d --- /dev/null +++ b/codex-rs/code-mode-protocol/src/host/codec.rs @@ -0,0 +1,170 @@ +use std::io; +use std::mem::size_of; + +use serde::Serialize; +use serde::de::DeserializeOwned; +use tokio::io::AsyncRead; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWrite; +use tokio::io::AsyncWriteExt; + +/// Maximum JSON payload size accepted for one code-mode host frame. +pub const MAX_FRAME_BYTES: usize = 64 * 1024 * 1024; + +/// A serialized IPC frame that has already passed the payload size limit. +#[derive(Clone, Debug)] +pub struct EncodedFrame { + payload: Vec, +} + +impl EncodedFrame { + pub fn encode(message: &T) -> io::Result + where + T: Serialize, + { + let payload = serde_json::to_vec(message).map_err(|err| { + io::Error::new( + io::ErrorKind::InvalidData, + format!("failed to encode code-mode IPC frame: {err}"), + ) + })?; + if payload.len() > MAX_FRAME_BYTES { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!( + "code-mode IPC frame length {} exceeds {MAX_FRAME_BYTES} bytes", + payload.len() + ), + )); + } + Ok(Self { payload }) + } + + /// Returns the complete length-prefixed representation of this frame. + pub fn into_framed_bytes(self) -> Vec { + let mut bytes = Vec::with_capacity(size_of::() + self.payload.len()); + bytes.extend_from_slice(&(self.payload.len() as u32).to_le_bytes()); + bytes.extend_from_slice(&self.payload); + bytes + } + + /// Decodes exactly one complete length-prefixed frame. + pub fn decode_framed(bytes: &[u8]) -> io::Result + where + T: DeserializeOwned, + { + let length_bytes: [u8; size_of::()] = bytes + .get(..size_of::()) + .and_then(|length_bytes| length_bytes.try_into().ok()) + .ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidData, + "code-mode IPC frame is missing its length prefix", + ) + })?; + let length = u32::from_le_bytes(length_bytes) as usize; + if length > MAX_FRAME_BYTES { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("code-mode IPC frame length {length} exceeds {MAX_FRAME_BYTES} bytes"), + )); + } + + let payload = &bytes[size_of::()..]; + if payload.len() != length { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!( + "code-mode IPC frame declares {length} payload bytes but contains {}", + payload.len() + ), + )); + } + + serde_json::from_slice(payload).map_err(|err| { + io::Error::new( + io::ErrorKind::InvalidData, + format!("failed to decode code-mode IPC frame: {err}"), + ) + }) + } +} + +/// Decodes JSON messages prefixed by a four-byte little-endian payload length. +pub struct FramedReader { + reader: R, +} + +impl FramedReader +where + R: AsyncRead + Unpin, +{ + pub fn new(reader: R) -> Self { + Self { reader } + } + + /// Reads the next frame, returning `None` only for EOF at a frame boundary. + pub async fn read(&mut self) -> io::Result> + where + T: DeserializeOwned, + { + let mut length_bytes = [0_u8; size_of::()]; + if self.reader.read(&mut length_bytes[..1]).await? == 0 { + return Ok(None); + } + self.reader.read_exact(&mut length_bytes[1..]).await?; + + let length = u32::from_le_bytes(length_bytes) as usize; + if length > MAX_FRAME_BYTES { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("code-mode IPC frame length {length} exceeds {MAX_FRAME_BYTES} bytes"), + )); + } + + let mut payload = vec![0; length]; + self.reader.read_exact(&mut payload).await?; + serde_json::from_slice(&payload).map(Some).map_err(|err| { + io::Error::new( + io::ErrorKind::InvalidData, + format!("failed to decode code-mode IPC frame: {err}"), + ) + }) + } +} + +/// Encodes JSON messages with a four-byte little-endian payload length. +pub struct FramedWriter { + writer: W, +} + +impl FramedWriter +where + W: AsyncWrite + Unpin, +{ + pub fn new(writer: W) -> Self { + Self { writer } + } + + /// Writes and flushes one complete frame. + pub async fn write(&mut self, message: &T) -> io::Result<()> + where + T: Serialize, + { + self.write_frame(&EncodedFrame::encode(message)?).await + } + + /// Writes and flushes a frame encoded before it entered an I/O queue. + pub async fn write_frame(&mut self, frame: &EncodedFrame) -> io::Result<()> { + let length = u32::try_from(frame.payload.len()).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidData, + "code-mode IPC frame length exceeds u32", + ) + })?; + + self.writer.write_all(&length.to_le_bytes()).await?; + self.writer.write_all(&frame.payload).await?; + self.writer.flush().await + } +} diff --git a/codex-rs/code-mode-protocol/src/host/codec_tests.rs b/codex-rs/code-mode-protocol/src/host/codec_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..332a93766ed96768e65f0f20298d26b14db743fc --- /dev/null +++ b/codex-rs/code-mode-protocol/src/host/codec_tests.rs @@ -0,0 +1,137 @@ +use pretty_assertions::assert_eq; +use serde_json::json; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; + +use super::EncodedFrame; +use super::FramedReader; +use super::FramedWriter; +use super::MAX_FRAME_BYTES; + +#[test] +fn complete_frame_round_trips_without_a_byte_stream() { + let value = json!({"type": "session/open", "sessionId": "session-1"}); + let bytes = EncodedFrame::encode(&value) + .expect("encode frame") + .into_framed_bytes(); + + assert_eq!( + EncodedFrame::decode_framed::(&bytes).expect("decode frame"), + value + ); +} + +#[test] +fn complete_frame_rejects_truncated_and_trailing_payloads() { + let value = json!({"value": 1}); + let bytes = EncodedFrame::encode(&value) + .expect("encode frame") + .into_framed_bytes(); + + let truncated = &bytes[..bytes.len() - 1]; + let truncated_error = EncodedFrame::decode_framed::(truncated) + .expect_err("truncated frame should fail"); + assert_eq!(truncated_error.kind(), std::io::ErrorKind::InvalidData); + + let mut trailing = bytes; + trailing.push(0); + let trailing_error = EncodedFrame::decode_framed::(&trailing) + .expect_err("frame with trailing bytes should fail"); + assert_eq!(trailing_error.kind(), std::io::ErrorKind::InvalidData); +} + +#[tokio::test] +async fn frame_wire_format_is_little_endian_length_prefixed_json() { + let (writer, mut reader) = tokio::io::duplex(/*max_buf_size*/ 128); + let write = tokio::spawn(async move { + FramedWriter::new(writer) + .write(&json!({"value": 1})) + .await + .expect("write frame"); + }); + + let mut bytes = Vec::new(); + reader.read_to_end(&mut bytes).await.expect("read bytes"); + write.await.expect("writer task"); + + let payload = br#"{"value":1}"#; + let mut expected = (payload.len() as u32).to_le_bytes().to_vec(); + expected.extend_from_slice(payload); + assert_eq!(bytes, expected); +} + +#[tokio::test] +async fn fragmented_frame_round_trips() { + let value = json!({"type": "session/open", "sessionId": "session-1"}); + let payload = serde_json::to_vec(&value).expect("serialize"); + let mut bytes = (payload.len() as u32).to_le_bytes().to_vec(); + bytes.extend(payload); + + let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 128); + let write = tokio::spawn(async move { + for byte in bytes { + writer.write_all(&[byte]).await.expect("write byte"); + tokio::task::yield_now().await; + } + }); + + assert_eq!( + FramedReader::new(reader) + .read::() + .await + .expect("read frame"), + Some(value) + ); + write.await.expect("writer task"); +} + +#[tokio::test] +async fn eof_is_clean_only_at_a_frame_boundary() { + let (writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16); + drop(writer); + assert_eq!( + FramedReader::new(reader) + .read::() + .await + .expect("clean eof"), + None + ); + + let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16); + writer + .write_all(&[1, 0]) + .await + .expect("write partial header"); + drop(writer); + let err = FramedReader::new(reader) + .read::() + .await + .expect_err("truncated header"); + assert_eq!(err.kind(), std::io::ErrorKind::UnexpectedEof); +} + +#[tokio::test] +async fn oversized_and_malformed_frames_are_rejected() { + let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16); + writer + .write_all(&((MAX_FRAME_BYTES as u32) + 1).to_le_bytes()) + .await + .expect("write oversized header"); + let err = FramedReader::new(reader) + .read::() + .await + .expect_err("oversized frame"); + assert_eq!(err.kind(), std::io::ErrorKind::InvalidData); + + let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16); + writer + .write_all(&(1_u32).to_le_bytes()) + .await + .expect("write length"); + writer.write_all(b"{").await.expect("write malformed json"); + let err = FramedReader::new(reader) + .read::() + .await + .expect_err("malformed frame"); + assert_eq!(err.kind(), std::io::ErrorKind::InvalidData); +} diff --git a/codex-rs/code-mode-protocol/src/host/error.rs b/codex-rs/code-mode-protocol/src/host/error.rs new file mode 100644 index 0000000000000000000000000000000000000000..423202e44bd1ad3c949a91d6e401fb36ffdcd804 --- /dev/null +++ b/codex-rs/code-mode-protocol/src/host/error.rs @@ -0,0 +1,19 @@ +use serde::Deserialize; +use serde::Serialize; + +use super::Capability; +use super::SupportedProtocolVersions; + +/// Explains why connection negotiation was rejected before any session opened. +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")] +pub enum HandshakeRejectReason { + #[serde(rename = "noCompatibleVersion")] + NoCompatibleVersion { + supported_versions: SupportedProtocolVersions, + }, + #[serde(rename = "missingRequiredCapability")] + MissingRequiredCapability { capability: Capability }, + #[serde(rename = "invalidHello")] + InvalidHello { message: String }, +} diff --git a/codex-rs/code-mode-protocol/src/host/host_tests.rs b/codex-rs/code-mode-protocol/src/host/host_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..591b29e5dd6871a4cc676b3b12809bde7e6f6dc0 --- /dev/null +++ b/codex-rs/code-mode-protocol/src/host/host_tests.rs @@ -0,0 +1,847 @@ +use std::fmt::Debug; + +use pretty_assertions::assert_eq; +use serde::Serialize; +use serde::de::DeserializeOwned; +use serde_json::Value; +use serde_json::json; + +use super::Capability; +use super::CapabilitySet; +use super::ClientHello; +use super::ClientToHost; +use super::DelegateRequest; +use super::DelegateRequestId; +use super::DelegateResponse; +use super::HandshakeRejectReason; +use super::HostHello; +use super::HostRequest; +use super::HostResponse; +use super::HostToClient; +use super::ProtocolVersion; +use super::RequestId; +use super::SessionId; +use super::SupportedProtocolVersions; +use super::WireCellId; +use super::WireContentItem; +use super::WireExecuteRequest; +use super::WireImageDetail; +use super::WireNestedToolCall; +use super::WireResult; +use super::WireRuntimeResponse; +use super::WireSessionCellExecutionLimits; +use super::WireToolDefinition; +use super::WireToolKind; +use super::WireToolName; +use super::WireWaitOutcome; +use super::WireWaitRequest; +use crate::CodeModeSessionCellExecutionLimits; +use crate::ExecuteRequest; + +fn session_id() -> SessionId { + SessionId::new("session-1").expect("valid session ID") +} + +fn cell_id(value: &str) -> WireCellId { + WireCellId::new(value) +} + +fn request_id(value: i64) -> RequestId { + RequestId::new(value) +} + +fn delegate_request_id(value: i64) -> DelegateRequestId { + DelegateRequestId::new(value) +} + +fn capability(value: &str) -> Capability { + Capability::new(value).expect("valid capability") +} + +fn supported_versions() -> SupportedProtocolVersions { + SupportedProtocolVersions::try_new([ProtocolVersion::V1]) + .expect("nonempty unique protocol versions") +} + +fn assert_wire_round_trip(message: T, encoded: Value) +where + T: Debug + DeserializeOwned + PartialEq + Serialize, +{ + assert_eq!(serde_json::to_value(&message).expect("serialize"), encoded); + assert_eq!( + serde_json::from_value::(encoded).expect("deserialize"), + message + ); +} + +fn execute_request() -> WireExecuteRequest { + WireExecuteRequest { + tool_call_id: "call-1".to_string(), + enabled_tools: vec![ + WireToolDefinition { + name: "function_tool".to_string(), + tool_name: WireToolName { + name: "function_tool".to_string(), + namespace: None, + }, + description: "function tool".to_string(), + kind: WireToolKind::Function, + input_schema: Some(json!({ "type": "object" })), + output_schema: None, + }, + WireToolDefinition { + name: "freeform_tool".to_string(), + tool_name: WireToolName { + name: "freeform_tool".to_string(), + namespace: Some("mcp__sample__".to_string()), + }, + description: "freeform tool".to_string(), + kind: WireToolKind::Freeform, + input_schema: None, + output_schema: Some(json!({ "type": "string" })), + }, + ], + source: "text('hello');".to_string(), + yield_time_ms: Some(25), + max_output_tokens: Some(100), + } +} + +fn content_items() -> Vec { + vec![ + WireContentItem::InputText { + text: "hello".to_string(), + }, + WireContentItem::InputImage { + image_url: "data:image/png;base64,none".to_string(), + detail: None, + }, + WireContentItem::InputImage { + image_url: "data:image/png;base64,auto".to_string(), + detail: Some(WireImageDetail::Auto), + }, + WireContentItem::InputImage { + image_url: "data:image/png;base64,low".to_string(), + detail: Some(WireImageDetail::Low), + }, + WireContentItem::InputImage { + image_url: "data:image/png;base64,high".to_string(), + detail: Some(WireImageDetail::High), + }, + WireContentItem::InputImage { + image_url: "data:image/png;base64,original".to_string(), + detail: Some(WireImageDetail::Original), + }, + WireContentItem::InputAudio { + audio_url: "data:audio/wav;base64,YXVkaW8=".to_string(), + }, + ] +} + +fn content_items_json() -> Value { + json!([ + { "type": "input_text", "text": "hello" }, + { "type": "input_image", "image_url": "data:image/png;base64,none" }, + { + "type": "input_image", + "image_url": "data:image/png;base64,auto", + "detail": "auto", + }, + { + "type": "input_image", + "image_url": "data:image/png;base64,low", + "detail": "low", + }, + { + "type": "input_image", + "image_url": "data:image/png;base64,high", + "detail": "high", + }, + { + "type": "input_image", + "image_url": "data:image/png;base64,original", + "detail": "original", + }, + { + "type": "input_audio", + "audio_url": "data:audio/wav;base64,YXVkaW8=", + }, + ]) +} + +#[test] +fn handshake_v1_variants_are_pinned() { + assert_wire_round_trip( + ClientToHost::ClientHello( + ClientHello::new( + supported_versions(), + CapabilitySet::try_new([capability("required")]).expect("valid required set"), + CapabilitySet::try_new([capability("optional")]).expect("valid optional set"), + ) + .expect("disjoint capabilities"), + ), + json!({ + "type": "connection/hello", + "supportedVersions": [1], + "requiredCapabilities": ["required"], + "optionalCapabilities": ["optional"], + }), + ); + assert_wire_round_trip( + HostToClient::HostHello(HostHello::new( + ProtocolVersion::V1, + CapabilitySet::try_new([capability("required")]).expect("valid capabilities"), + )), + json!({ + "type": "connection/ready", + "selectedVersion": 1, + "capabilities": ["required"], + }), + ); + for (reason, encoded) in [ + ( + HandshakeRejectReason::NoCompatibleVersion { + supported_versions: supported_versions(), + }, + json!({ + "type": "connection/rejected", + "reason": { + "type": "noCompatibleVersion", + "supportedVersions": [1], + }, + }), + ), + ( + HandshakeRejectReason::MissingRequiredCapability { + capability: capability("required"), + }, + json!({ + "type": "connection/rejected", + "reason": { + "type": "missingRequiredCapability", + "capability": "required", + }, + }), + ), + ( + HandshakeRejectReason::InvalidHello { + message: "invalid hello".to_string(), + }, + json!({ + "type": "connection/rejected", + "reason": { + "type": "invalidHello", + "message": "invalid hello", + }, + }), + ), + ] { + assert_wire_round_trip(HostToClient::HandshakeRejected { reason }, encoded); + } +} + +#[test] +fn open_session_serializes_optional_cell_execution_limits() { + assert_wire_round_trip( + HostRequest::OpenSession { + session_id: session_id(), + cell_execution_limits: Some(WireSessionCellExecutionLimits { + max_yield_time_ms: Some(250), + max_heap_size_bytes: Some(16 * 1024 * 1024), + }), + }, + json!({ + "method": "session/open", + "sessionId": "session-1", + "cellExecutionLimits": { + "maxYieldTimeMs": 250, + "maxHeapSizeBytes": 16 * 1024 * 1024, + }, + }), + ); +} + +#[test] +fn session_cell_execution_limits_convert_between_domain_and_wire() { + let domain_limits = CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(250), + max_heap_size_bytes: Some(16_usize * 1024 * 1024), + }; + let wire_limits = WireSessionCellExecutionLimits { + max_yield_time_ms: Some(250), + max_heap_size_bytes: Some(16_u64 * 1024 * 1024), + }; + + assert_eq!( + WireSessionCellExecutionLimits::try_from(domain_limits.clone()) + .expect("domain limits convert to wire limits"), + wire_limits + ); + assert_eq!( + CodeModeSessionCellExecutionLimits::try_from(wire_limits) + .expect("wire limits convert to domain limits"), + domain_limits + ); +} + +#[cfg(target_pointer_width = "32")] +#[test] +fn session_cell_execution_limits_reject_heap_sizes_that_exceed_usize() { + let wire_limits = WireSessionCellExecutionLimits { + max_yield_time_ms: None, + max_heap_size_bytes: Some(u64::from(u32::MAX) + 1), + }; + + assert!(CodeModeSessionCellExecutionLimits::try_from(wire_limits).is_err()); +} + +#[test] +fn client_to_host_v1_variants_are_pinned() { + let execute_request = execute_request(); + for (id, request, encoded_request) in [ + ( + request_id(/*value*/ 1), + HostRequest::OpenSession { + session_id: session_id(), + cell_execution_limits: None, + }, + json!({ "method": "session/open", "sessionId": "session-1" }), + ), + ( + request_id(/*value*/ 2), + HostRequest::Execute { + session_id: session_id(), + request: execute_request, + }, + json!({ + "method": "session/execute", + "sessionId": "session-1", + "request": { + "tool_call_id": "call-1", + "enabled_tools": [ + { + "name": "function_tool", + "tool_name": { "name": "function_tool", "namespace": null }, + "description": "function tool", + "kind": "function", + "input_schema": { "type": "object" }, + "output_schema": null, + }, + { + "name": "freeform_tool", + "tool_name": { + "name": "freeform_tool", + "namespace": "mcp__sample__", + }, + "description": "freeform tool", + "kind": "freeform", + "input_schema": null, + "output_schema": { "type": "string" }, + }, + ], + "source": "text('hello');", + "yield_time_ms": 25, + "max_output_tokens": 100, + }, + }), + ), + ( + request_id(/*value*/ 3), + HostRequest::Wait { + session_id: session_id(), + request: WireWaitRequest { + cell_id: cell_id("cell-1"), + yield_time_ms: 50, + }, + }, + json!({ + "method": "session/wait", + "sessionId": "session-1", + "request": { "cell_id": "cell-1", "yield_time_ms": 50 }, + }), + ), + ( + request_id(/*value*/ 4), + HostRequest::Terminate { + session_id: session_id(), + cell_id: cell_id("cell-1"), + }, + json!({ + "method": "session/terminate", + "sessionId": "session-1", + "cellId": "cell-1", + }), + ), + ( + request_id(/*value*/ 5), + HostRequest::ShutdownSession { + session_id: session_id(), + }, + json!({ "method": "session/shutdown", "sessionId": "session-1" }), + ), + ] { + assert_wire_round_trip( + ClientToHost::Request { id, request }, + json!({ + "type": "operation/request", + "id": id, + "request": encoded_request, + }), + ); + } + + for (id, result, encoded_result) in [ + ( + delegate_request_id(/*value*/ 6), + WireResult::Ok { + value: DelegateResponse::ToolResult { + result: json!({ "answer": 42 }), + }, + }, + json!({ + "status": "ok", + "value": { "type": "tool/result", "result": { "answer": 42 } }, + }), + ), + ( + delegate_request_id(/*value*/ 7), + WireResult::Ok { + value: DelegateResponse::NotificationDelivered, + }, + json!({ + "status": "ok", + "value": { "type": "notification/delivered" }, + }), + ), + ( + delegate_request_id(/*value*/ 8), + WireResult::Err { + message: "delegate failed".to_string(), + }, + json!({ "status": "error", "message": "delegate failed" }), + ), + ] { + assert_wire_round_trip( + ClientToHost::DelegateResponse { id, result }, + json!({ + "type": "delegate/response", + "id": id, + "result": encoded_result, + }), + ); + } + + assert_wire_round_trip( + ClientToHost::CancelRequest { + id: request_id(/*value*/ 9), + }, + json!({ + "type": "operation/cancel", + "id": 9, + }), + ); +} + +#[test] +fn host_to_client_v1_variants_are_pinned() { + for (id, response, encoded_response) in [ + ( + request_id(/*value*/ 1), + HostResponse::SessionReady { + session_id: session_id(), + }, + json!({ "type": "session/ready", "sessionId": "session-1" }), + ), + ( + request_id(/*value*/ 2), + HostResponse::ExecutionStarted { + cell_id: cell_id("cell-1"), + }, + json!({ "type": "execution/started", "cellId": "cell-1" }), + ), + ( + request_id(/*value*/ 3), + HostResponse::WaitCompleted { + outcome: WireWaitOutcome::LiveCell(WireRuntimeResponse::Yielded { + code_mode_host_duration_ns: 0, + cell_id: cell_id("cell-1"), + content_items: content_items(), + }), + }, + json!({ + "type": "wait/completed", + "outcome": { + "LiveCell": { + "Yielded": { + "cell_id": "cell-1", + "content_items": content_items_json(), + "code_mode_host_duration_ns": 0, + }, + }, + }, + }), + ), + ( + request_id(/*value*/ 4), + HostResponse::WaitCompleted { + outcome: WireWaitOutcome::MissingCell(WireRuntimeResponse::Result { + code_mode_host_duration_ns: 0, + cell_id: cell_id("missing-cell"), + content_items: Vec::new(), + error_text: Some("cell not found".to_string()), + }), + }, + json!({ + "type": "wait/completed", + "outcome": { + "MissingCell": { + "Result": { + "cell_id": "missing-cell", + "content_items": [], + "error_text": "cell not found", + "code_mode_host_duration_ns": 0, + }, + }, + }, + }), + ), + ( + request_id(/*value*/ 5), + HostResponse::SessionClosed { + session_id: session_id(), + }, + json!({ "type": "session/closed", "sessionId": "session-1" }), + ), + ] { + assert_wire_round_trip( + HostToClient::Response { + id, + result: WireResult::Ok { value: response }, + }, + json!({ + "type": "operation/response", + "id": id, + "result": { "status": "ok", "value": encoded_response }, + }), + ); + } + assert_wire_round_trip( + HostToClient::Response { + id: request_id(/*value*/ 6), + result: WireResult::Err { + message: "operation failed".to_string(), + }, + }, + json!({ + "type": "operation/response", + "id": 6, + "result": { "status": "error", "message": "operation failed" }, + }), + ); + + assert_wire_round_trip( + HostToClient::InitialResponse { + id: request_id(/*value*/ 7), + result: WireResult::Ok { + value: WireRuntimeResponse::Terminated { + code_mode_host_duration_ns: 0, + cell_id: cell_id("cell-1"), + content_items: Vec::new(), + }, + }, + }, + json!({ + "type": "execute/initialResponse", + "id": 7, + "result": { + "status": "ok", + "value": { + "Terminated": { + "cell_id": "cell-1", + "content_items": [], + "code_mode_host_duration_ns": 0, + }, + }, + }, + }), + ); + assert_wire_round_trip( + HostToClient::InitialResponse { + id: request_id(/*value*/ 8), + result: WireResult::Err { + message: "execution failed".to_string(), + }, + }, + json!({ + "type": "execute/initialResponse", + "id": 8, + "result": { "status": "error", "message": "execution failed" }, + }), + ); + + assert_wire_round_trip( + HostToClient::DelegateRequest { + id: delegate_request_id(/*value*/ 9), + session_id: session_id(), + request: DelegateRequest::InvokeTool { + invocation: WireNestedToolCall { + cell_id: cell_id("cell-1"), + runtime_tool_call_id: "runtime-call-1".to_string(), + tool_name: WireToolName { + name: "freeform_tool".to_string(), + namespace: Some("mcp__sample__".to_string()), + }, + tool_kind: WireToolKind::Freeform, + input: Some(json!({ "value": 1 })), + }, + }, + }, + json!({ + "type": "delegate/request", + "id": 9, + "sessionId": "session-1", + "request": { + "type": "tool/invoke", + "invocation": { + "cell_id": "cell-1", + "runtime_tool_call_id": "runtime-call-1", + "tool_name": { + "name": "freeform_tool", + "namespace": "mcp__sample__", + }, + "tool_kind": "freeform", + "input": { "value": 1 }, + }, + }, + }), + ); + assert_wire_round_trip( + HostToClient::DelegateRequest { + id: delegate_request_id(/*value*/ 10), + session_id: session_id(), + request: DelegateRequest::Notify { + call_id: "call-1".to_string(), + cell_id: cell_id("cell-1"), + text: "important".to_string(), + }, + }, + json!({ + "type": "delegate/request", + "id": 10, + "sessionId": "session-1", + "request": { + "type": "notification/send", + "callId": "call-1", + "cellId": "cell-1", + "text": "important", + }, + }), + ); + assert_wire_round_trip( + HostToClient::CancelDelegateRequest { + id: delegate_request_id(/*value*/ 11), + }, + json!({ "type": "delegate/cancel", "id": 11 }), + ); + assert_wire_round_trip( + HostToClient::CellClosed { + session_id: session_id(), + cell_id: cell_id("cell-1"), + }, + json!({ + "type": "cell/closed", + "sessionId": "session-1", + "cellId": "cell-1", + }), + ); +} + +#[test] +fn execute_request_integer_bounds_are_enforced() { + let wire_request = execute_request(); + let domain_request = ExecuteRequest::try_from(wire_request.clone()) + .expect("valid wire request converts to the domain"); + assert_eq!( + WireExecuteRequest::try_from(domain_request.clone()) + .expect("valid domain request converts to the wire"), + wire_request + ); + + let too_large = ExecuteRequest { + max_output_tokens: Some(usize::try_from(i32::MAX).expect("i32::MAX fits usize") + 1), + ..domain_request + }; + assert!(WireExecuteRequest::try_from(too_large).is_err()); + + let negative = WireExecuteRequest { + max_output_tokens: Some(-1), + ..wire_request + }; + assert!(ExecuteRequest::try_from(negative).is_err()); +} + +#[test] +fn invalid_protocol_states_cannot_be_constructed_or_decoded() { + assert!(SessionId::new("").is_err()); + assert!(Capability::new(" ").is_err()); + assert!(ProtocolVersion::new(/*value*/ 0).is_none()); + assert!(SupportedProtocolVersions::try_new([]).is_err()); + assert!( + SupportedProtocolVersions::try_new([ProtocolVersion::V1, ProtocolVersion::V1]).is_err() + ); + assert!(CapabilitySet::try_new([capability("same"), capability("same")]).is_err()); + + let version_two = ProtocolVersion::new(/*value*/ 2).expect("valid protocol version"); + let versions = SupportedProtocolVersions::try_new([ProtocolVersion::V1, version_two]) + .expect("valid versions"); + assert!(versions.contains(ProtocolVersion::V1)); + assert_eq!( + versions.iter().collect::>(), + vec![ProtocolVersion::V1, version_two] + ); + + let overlapping = capability("overlapping"); + assert!( + ClientHello::new( + supported_versions(), + CapabilitySet::try_new([overlapping.clone()]).expect("valid required set"), + CapabilitySet::try_new([overlapping]).expect("valid optional set"), + ) + .is_err() + ); + + for invalid in [ + json!({ + "type": "operation/request", + "id": 1, + "request": { "method": "session/open", "sessionId": "" }, + }), + json!({ + "type": "connection/hello", + "supportedVersions": [], + "requiredCapabilities": [], + "optionalCapabilities": [], + }), + json!({ + "type": "connection/hello", + "supportedVersions": [1], + "requiredCapabilities": ["overlapping"], + "optionalCapabilities": ["overlapping"], + }), + ] { + assert!(serde_json::from_value::(invalid).is_err()); + } +} + +#[test] +fn every_nested_v1_object_rejects_unknown_fields() { + assert!( + serde_json::from_value::(json!({ + "type": "operation/request", + "id": 1, + "request": { "method": "session/open", "sessionId": "session-1" }, + "unexpected": true, + })) + .is_err() + ); + assert!( + serde_json::from_value::(json!({ + "method": "session/open", + "sessionId": "session-1", + "unexpected": true, + })) + .is_err() + ); + assert!( + serde_json::from_value::(json!({ + "method": "session/open", + "sessionId": "session-1", + "cellExecutionLimits": { + "maxYieldTimeMs": 250, + "unexpected": true, + }, + })) + .is_err() + ); + assert!( + serde_json::from_value::(json!({ + "tool_call_id": "call-1", + "enabled_tools": [], + "source": "text('hello');", + "yield_time_ms": null, + "max_output_tokens": null, + "unexpected": true, + })) + .is_err() + ); + assert!( + serde_json::from_value::(json!({ + "name": "tool", + "tool_name": { "name": "tool", "namespace": null }, + "description": "tool", + "kind": "function", + "input_schema": null, + "output_schema": null, + "unexpected": true, + })) + .is_err() + ); + assert!( + serde_json::from_value::(json!({ + "name": "tool", + "namespace": null, + "unexpected": true, + })) + .is_err() + ); + assert!( + serde_json::from_value::(json!({ + "cell_id": "cell-1", + "yield_time_ms": 50, + "unexpected": true, + })) + .is_err() + ); + assert!( + serde_json::from_value::(json!({ + "Yielded": { + "cell_id": "cell-1", + "content_items": [], + "code_mode_host_duration_ns": 0, + "unexpected": true, + }, + })) + .is_err() + ); + assert!( + serde_json::from_value::(json!({ + "type": "input_text", + "text": "hello", + "unexpected": true, + })) + .is_err() + ); + assert!( + serde_json::from_value::(json!({ + "cell_id": "cell-1", + "runtime_tool_call_id": "runtime-call-1", + "tool_name": { "name": "tool", "namespace": null }, + "tool_kind": "function", + "input": null, + "unexpected": true, + })) + .is_err() + ); + assert!( + serde_json::from_value::(json!({ + "type": "operation/response", + "id": 1, + "result": { + "status": "ok", + "value": { "type": "session/ready", "sessionId": "session-1" }, + }, + "unexpected": true, + })) + .is_err() + ); +} diff --git a/codex-rs/code-mode-protocol/src/host/message.rs b/codex-rs/code-mode-protocol/src/host/message.rs new file mode 100644 index 0000000000000000000000000000000000000000..bc442426a3b0b1d219c9da42d6f120c995d34c43 --- /dev/null +++ b/codex-rs/code-mode-protocol/src/host/message.rs @@ -0,0 +1,264 @@ +use std::fmt; + +use serde::Deserialize; +use serde::Serialize; +use serde_json::Value as JsonValue; + +use super::Capability; +use super::CapabilitySet; +use super::DelegateRequestId; +use super::HandshakeRejectReason; +use super::ProtocolVersion; +use super::RequestId; +use super::SessionId; +use super::SupportedProtocolVersions; +use super::WireCellId; +use super::WireExecuteRequest; +use super::WireNestedToolCall; +use super::WireRuntimeResponse; +use super::WireSessionCellExecutionLimits; +use super::WireWaitOutcome; +use super::WireWaitRequest; + +#[derive(Clone, Debug, PartialEq, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct ClientHello { + supported_versions: SupportedProtocolVersions, + required_capabilities: CapabilitySet, + optional_capabilities: CapabilitySet, +} + +impl ClientHello { + pub fn new( + supported_versions: SupportedProtocolVersions, + required_capabilities: CapabilitySet, + optional_capabilities: CapabilitySet, + ) -> Result { + if let Some(capability) = required_capabilities + .iter() + .find(|capability| optional_capabilities.contains(capability)) + { + return Err(ClientHelloError::OverlappingCapability(capability.clone())); + } + Ok(Self { + supported_versions, + required_capabilities, + optional_capabilities, + }) + } + + pub fn supported_versions(&self) -> &SupportedProtocolVersions { + &self.supported_versions + } + + pub fn required_capabilities(&self) -> &CapabilitySet { + &self.required_capabilities + } + + pub fn optional_capabilities(&self) -> &CapabilitySet { + &self.optional_capabilities + } +} + +#[derive(Deserialize)] +#[serde(deny_unknown_fields, rename_all = "camelCase")] +struct ClientHelloWire { + supported_versions: SupportedProtocolVersions, + required_capabilities: CapabilitySet, + optional_capabilities: CapabilitySet, +} + +impl<'de> Deserialize<'de> for ClientHello { + fn deserialize(deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + let wire = ClientHelloWire::deserialize(deserializer)?; + Self::new( + wire.supported_versions, + wire.required_capabilities, + wire.optional_capabilities, + ) + .map_err(serde::de::Error::custom) + } +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum ClientHelloError { + OverlappingCapability(Capability), +} + +impl fmt::Display for ClientHelloError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::OverlappingCapability(capability) => write!( + formatter, + "capability `{capability}` cannot be both required and optional" + ), + } + } +} + +impl std::error::Error for ClientHelloError {} + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(deny_unknown_fields, rename_all = "camelCase")] +pub struct HostHello { + selected_version: ProtocolVersion, + capabilities: CapabilitySet, +} + +impl HostHello { + pub fn new(selected_version: ProtocolVersion, capabilities: CapabilitySet) -> Self { + Self { + selected_version, + capabilities, + } + } + + pub fn selected_version(&self) -> ProtocolVersion { + self.selected_version + } + + pub fn capabilities(&self) -> &CapabilitySet { + &self.capabilities + } +} + +/// Messages sent from a client to the code-mode host. +#[derive(Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")] +pub enum ClientToHost { + #[serde(rename = "connection/hello")] + ClientHello(ClientHello), + #[serde(rename = "operation/request")] + Request { id: RequestId, request: HostRequest }, + #[serde(rename = "operation/cancel")] + CancelRequest { id: RequestId }, + #[serde(rename = "delegate/response")] + DelegateResponse { + id: DelegateRequestId, + result: WireResult, + }, +} + +/// Messages sent from the code-mode host to a client. +#[derive(Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")] +pub enum HostToClient { + #[serde(rename = "connection/ready")] + HostHello(HostHello), + #[serde(rename = "connection/rejected")] + HandshakeRejected { reason: HandshakeRejectReason }, + #[serde(rename = "operation/response")] + Response { + id: RequestId, + result: WireResult, + }, + #[serde(rename = "execute/initialResponse")] + InitialResponse { + id: RequestId, + result: WireResult, + }, + #[serde(rename = "delegate/request")] + DelegateRequest { + id: DelegateRequestId, + session_id: SessionId, + request: DelegateRequest, + }, + #[serde(rename = "delegate/cancel")] + CancelDelegateRequest { id: DelegateRequestId }, + #[serde(rename = "cell/closed")] + CellClosed { + session_id: SessionId, + cell_id: WireCellId, + }, +} + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields, tag = "method", rename_all_fields = "camelCase")] +pub enum HostRequest { + #[serde(rename = "session/open")] + OpenSession { + session_id: SessionId, + #[serde(default, skip_serializing_if = "Option::is_none")] + cell_execution_limits: Option, + }, + #[serde(rename = "session/execute")] + Execute { + session_id: SessionId, + request: WireExecuteRequest, + }, + #[serde(rename = "session/wait")] + Wait { + session_id: SessionId, + request: WireWaitRequest, + }, + #[serde(rename = "session/terminate")] + Terminate { + session_id: SessionId, + cell_id: WireCellId, + }, + #[serde(rename = "session/shutdown")] + ShutdownSession { session_id: SessionId }, +} + +#[derive(Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")] +pub enum HostResponse { + #[serde(rename = "session/ready")] + SessionReady { session_id: SessionId }, + #[serde(rename = "execution/started")] + ExecutionStarted { cell_id: WireCellId }, + #[serde(rename = "wait/completed")] + WaitCompleted { outcome: WireWaitOutcome }, + #[serde(rename = "session/closed")] + SessionClosed { session_id: SessionId }, +} + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")] +pub enum DelegateRequest { + #[serde(rename = "tool/invoke")] + InvokeTool { invocation: WireNestedToolCall }, + #[serde(rename = "notification/send")] + Notify { + call_id: String, + cell_id: WireCellId, + text: String, + }, +} + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")] +pub enum DelegateResponse { + #[serde(rename = "tool/result")] + ToolResult { result: JsonValue }, + #[serde(rename = "notification/delivered")] + NotificationDelivered, +} + +#[derive(Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields, tag = "status", rename_all_fields = "camelCase")] +pub enum WireResult { + #[serde(rename = "ok")] + Ok { value: T }, + #[serde(rename = "error")] + Err { message: String }, +} + +impl WireResult { + pub fn from_result(result: Result) -> Self { + match result { + Ok(value) => Self::Ok { value }, + Err(message) => Self::Err { message }, + } + } + + pub fn into_result(self) -> Result { + match self { + Self::Ok { value } => Ok(value), + Self::Err { message } => Err(message), + } + } +} diff --git a/codex-rs/code-mode-protocol/src/host/mod.rs b/codex-rs/code-mode-protocol/src/host/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..43bb7111625e3afa0ffe8dd3ddd66318347ad06e --- /dev/null +++ b/codex-rs/code-mode-protocol/src/host/mod.rs @@ -0,0 +1,62 @@ +//! Messages and framing for the code-mode host boundary. +//! +//! Protocol version 1 multiplexes session operations and delegate callbacks by +//! request ID over one ordered connection. + +mod codec; +mod error; +mod message; +mod payload; +mod types; + +/// Maximum number of unresolved delegate callbacks allowed per host connection. +pub const MAX_PENDING_DELEGATE_CALLS: usize = 1_024; + +/// Negotiated support for cell execution resource limits on `session/open`. +pub const SESSION_RESOURCE_LIMITS_CAPABILITY: &str = "session-cell-execution-resource-limits"; + +pub use codec::EncodedFrame; +pub use codec::FramedReader; +pub use codec::FramedWriter; +pub use codec::MAX_FRAME_BYTES; +pub use error::HandshakeRejectReason; +pub use message::ClientHello; +pub use message::ClientHelloError; +pub use message::ClientToHost; +pub use message::DelegateRequest; +pub use message::DelegateResponse; +pub use message::HostHello; +pub use message::HostRequest; +pub use message::HostResponse; +pub use message::HostToClient; +pub use message::WireResult; +pub use payload::WireCellId; +pub use payload::WireContentItem; +pub use payload::WireExecuteRequest; +pub use payload::WireImageDetail; +pub use payload::WireNestedToolCall; +pub use payload::WireRuntimeResponse; +pub use payload::WireSessionCellExecutionLimits; +pub use payload::WireToolDefinition; +pub use payload::WireToolKind; +pub use payload::WireToolName; +pub use payload::WireWaitOutcome; +pub use payload::WireWaitRequest; +pub use types::Capability; +pub use types::CapabilitySet; +pub use types::DelegateRequestId; +pub use types::DuplicateCapability; +pub use types::InvalidIdentifier; +pub use types::InvalidSupportedProtocolVersions; +pub use types::ProtocolVersion; +pub use types::RequestId; +pub use types::SessionId; +pub use types::SupportedProtocolVersions; + +#[cfg(test)] +#[path = "host_tests.rs"] +mod tests; + +#[cfg(test)] +#[path = "codec_tests.rs"] +mod codec_tests; diff --git a/codex-rs/code-mode-protocol/src/host/payload.rs b/codex-rs/code-mode-protocol/src/host/payload.rs new file mode 100644 index 0000000000000000000000000000000000000000..392a80d24acb2f63ef339dae82b44654c6066b1c --- /dev/null +++ b/codex-rs/code-mode-protocol/src/host/payload.rs @@ -0,0 +1,492 @@ +use std::num::TryFromIntError; +use std::time::Duration; + +use codex_protocol::ToolName; +use serde::Deserialize; +use serde::Serialize; +use serde_json::Value as JsonValue; + +use crate::CellId; +use crate::CodeModeNestedToolCall; +use crate::CodeModeSessionCellExecutionLimits; +use crate::CodeModeToolKind; +use crate::ExecuteRequest; +use crate::FunctionCallOutputContentItem; +use crate::ImageDetail; +use crate::MissingCodeModeHostDuration; +use crate::RuntimeResponse; +use crate::ToolDefinition; +use crate::WaitOutcome; +use crate::WaitRequest; + +/// The per-cell execution limits carried by a V1 session-open request. +#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)] +#[serde(deny_unknown_fields, rename_all = "camelCase")] +pub struct WireSessionCellExecutionLimits { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_yield_time_ms: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_heap_size_bytes: Option, +} + +impl TryFrom for WireSessionCellExecutionLimits { + type Error = TryFromIntError; + + fn try_from(value: CodeModeSessionCellExecutionLimits) -> Result { + Ok(Self { + max_yield_time_ms: value.max_yield_time_ms, + max_heap_size_bytes: value.max_heap_size_bytes.map(u64::try_from).transpose()?, + }) + } +} + +impl TryFrom for CodeModeSessionCellExecutionLimits { + type Error = TryFromIntError; + + fn try_from(value: WireSessionCellExecutionLimits) -> Result { + Ok(Self { + max_yield_time_ms: value.max_yield_time_ms, + max_heap_size_bytes: value.max_heap_size_bytes.map(usize::try_from).transpose()?, + }) + } +} + +/// A cell identifier with a wire representation owned by protocol V1. +#[derive(Clone, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)] +#[serde(transparent)] +pub struct WireCellId(String); + +impl WireCellId { + pub fn new(value: impl Into) -> Self { + Self(value.into()) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl From for WireCellId { + fn from(value: CellId) -> Self { + Self(value.as_str().to_string()) + } +} + +impl From<&CellId> for WireCellId { + fn from(value: &CellId) -> Self { + Self(value.as_str().to_string()) + } +} + +impl From for CellId { + fn from(value: WireCellId) -> Self { + Self::new(value.0) + } +} + +/// The V1 wire representation of a tool's stable name. +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct WireToolName { + pub name: String, + pub namespace: Option, +} + +impl From for WireToolName { + fn from(value: ToolName) -> Self { + Self { + name: value.name, + namespace: value.namespace, + } + } +} + +impl From for ToolName { + fn from(value: WireToolName) -> Self { + Self::new(value.namespace, value.name) + } +} + +/// The tool invocation shape supported by protocol V1. +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum WireToolKind { + Function, + Freeform, +} + +impl From for WireToolKind { + fn from(value: CodeModeToolKind) -> Self { + match value { + CodeModeToolKind::Function => Self::Function, + CodeModeToolKind::Freeform => Self::Freeform, + } + } +} + +impl From for CodeModeToolKind { + fn from(value: WireToolKind) -> Self { + match value { + WireToolKind::Function => Self::Function, + WireToolKind::Freeform => Self::Freeform, + } + } +} + +/// A V1 tool definition embedded in an execute request. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct WireToolDefinition { + pub name: String, + pub tool_name: WireToolName, + pub description: String, + pub kind: WireToolKind, + pub input_schema: Option, + pub output_schema: Option, +} + +impl From for WireToolDefinition { + fn from(value: ToolDefinition) -> Self { + Self { + name: value.name, + tool_name: value.tool_name.into(), + description: value.description, + kind: value.kind.into(), + input_schema: value.input_schema, + output_schema: value.output_schema, + } + } +} + +impl From for ToolDefinition { + fn from(value: WireToolDefinition) -> Self { + Self { + name: value.name, + tool_name: value.tool_name.into(), + description: value.description, + kind: value.kind.into(), + input_schema: value.input_schema, + output_schema: value.output_schema, + } + } +} + +/// The complete execute request shape supported by protocol V1. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct WireExecuteRequest { + pub tool_call_id: String, + pub enabled_tools: Vec, + pub source: String, + pub yield_time_ms: Option, + pub max_output_tokens: Option, +} + +impl TryFrom for WireExecuteRequest { + type Error = TryFromIntError; + + fn try_from(value: ExecuteRequest) -> Result { + Ok(Self { + tool_call_id: value.tool_call_id, + enabled_tools: value.enabled_tools.into_iter().map(Into::into).collect(), + source: value.source, + yield_time_ms: value.yield_time_ms, + max_output_tokens: value.max_output_tokens.map(i32::try_from).transpose()?, + }) + } +} + +impl TryFrom for ExecuteRequest { + type Error = TryFromIntError; + + fn try_from(value: WireExecuteRequest) -> Result { + Ok(Self { + tool_call_id: value.tool_call_id, + enabled_tools: value.enabled_tools.into_iter().map(Into::into).collect(), + source: value.source, + yield_time_ms: value.yield_time_ms, + max_output_tokens: value.max_output_tokens.map(usize::try_from).transpose()?, + }) + } +} + +/// The complete wait request shape supported by protocol V1. +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct WireWaitRequest { + pub cell_id: WireCellId, + pub yield_time_ms: u64, +} + +impl From for WireWaitRequest { + fn from(value: WaitRequest) -> Self { + Self { + cell_id: value.cell_id.into(), + yield_time_ms: value.yield_time_ms, + } + } +} + +impl From for WaitRequest { + fn from(value: WireWaitRequest) -> Self { + Self { + cell_id: value.cell_id.into(), + yield_time_ms: value.yield_time_ms, + } + } +} + +/// Image detail values accepted in a V1 runtime response. +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(rename_all = "lowercase")] +pub enum WireImageDetail { + Auto, + Low, + High, + Original, +} + +impl From for WireImageDetail { + fn from(value: ImageDetail) -> Self { + match value { + ImageDetail::Auto => Self::Auto, + ImageDetail::Low => Self::Low, + ImageDetail::High => Self::High, + ImageDetail::Original => Self::Original, + } + } +} + +impl From for ImageDetail { + fn from(value: WireImageDetail) -> Self { + match value { + WireImageDetail::Auto => Self::Auto, + WireImageDetail::Low => Self::Low, + WireImageDetail::High => Self::High, + WireImageDetail::Original => Self::Original, + } + } +} + +/// One output item emitted by a V1 runtime response. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields, tag = "type", rename_all = "snake_case")] +pub enum WireContentItem { + InputText { + text: String, + }, + InputImage { + image_url: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + detail: Option, + }, + InputAudio { + audio_url: String, + }, +} + +impl From for WireContentItem { + fn from(value: FunctionCallOutputContentItem) -> Self { + match value { + FunctionCallOutputContentItem::InputText { text } => Self::InputText { text }, + FunctionCallOutputContentItem::InputImage { image_url, detail } => Self::InputImage { + image_url, + detail: detail.map(Into::into), + }, + FunctionCallOutputContentItem::InputAudio { audio_url } => { + Self::InputAudio { audio_url } + } + } + } +} + +impl From for FunctionCallOutputContentItem { + fn from(value: WireContentItem) -> Self { + match value { + WireContentItem::InputText { text } => Self::InputText { text }, + WireContentItem::InputImage { image_url, detail } => Self::InputImage { + image_url, + detail: detail.map(Into::into), + }, + WireContentItem::InputAudio { audio_url } => Self::InputAudio { audio_url }, + } + } +} + +/// Runtime output returned over the V1 host connection. +/// +/// Host time is required and covers this request, not the cell lifetime. The +/// app-server and host run at the same version; no negotiation is needed. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub enum WireRuntimeResponse { + Yielded { + cell_id: WireCellId, + content_items: Vec, + code_mode_host_duration_ns: u64, + }, + Terminated { + cell_id: WireCellId, + content_items: Vec, + code_mode_host_duration_ns: u64, + }, + Result { + cell_id: WireCellId, + content_items: Vec, + error_text: Option, + code_mode_host_duration_ns: u64, + }, +} + +impl TryFrom for WireRuntimeResponse { + type Error = MissingCodeModeHostDuration; + + /// Preserves the response's timing; the host handler must record it first. + fn try_from(value: RuntimeResponse) -> Result { + Ok(match value { + RuntimeResponse::Yielded { + cell_id, + content_items, + code_mode_host_duration, + } => { + let code_mode_host_duration = + code_mode_host_duration.ok_or(MissingCodeModeHostDuration)?; + Self::Yielded { + cell_id: cell_id.into(), + content_items: content_items.into_iter().map(Into::into).collect(), + code_mode_host_duration_ns: u64::try_from(code_mode_host_duration.as_nanos()) + .unwrap_or(u64::MAX), + } + } + RuntimeResponse::Terminated { + cell_id, + content_items, + code_mode_host_duration, + } => { + let code_mode_host_duration = + code_mode_host_duration.ok_or(MissingCodeModeHostDuration)?; + Self::Terminated { + cell_id: cell_id.into(), + content_items: content_items.into_iter().map(Into::into).collect(), + code_mode_host_duration_ns: u64::try_from(code_mode_host_duration.as_nanos()) + .unwrap_or(u64::MAX), + } + } + RuntimeResponse::Result { + cell_id, + content_items, + error_text, + code_mode_host_duration, + } => { + let code_mode_host_duration = + code_mode_host_duration.ok_or(MissingCodeModeHostDuration)?; + Self::Result { + cell_id: cell_id.into(), + content_items: content_items.into_iter().map(Into::into).collect(), + error_text, + code_mode_host_duration_ns: u64::try_from(code_mode_host_duration.as_nanos()) + .unwrap_or(u64::MAX), + } + } + }) + } +} + +impl From for RuntimeResponse { + fn from(value: WireRuntimeResponse) -> Self { + match value { + WireRuntimeResponse::Yielded { + cell_id, + content_items, + code_mode_host_duration_ns, + } => Self::Yielded { + cell_id: cell_id.into(), + content_items: content_items.into_iter().map(Into::into).collect(), + code_mode_host_duration: Some(Duration::from_nanos(code_mode_host_duration_ns)), + }, + WireRuntimeResponse::Terminated { + cell_id, + content_items, + code_mode_host_duration_ns, + } => Self::Terminated { + cell_id: cell_id.into(), + content_items: content_items.into_iter().map(Into::into).collect(), + code_mode_host_duration: Some(Duration::from_nanos(code_mode_host_duration_ns)), + }, + WireRuntimeResponse::Result { + cell_id, + content_items, + error_text, + code_mode_host_duration_ns, + } => Self::Result { + cell_id: cell_id.into(), + content_items: content_items.into_iter().map(Into::into).collect(), + error_text, + code_mode_host_duration: Some(Duration::from_nanos(code_mode_host_duration_ns)), + }, + } + } +} + +/// Whether a waited-for cell remained live in protocol V1. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub enum WireWaitOutcome { + LiveCell(WireRuntimeResponse), + MissingCell(WireRuntimeResponse), +} + +impl TryFrom for WireWaitOutcome { + type Error = MissingCodeModeHostDuration; + + fn try_from(value: WaitOutcome) -> Result { + Ok(match value { + WaitOutcome::LiveCell(response) => Self::LiveCell(response.try_into()?), + WaitOutcome::MissingCell(response) => Self::MissingCell(response.try_into()?), + }) + } +} + +impl From for WaitOutcome { + fn from(value: WireWaitOutcome) -> Self { + match value { + WireWaitOutcome::LiveCell(response) => Self::LiveCell(response.into()), + WireWaitOutcome::MissingCell(response) => Self::MissingCell(response.into()), + } + } +} + +/// A nested tool invocation sent over the V1 host connection. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct WireNestedToolCall { + pub cell_id: WireCellId, + pub runtime_tool_call_id: String, + pub tool_name: WireToolName, + pub tool_kind: WireToolKind, + pub input: Option, +} + +impl From for WireNestedToolCall { + fn from(value: CodeModeNestedToolCall) -> Self { + Self { + cell_id: value.cell_id.into(), + runtime_tool_call_id: value.runtime_tool_call_id, + tool_name: value.tool_name.into(), + tool_kind: value.tool_kind.into(), + input: value.input, + } + } +} + +impl From for CodeModeNestedToolCall { + fn from(value: WireNestedToolCall) -> Self { + Self { + cell_id: value.cell_id.into(), + runtime_tool_call_id: value.runtime_tool_call_id, + tool_name: value.tool_name.into(), + tool_kind: value.tool_kind.into(), + input: value.input, + } + } +} diff --git a/codex-rs/code-mode-protocol/src/host/types.rs b/codex-rs/code-mode-protocol/src/host/types.rs new file mode 100644 index 0000000000000000000000000000000000000000..40c69df90b90a622b166f38bdc4806695f7992f8 --- /dev/null +++ b/codex-rs/code-mode-protocol/src/host/types.rs @@ -0,0 +1,248 @@ +use std::collections::BTreeSet; +use std::fmt; +use std::num::NonZeroU32; + +use serde::Deserialize; +use serde::Deserializer; +use serde::Serialize; +use serde::Serializer; +use serde::de::Error as _; + +/// Correlates one client operation request with the host's response. +#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)] +#[serde(transparent)] +pub struct RequestId(i64); + +impl RequestId { + pub const fn new(value: i64) -> Self { + Self(value) + } +} + +/// Correlates one host delegate request with the client's response. +#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)] +#[serde(transparent)] +pub struct DelegateRequestId(i64); + +impl DelegateRequestId { + pub const fn new(value: i64) -> Self { + Self(value) + } +} + +#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)] +#[serde(transparent)] +pub struct ProtocolVersion(NonZeroU32); + +impl ProtocolVersion { + pub const V1: Self = Self(NonZeroU32::MIN); + + pub const fn new(value: u32) -> Option { + match NonZeroU32::new(value) { + Some(value) => Some(Self(value)), + None => None, + } + } + + pub const fn get(self) -> u32 { + self.0.get() + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct InvalidIdentifier; + +impl fmt::Display for InvalidIdentifier { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("identifier must not be empty") + } +} + +impl std::error::Error for InvalidIdentifier {} + +#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +struct NonEmptyString(String); + +impl NonEmptyString { + fn new(value: impl Into) -> Result { + let value = value.into(); + if value.trim().is_empty() { + Err(InvalidIdentifier) + } else { + Ok(Self(value)) + } + } + + fn as_str(&self) -> &str { + &self.0 + } +} + +impl Serialize for NonEmptyString { + fn serialize(&self, serializer: S) -> Result + where + S: Serializer, + { + self.0.serialize(serializer) + } +} + +impl<'de> Deserialize<'de> for NonEmptyString { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + Self::new(String::deserialize(deserializer)?).map_err(D::Error::custom) + } +} + +/// A named protocol feature advertised during connection negotiation. +#[derive(Clone, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)] +#[serde(transparent)] +pub struct Capability(NonEmptyString); + +impl Capability { + pub fn new(value: impl Into) -> Result { + NonEmptyString::new(value).map(Self) + } + + pub fn as_str(&self) -> &str { + self.0.as_str() + } +} + +impl fmt::Display for Capability { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str(self.as_str()) + } +} + +/// Identifies one logical code-mode session on a connection. +#[derive(Clone, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)] +#[serde(transparent)] +pub struct SessionId(NonEmptyString); + +impl SessionId { + pub fn new(value: impl Into) -> Result { + NonEmptyString::new(value).map(Self) + } + + pub fn as_str(&self) -> &str { + self.0.as_str() + } +} + +impl fmt::Display for SessionId { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str(self.as_str()) + } +} + +#[derive(Clone, Debug, Default, Eq, PartialEq, Serialize)] +#[serde(transparent)] +pub struct CapabilitySet(BTreeSet); + +impl CapabilitySet { + pub fn empty() -> Self { + Self::default() + } + + pub fn try_new( + capabilities: impl IntoIterator, + ) -> Result { + let mut unique = BTreeSet::new(); + for capability in capabilities { + if !unique.insert(capability.clone()) { + return Err(DuplicateCapability { capability }); + } + } + Ok(Self(unique)) + } + + pub fn contains(&self, capability: &Capability) -> bool { + self.0.contains(capability) + } + + pub fn iter(&self) -> impl Iterator { + self.0.iter() + } +} + +impl<'de> Deserialize<'de> for CapabilitySet { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + Self::try_new(Vec::::deserialize(deserializer)?).map_err(D::Error::custom) + } +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct DuplicateCapability { + capability: Capability, +} + +impl fmt::Display for DuplicateCapability { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(formatter, "duplicate capability `{}`", self.capability) + } +} + +impl std::error::Error for DuplicateCapability {} + +#[derive(Clone, Debug, Eq, PartialEq, Serialize)] +#[serde(transparent)] +pub struct SupportedProtocolVersions(BTreeSet); + +impl SupportedProtocolVersions { + pub fn try_new( + versions: impl IntoIterator, + ) -> Result { + let mut unique = BTreeSet::new(); + for version in versions { + if !unique.insert(version) { + return Err(InvalidSupportedProtocolVersions::Duplicate(version)); + } + } + if unique.is_empty() { + return Err(InvalidSupportedProtocolVersions::Empty); + } + Ok(Self(unique)) + } + + pub fn contains(&self, version: ProtocolVersion) -> bool { + self.0.contains(&version) + } + + pub fn iter(&self) -> impl Iterator + '_ { + self.0.iter().copied() + } +} + +impl<'de> Deserialize<'de> for SupportedProtocolVersions { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + Self::try_new(Vec::::deserialize(deserializer)?).map_err(D::Error::custom) + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum InvalidSupportedProtocolVersions { + Empty, + Duplicate(ProtocolVersion), +} + +impl fmt::Display for InvalidSupportedProtocolVersions { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Empty => formatter.write_str("at least one protocol version is required"), + Self::Duplicate(version) => { + write!(formatter, "duplicate protocol version {}", version.get()) + } + } + } +} + +impl std::error::Error for InvalidSupportedProtocolVersions {} diff --git a/codex-rs/code-mode-protocol/src/json_schema_types.rs b/codex-rs/code-mode-protocol/src/json_schema_types.rs new file mode 100644 index 0000000000000000000000000000000000000000..0c519528a4b6896097906393123ef410beccc3d1 --- /dev/null +++ b/codex-rs/code-mode-protocol/src/json_schema_types.rs @@ -0,0 +1,538 @@ +use serde_json::Value as JsonValue; +use std::collections::BTreeMap; + +use crate::description::normalize_code_mode_identifier; + +// Expose one nested recursive shape, then fall back to `unknown` on the next +// occurrence so generated tool declarations remain finite. +const MAX_LOCAL_REF_EXPANSIONS_PER_PATH: usize = 2; +// Bound repeated refs and DAG fan-out separately from cycle depth so one +// compact schema cannot expand into an arbitrarily large model-visible item. +const MAX_TOTAL_LOCAL_REF_EXPANSIONS: usize = 32; +// Bound individual schema rendering before assembling the final declaration. +const MAX_RENDERED_SCHEMA_BYTES: usize = 16_000; +// Charge intermediate render strings as they are built so repeated local refs +// cannot allocate unbounded expanded copies before the final schema cap runs. +const MAX_RENDER_WORK_BYTES: usize = MAX_RENDERED_SCHEMA_BYTES * 4; + +pub fn render_json_schema_to_typescript(schema: &JsonValue) -> String { + let rendered = JsonSchemaTypeRenderer::new(schema).render(schema); + if rendered.len() > MAX_RENDERED_SCHEMA_BYTES { + "unknown".to_string() + } else { + rendered + } +} + +struct JsonSchemaTypeRenderer<'a> { + root: &'a JsonValue, + nested_schema_resource_depth: usize, + active_local_ref_expansions: BTreeMap, + remaining_local_ref_expansions: usize, + remaining_render_work_bytes: usize, + render_work_budget_exhausted: bool, +} + +impl<'a> JsonSchemaTypeRenderer<'a> { + fn new(root: &'a JsonValue) -> Self { + Self { + root, + nested_schema_resource_depth: 0, + active_local_ref_expansions: BTreeMap::new(), + remaining_local_ref_expansions: MAX_TOTAL_LOCAL_REF_EXPANSIONS, + remaining_render_work_bytes: MAX_RENDER_WORK_BYTES, + render_work_budget_exhausted: false, + } + } + + fn render(&mut self, schema: &JsonValue) -> String { + if self.render_work_budget_exhausted { + return "unknown".to_string(); + } + + // A nested `$id` starts a new schema resource. Fragment-only refs below + // it are scoped to that resource, not to the outer document root. + let enters_nested_schema_resource = !std::ptr::eq(schema, self.root) + && schema + .as_object() + .is_some_and(|map| map.contains_key("$id")); + if enters_nested_schema_resource { + self.nested_schema_resource_depth += 1; + } + + let rendered = match schema { + JsonValue::Bool(true) => "unknown".to_string(), + JsonValue::Bool(false) => "never".to_string(), + JsonValue::Object(map) => self.render_map(map), + _ => "unknown".to_string(), + }; + if enters_nested_schema_resource { + self.nested_schema_resource_depth -= 1; + } + self.finish_render(rendered) + } + + fn render_map(&mut self, map: &serde_json::Map) -> String { + if self.render_work_budget_exhausted { + return "unknown".to_string(); + } + + if map.contains_key("$ref") { + return self.render_ref(map); + } + + if let Some(value) = map.get("const") { + return self.render_literal(value); + } + + if let Some(values) = map.get("enum").and_then(JsonValue::as_array) { + let mut rendered = Vec::new(); + for value in values { + let literal = self.render_literal(value); + if self.render_work_budget_exhausted { + return "unknown".to_string(); + } + if !self.consume_render_work(literal.len()) { + return "unknown".to_string(); + } + rendered.push(literal); + } + if !rendered.is_empty() { + return rendered.join(" | "); + } + } + + for key in ["anyOf", "oneOf"] { + if let Some(variants) = map.get(key).and_then(JsonValue::as_array) { + let mut rendered = Vec::new(); + for variant in variants { + if self.render_work_budget_exhausted { + return "unknown".to_string(); + } + rendered.push(self.render(variant)); + } + if !rendered.is_empty() { + return rendered.join(" | "); + } + } + } + + if let Some(variants) = map.get("allOf").and_then(JsonValue::as_array) { + let mut rendered = Vec::new(); + for variant in variants { + if self.render_work_budget_exhausted { + return "unknown".to_string(); + } + rendered.push(parenthesize_union_for_intersection(self.render(variant))); + } + if !rendered.is_empty() { + return rendered.join(" & "); + } + } + + if let Some(schema_type) = map.get("type") { + if let Some(types) = schema_type.as_array() { + let mut rendered = Vec::new(); + for schema_type in types.iter().filter_map(JsonValue::as_str) { + if self.render_work_budget_exhausted { + return "unknown".to_string(); + } + rendered.push(self.render_type_keyword(map, schema_type)); + } + if !rendered.is_empty() { + return rendered.join(" | "); + } + } + + if let Some(schema_type) = schema_type.as_str() { + return self.render_type_keyword(map, schema_type); + } + } + + if map.contains_key("properties") + || map.contains_key("additionalProperties") + || map.contains_key("required") + { + return self.render_object(map); + } + + if map.contains_key("items") || map.contains_key("prefixItems") { + return self.render_array(map); + } + + "unknown".to_string() + } + + fn render_ref(&mut self, map: &serde_json::Map) -> String { + let referenced_type = if self.nested_schema_resource_depth > 0 { + None + } else { + map.get("$ref") + .and_then(JsonValue::as_str) + .and_then(local_json_pointer) + .and_then(|pointer| { + let active_expansions = self + .active_local_ref_expansions + .get(&pointer) + .copied() + .unwrap_or_default(); + if active_expansions >= MAX_LOCAL_REF_EXPANSIONS_PER_PATH { + return Some("unknown".to_string()); + } + if self.remaining_local_ref_expansions == 0 { + return Some("unknown".to_string()); + } + + let root = self.root; + let target = if pointer.is_empty() { + Some(root) + } else { + root.pointer(&pointer) + }?; + self.remaining_local_ref_expansions -= 1; + self.active_local_ref_expansions + .insert(pointer.clone(), active_expansions + 1); + + let rendered = self.render(target); + if active_expansions == 0 { + self.active_local_ref_expansions.remove(&pointer); + } else { + self.active_local_ref_expansions + .insert(pointer, active_expansions); + } + Some(rendered) + }) + } + .unwrap_or_else(|| "unknown".to_string()); + if self.render_work_budget_exhausted { + return "unknown".to_string(); + } + + let siblings = map + .iter() + .filter(|(key, _)| !matches!(key.as_str(), "$ref" | "$defs" | "definitions")) + .map(|(key, value)| (key.clone(), value.clone())) + .collect(); + if !has_renderable_schema_keywords(&siblings) { + return referenced_type; + } + + let sibling_type = self.render_map(&siblings); + match (referenced_type.as_str(), sibling_type.as_str()) { + ("unknown", _) => sibling_type, + (_, "unknown") => referenced_type, + _ => format!("({referenced_type}) & ({sibling_type})"), + } + } + + fn render_type_keyword( + &mut self, + map: &serde_json::Map, + schema_type: &str, + ) -> String { + match schema_type { + "string" => "string".to_string(), + "number" | "integer" => "number".to_string(), + "boolean" => "boolean".to_string(), + "null" => "null".to_string(), + "array" => self.render_array(map), + "object" => self.render_object(map), + _ => "unknown".to_string(), + } + } + + fn render_array(&mut self, map: &serde_json::Map) -> String { + if let Some(items) = map.get("items") { + let item_type = self.render(items); + if self.render_work_budget_exhausted { + return "unknown".to_string(); + } + return format!("Array<{item_type}>"); + } + + if let Some(items) = map.get("prefixItems").and_then(JsonValue::as_array) { + let mut item_types = Vec::new(); + for item in items { + if self.render_work_budget_exhausted { + return "unknown".to_string(); + } + item_types.push(self.render(item)); + } + if !item_types.is_empty() { + return format!("[{}]", item_types.join(", ")); + } + } + + "unknown[]".to_string() + } + + fn append_additional_properties_line( + &mut self, + lines: &mut Vec, + map: &serde_json::Map, + properties: &serde_json::Map, + line_prefix: &str, + ) -> bool { + if let Some(additional_properties) = map.get("additionalProperties") { + let property_type = match additional_properties { + JsonValue::Bool(true) => Some("unknown".to_string()), + JsonValue::Bool(false) => None, + value => Some(self.render(value)), + }; + + if let Some(property_type) = property_type { + return self.push_render_line( + lines, + format!("{line_prefix}[key: string]: {property_type};"), + ); + } + } else if properties.is_empty() { + return self.push_render_line(lines, format!("{line_prefix}[key: string]: unknown;")); + } + true + } + + fn render_object_property( + &mut self, + name: &str, + value: &JsonValue, + required: &[&str], + ) -> String { + if name.len() > self.remaining_render_work_bytes { + self.render_work_budget_exhausted = true; + return "unknown".to_string(); + } + let optional = if required.iter().any(|required_name| required_name == &name) { + "" + } else { + "?" + }; + let property_name = render_json_schema_property_name(name); + let property_type = self.render(value); + if self.render_work_budget_exhausted { + return "unknown".to_string(); + } + format!("{property_name}{optional}: {property_type};") + } + + fn render_object(&mut self, map: &serde_json::Map) -> String { + let required = map + .get("required") + .and_then(JsonValue::as_array) + .map(|items| { + items + .iter() + .filter_map(JsonValue::as_str) + .collect::>() + }) + .unwrap_or_default(); + let empty_properties = serde_json::Map::new(); + let properties = map + .get("properties") + .and_then(JsonValue::as_object) + .unwrap_or(&empty_properties); + + let mut sorted_properties = properties.iter().collect::>(); + sorted_properties.sort_unstable_by_key(|(name_a, _)| *name_a); + if sorted_properties + .iter() + .any(|(_, value)| has_property_description(value)) + { + let mut lines = Vec::new(); + if !self.push_render_line(&mut lines, "{".to_string()) { + return "unknown".to_string(); + } + for (name, value) in sorted_properties { + if let Some(description) = value.get("description").and_then(JsonValue::as_str) { + for description_line in description + .lines() + .map(str::trim) + .filter(|line| !line.is_empty()) + { + if description_line.len().saturating_add(5) + > self.remaining_render_work_bytes + || !self + .push_render_line(&mut lines, format!(" // {description_line}")) + { + return "unknown".to_string(); + } + } + } + + let property = self.render_object_property(name, value, &required); + if self.render_work_budget_exhausted + || !self.push_render_line(&mut lines, format!(" {property}")) + { + return "unknown".to_string(); + } + } + + if !self.append_additional_properties_line(&mut lines, map, properties, " ") + || !self.push_render_line(&mut lines, "}".to_string()) + { + return "unknown".to_string(); + } + return lines.join("\n"); + } + + let mut lines = Vec::new(); + for (name, value) in sorted_properties { + let property = self.render_object_property(name, value, &required); + if self.render_work_budget_exhausted || !self.push_render_line(&mut lines, property) { + return "unknown".to_string(); + } + } + + if !self.append_additional_properties_line(&mut lines, map, properties, "") { + return "unknown".to_string(); + } + + if lines.is_empty() { + return "{}".to_string(); + } + + format!("{{ {} }}", lines.join(" ")) + } + + fn finish_render(&mut self, rendered: String) -> String { + if self.consume_render_work(rendered.len()) { + rendered + } else { + "unknown".to_string() + } + } + + fn render_literal(&mut self, value: &JsonValue) -> String { + if json_literal_serialization_upper_bound(value) > self.remaining_render_work_bytes { + self.render_work_budget_exhausted = true; + "unknown".to_string() + } else { + render_json_schema_literal(value) + } + } + + fn consume_render_work(&mut self, rendered_bytes: usize) -> bool { + if rendered_bytes > self.remaining_render_work_bytes { + self.render_work_budget_exhausted = true; + false + } else { + self.remaining_render_work_bytes -= rendered_bytes; + true + } + } + + fn push_render_line(&mut self, lines: &mut Vec, line: String) -> bool { + if !self.consume_render_work(line.len()) { + return false; + } + lines.push(line); + true + } +} + +fn parenthesize_union_for_intersection(rendered: String) -> String { + if rendered.contains(" | ") { + format!("({rendered})") + } else { + rendered + } +} + +fn local_json_pointer(reference: &str) -> Option { + let fragment = reference.strip_prefix('#')?; + let pointer = percent_decode_uri_fragment(fragment)?; + if pointer.is_empty() || pointer.starts_with('/') { + Some(pointer) + } else { + None + } +} + +fn percent_decode_uri_fragment(fragment: &str) -> Option { + let bytes = fragment.as_bytes(); + let mut decoded = Vec::with_capacity(bytes.len()); + let mut index = 0; + while index < bytes.len() { + if bytes[index] == b'%' { + let high = decode_hex_digit(*bytes.get(index + 1)?)?; + let low = decode_hex_digit(*bytes.get(index + 2)?)?; + decoded.push((high << 4) | low); + index += 3; + } else { + decoded.push(bytes[index]); + index += 1; + } + } + String::from_utf8(decoded).ok() +} + +fn decode_hex_digit(digit: u8) -> Option { + match digit { + b'0'..=b'9' => Some(digit - b'0'), + b'a'..=b'f' => Some(digit - b'a' + 10), + b'A'..=b'F' => Some(digit - b'A' + 10), + _ => None, + } +} + +fn has_renderable_schema_keywords(map: &serde_json::Map) -> bool { + [ + "const", + "enum", + "anyOf", + "oneOf", + "allOf", + "type", + "properties", + "additionalProperties", + "required", + "items", + "prefixItems", + ] + .iter() + .any(|key| map.contains_key(*key)) +} + +fn has_property_description(value: &JsonValue) -> bool { + value + .get("description") + .and_then(JsonValue::as_str) + .is_some_and(|description| !description.is_empty()) +} + +fn render_json_schema_property_name(name: &str) -> String { + if normalize_code_mode_identifier(name) == name { + name.to_string() + } else { + serde_json::to_string(name).unwrap_or_else(|_| format!("\"{}\"", name.replace('"', "\\\""))) + } +} + +fn render_json_schema_literal(value: &JsonValue) -> String { + serde_json::to_string(value).unwrap_or_else(|_| "unknown".to_string()) +} + +fn json_literal_serialization_upper_bound(value: &JsonValue) -> usize { + match value { + JsonValue::Null => 4, + JsonValue::Bool(false) => 5, + JsonValue::Bool(true) => 4, + JsonValue::Number(number) => number.to_string().len(), + // JSON escaping can expand one UTF-8 byte to at most one six-byte + // Unicode escape, so this bounds allocation before serialization. + JsonValue::String(string) => string.len().saturating_mul(6).saturating_add(2), + JsonValue::Array(values) => values.iter().fold(2, |size, value| { + size.saturating_add(1) + .saturating_add(json_literal_serialization_upper_bound(value)) + }), + JsonValue::Object(map) => map.iter().fold(2, |size, (key, value)| { + size.saturating_add(4) + .saturating_add(key.len().saturating_mul(6)) + .saturating_add(json_literal_serialization_upper_bound(value)) + }), + } +} + +#[cfg(test)] +#[path = "json_schema_types_tests.rs"] +mod tests; diff --git a/codex-rs/code-mode-protocol/src/json_schema_types_tests.rs b/codex-rs/code-mode-protocol/src/json_schema_types_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..449865f60e3d1ae77a2fc2ee66fe565402f878de --- /dev/null +++ b/codex-rs/code-mode-protocol/src/json_schema_types_tests.rs @@ -0,0 +1,200 @@ +use super::*; +use pretty_assertions::assert_eq; +use serde_json::json; + +#[test] +fn renders_recursive_local_refs_with_escaped_pointer_segments() { + let schema = json!({ + "type": "object", + "properties": { + "clauses": { + "type": "array", + "items": { "$ref": "#/$defs/Boolean~1Clause~0v1" } + } + }, + "$defs": { + "Boolean/Clause~v1": { + "type": "object", + "properties": { + "query": { "$ref": "#/$defs/Query" } + } + }, + "Query": { + "oneOf": [ + { "type": "string" }, + { + "type": "object", + "properties": { + "clauses": { + "type": "array", + "items": { "$ref": "#/$defs/Boolean~1Clause~0v1" } + } + } + } + ] + } + } + }); + + let rendered = render_json_schema_to_typescript(&schema); + assert!(rendered.contains("clauses?: Array<{ query?: string | { clauses?: Array<{")); + assert!(rendered.contains("query?: string | { clauses?: Array; };")); +} + +#[test] +fn renders_ref_siblings_uri_fragments_and_all_of_precedence() { + assert_eq!( + render_json_schema_to_typescript(&json!({ + "$ref": "#/$defs/Label", + "enum": ["A"], + "$defs": { "Label": { "type": "string" } } + })), + r#"(string) & ("A")"# + ); + assert_eq!( + render_json_schema_to_typescript(&json!({ + "$ref": "#/$defs/Foo%20Bar", + "$defs": { "Foo Bar": { "type": "string" } } + })), + "string" + ); + assert_eq!( + render_json_schema_to_typescript(&json!({ + "allOf": [ + { "$ref": "#/$defs/Choice" }, + { "type": "object", "properties": { "value": { "type": "string" } } } + ], + "$defs": { + "Choice": { "oneOf": [{ "type": "string" }, { "type": "number" }] } + } + })), + "(string | number) & { value?: string; }" + ); +} + +#[test] +fn leaves_local_refs_under_nested_schema_resources_unresolved() { + let schema = json!({ + "$defs": { + "Choice": { "type": "string" } + }, + "type": "object", + "properties": { + "nested": { + "$id": "urn:nested", + "$defs": { + "Choice": { "type": "number" } + }, + "type": "object", + "properties": { + "value": { "$ref": "#/$defs/Choice" } + } + } + } + }); + + assert_eq!( + render_json_schema_to_typescript(&schema), + "{ nested?: { value?: unknown; }; }" + ); +} + +#[test] +fn bounds_expansions_without_charging_dangling_refs() { + let mut properties = (0..MAX_TOTAL_LOCAL_REF_EXPANSIONS) + .map(|index| { + ( + format!("a_missing_{index}"), + json!({ "$ref": format!("#/$defs/Missing{index}") }), + ) + }) + .collect::>(); + properties.insert("z_valid".to_string(), json!({ "$ref": "#/$defs/Valid" })); + let schema = json!({ + "type": "object", + "properties": properties, + "$defs": { "Valid": { "type": "string" } } + }); + + let rendered = render_json_schema_to_typescript(&schema); + assert!(rendered.contains("z_valid?: string;")); + + let properties = (0..MAX_TOTAL_LOCAL_REF_EXPANSIONS + 2) + .map(|index| { + ( + format!("property_{index}"), + json!({ "$ref": "#/$defs/Item" }), + ) + }) + .collect::>(); + let rendered = render_json_schema_to_typescript(&json!({ + "type": "object", + "properties": properties, + "$defs": { "Item": { "type": "string" } } + })); + assert_eq!( + rendered.matches("string").count(), + MAX_TOTAL_LOCAL_REF_EXPANSIONS + ); + assert_eq!(rendered.matches("unknown").count(), 2); +} + +#[test] +fn repeated_large_ref_expansions_exhaust_render_work_budget() { + let properties = (0..MAX_TOTAL_LOCAL_REF_EXPANSIONS) + .map(|index| { + ( + format!("property_{index}"), + json!({ "$ref": "#/$defs/Item" }), + ) + }) + .collect::>(); + let schema = json!({ + "type": "object", + "properties": properties, + "$defs": { + "Item": { + "type": "object", + "properties": { + "value": { + "type": "string", + "description": "x".repeat(MAX_RENDERED_SCHEMA_BYTES / 2) + } + } + } + } + }); + + let mut renderer = JsonSchemaTypeRenderer::new(&schema); + assert_eq!(renderer.render(&schema), "unknown"); + assert!(renderer.render_work_budget_exhausted); +} + +#[test] +fn oversized_ref_literal_exhausts_render_work_budget() { + let schema = json!({ + "$ref": "#/$defs/Value", + "$defs": { + "Value": { + "const": "x".repeat(MAX_RENDER_WORK_BYTES) + } + } + }); + + let mut renderer = JsonSchemaTypeRenderer::new(&schema); + assert_eq!(renderer.render(&schema), "unknown"); + assert!(renderer.render_work_budget_exhausted); +} + +#[test] +fn rendered_schema_has_a_hard_size_cap() { + let description = "x".repeat(MAX_RENDERED_SCHEMA_BYTES); + let schema = json!({ + "type": "object", + "properties": { + "value": { "type": "string", "description": description } + } + }); + + assert_eq!(render_json_schema_to_typescript(&schema), "unknown"); +} diff --git a/codex-rs/code-mode-protocol/src/lib.rs b/codex-rs/code-mode-protocol/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..c253f5bcc789e05bbbadd7fcbcff9cb86d2e9670 --- /dev/null +++ b/codex-rs/code-mode-protocol/src/lib.rs @@ -0,0 +1,52 @@ +mod description; +pub mod grpc; +pub mod host; +mod json_schema_types; +mod response; +mod runtime; +mod session; + +pub use description::CODE_MODE_PRAGMA_PREFIX; +pub use description::CodeModeToolKind; +pub use description::EnabledToolMetadata; +pub use description::ImageDetailVisibility; +pub use description::ToolDefinition; +pub use description::ToolNamespaceDescription; +pub use description::augment_tool_definition; +pub use description::build_exec_tool_description; +pub use description::build_wait_tool_description; +pub use description::enabled_tool_metadata; +pub use description::is_code_mode_nested_tool; +pub use description::normalize_code_mode_identifier; +pub use description::parse_exec_source; +pub use description::render_code_mode_sample; +pub use json_schema_types::render_json_schema_to_typescript; +pub use response::DEFAULT_IMAGE_DETAIL; +pub use response::FunctionCallOutputContentItem; +pub use response::ImageDetail; +pub use runtime::CodeModeNestedToolCall; +pub use runtime::DEFAULT_EXEC_YIELD_TIME_MS; +pub use runtime::DEFAULT_MAX_OUTPUT_TOKENS_PER_EXEC_CALL; +pub use runtime::DEFAULT_WAIT_YIELD_TIME_MS; +pub use runtime::ExecuteRequest; +pub use runtime::ExecuteToPendingOutcome; +pub use runtime::MissingCodeModeHostDuration; +pub use runtime::RuntimeResponse; +pub use runtime::WaitOutcome; +pub use runtime::WaitRequest; +pub use runtime::WaitToPendingOutcome; +pub use runtime::WaitToPendingRequest; +pub use session::CellId; +pub use session::CodeModeSession; +pub use session::CodeModeSessionCellExecutionLimits; +pub use session::CodeModeSessionDelegate; +pub use session::CodeModeSessionProvider; +pub use session::CodeModeSessionProviderFuture; +pub use session::CodeModeSessionResultFuture; +pub use session::NoopCodeModeSessionDelegate; +pub use session::NotificationFuture; +pub use session::StartedCell; +pub use session::ToolInvocationFuture; + +pub const PUBLIC_TOOL_NAME: &str = "exec"; +pub const WAIT_TOOL_NAME: &str = "wait"; diff --git a/codex-rs/code-mode-protocol/src/response.rs b/codex-rs/code-mode-protocol/src/response.rs new file mode 100644 index 0000000000000000000000000000000000000000..9b45032d9461d6057831e2fa9648dddcfeb05edc --- /dev/null +++ b/codex-rs/code-mode-protocol/src/response.rs @@ -0,0 +1,29 @@ +use serde::Deserialize; +use serde::Serialize; + +#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "lowercase")] +pub enum ImageDetail { + Auto, + Low, + High, + Original, +} + +pub const DEFAULT_IMAGE_DETAIL: ImageDetail = ImageDetail::High; + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum FunctionCallOutputContentItem { + InputText { + text: String, + }, + InputImage { + image_url: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + detail: Option, + }, + InputAudio { + audio_url: String, + }, +} diff --git a/codex-rs/code-mode-protocol/src/runtime.rs b/codex-rs/code-mode-protocol/src/runtime.rs new file mode 100644 index 0000000000000000000000000000000000000000..afa96ebf29d348ff47869cb36a3893622494bf1a --- /dev/null +++ b/codex-rs/code-mode-protocol/src/runtime.rs @@ -0,0 +1,183 @@ +use std::error::Error; +use std::fmt; +use std::time::Duration; + +use codex_protocol::ToolName; +use serde::Deserialize; +use serde::Serialize; +use serde_json::Value as JsonValue; + +use crate::CellId; +use crate::CodeModeToolKind; +use crate::FunctionCallOutputContentItem; +use crate::ToolDefinition; + +pub const DEFAULT_EXEC_YIELD_TIME_MS: u64 = 10_000; +pub const DEFAULT_WAIT_YIELD_TIME_MS: u64 = 10_000; +pub const DEFAULT_MAX_OUTPUT_TOKENS_PER_EXEC_CALL: usize = 10_000; + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +pub struct ExecuteRequest { + pub tool_call_id: String, + pub enabled_tools: Vec, + pub source: String, + pub yield_time_ms: Option, + pub max_output_tokens: Option, +} + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +pub struct WaitRequest { + pub cell_id: CellId, + pub yield_time_ms: u64, +} + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +pub struct WaitToPendingRequest { + pub cell_id: CellId, +} + +#[derive(Debug, Deserialize, PartialEq, Serialize)] +pub enum WaitOutcome { + LiveCell(RuntimeResponse), + MissingCell(RuntimeResponse), +} + +impl WaitOutcome { + /// Returns timing for this wait or termination request, when supplied by its host. + pub fn code_mode_host_duration(&self) -> Option { + match self { + Self::LiveCell(response) | Self::MissingCell(response) => { + response.code_mode_host_duration() + } + } + } + + /// Records the enclosing host request's duration before wire conversion. + pub fn with_code_mode_host_duration(self, code_mode_host_duration: Duration) -> Self { + match self { + Self::LiveCell(response) => { + Self::LiveCell(response.with_code_mode_host_duration(code_mode_host_duration)) + } + Self::MissingCell(response) => { + Self::MissingCell(response.with_code_mode_host_duration(code_mode_host_duration)) + } + } + } +} + +#[derive(Debug, Deserialize, PartialEq, Serialize)] +pub enum ExecuteToPendingOutcome { + Pending { + cell_id: CellId, + content_items: Vec, + pending_tool_call_ids: Vec, + }, + Completed(RuntimeResponse), +} + +#[derive(Debug, Deserialize, PartialEq, Serialize)] +pub enum WaitToPendingOutcome { + LiveCell(ExecuteToPendingOutcome), + MissingCell(RuntimeResponse), +} + +impl From for RuntimeResponse { + fn from(outcome: WaitOutcome) -> Self { + match outcome { + WaitOutcome::LiveCell(response) | WaitOutcome::MissingCell(response) => response, + } + } +} + +/// Runtime output with optional timing for the host request that observed it. +/// +/// The JavaScript session returns untimed output. The host handler records its +/// complete request duration before conversion; wire conversions preserve that +/// field and reject untimed output. Decoded host responses always contain timing, +/// including measured zero. Raw traces also retain timing when present. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +pub enum RuntimeResponse { + Yielded { + cell_id: CellId, + content_items: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + code_mode_host_duration: Option, + }, + Terminated { + cell_id: CellId, + content_items: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + code_mode_host_duration: Option, + }, + Result { + cell_id: CellId, + content_items: Vec, + error_text: Option, + #[serde(skip_serializing_if = "Option::is_none")] + code_mode_host_duration: Option, + }, +} + +impl RuntimeResponse { + /// Returns timing for this observation, excluding background work between requests. + pub fn code_mode_host_duration(&self) -> Option { + match self { + Self::Yielded { + code_mode_host_duration, + .. + } + | Self::Terminated { + code_mode_host_duration, + .. + } + | Self::Result { + code_mode_host_duration, + .. + } => *code_mode_host_duration, + } + } + + /// Records the enclosing host request's duration before wire conversion. + pub fn with_code_mode_host_duration(mut self, code_mode_host_duration: Duration) -> Self { + match &mut self { + Self::Yielded { + code_mode_host_duration: value, + .. + } + | Self::Terminated { + code_mode_host_duration: value, + .. + } + | Self::Result { + code_mode_host_duration: value, + .. + } => *value = Some(code_mode_host_duration), + } + self + } +} + +/// An untimed runtime response cannot be encoded for delivery to a host client. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct MissingCodeModeHostDuration; + +impl fmt::Display for MissingCodeModeHostDuration { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("code-mode response is missing host duration") + } +} + +impl Error for MissingCodeModeHostDuration {} + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +pub struct CodeModeNestedToolCall { + pub cell_id: CellId, + pub runtime_tool_call_id: String, + pub tool_name: ToolName, + pub tool_kind: CodeModeToolKind, + pub input: Option, +} + +#[cfg(test)] +#[path = "runtime_tests.rs"] +mod tests; diff --git a/codex-rs/code-mode-protocol/src/runtime_tests.rs b/codex-rs/code-mode-protocol/src/runtime_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..9f20067b2c72ab6c7eb639b3199bb1503069b0a8 --- /dev/null +++ b/codex-rs/code-mode-protocol/src/runtime_tests.rs @@ -0,0 +1,107 @@ +//! Regression coverage for timing metadata at serialization boundaries. + +use std::time::Duration; + +use pretty_assertions::assert_eq; + +use super::MissingCodeModeHostDuration; +use super::RuntimeResponse; +use super::WaitOutcome; +use crate::CellId; +use crate::FunctionCallOutputContentItem; +use crate::host::WireRuntimeResponse; +use crate::host::WireWaitOutcome; + +/// Raw runtime serialization preserves absent timing; stdio always carries a +/// measured duration, including zero, without changing the response payload. +#[test] +fn code_mode_host_duration_survives_runtime_and_stdio_serialization() { + let content_items = vec![FunctionCallOutputContentItem::InputText { + text: "output".to_string(), + }]; + for response in [ + RuntimeResponse::Yielded { + cell_id: CellId::new("yielded-cell".to_string()), + content_items: content_items.clone(), + code_mode_host_duration: None, + }, + RuntimeResponse::Terminated { + cell_id: CellId::new("terminated-cell".to_string()), + content_items: content_items.clone(), + code_mode_host_duration: None, + }, + RuntimeResponse::Result { + cell_id: CellId::new("completed-cell".to_string()), + content_items, + error_text: Some("execution failed".to_string()), + code_mode_host_duration: None, + }, + ] { + for duration in [ + None, + Some(Duration::ZERO), + Some(Duration::from_nanos(/*nanos*/ 1_234_567_890)), + Some(Duration::from_nanos(u64::MAX)), + ] { + let mut expected = response.clone(); + match &mut expected { + RuntimeResponse::Yielded { + code_mode_host_duration, + .. + } + | RuntimeResponse::Terminated { + code_mode_host_duration, + .. + } + | RuntimeResponse::Result { + code_mode_host_duration, + .. + } => *code_mode_host_duration = duration, + } + + let payload = serde_json::to_value(&expected).expect("serialize response"); + assert_eq!( + serde_json::from_value::(payload).expect("deserialize response"), + expected + ); + + if duration.is_some() { + let wire_payload = serde_json::to_value( + WireRuntimeResponse::try_from(expected.clone()).expect("timed response"), + ) + .expect("serialize response over stdio"); + assert_eq!( + RuntimeResponse::from( + serde_json::from_value::(wire_payload) + .expect("deserialize stdio response") + ), + expected + ); + } + } + } +} + +/// Encoding must not turn a missing request measurement into measured zero, +/// including when no live cell remains to supply output. +#[test] +fn stdio_encoding_rejects_untimed_runtime_output() { + let response = RuntimeResponse::Terminated { + cell_id: CellId::new("cell".to_string()), + content_items: Vec::new(), + code_mode_host_duration: None, + }; + assert_eq!( + WireRuntimeResponse::try_from(response.clone()), + Err(MissingCodeModeHostDuration) + ); + for outcome in [ + WaitOutcome::LiveCell(response.clone()), + WaitOutcome::MissingCell(response), + ] { + assert_eq!( + WireWaitOutcome::try_from(outcome), + Err(MissingCodeModeHostDuration) + ); + } +} diff --git a/codex-rs/code-mode-protocol/src/session.rs b/codex-rs/code-mode-protocol/src/session.rs new file mode 100644 index 0000000000000000000000000000000000000000..36b88b916f494fc7b73bc72caf8fc4d1ddc95e35 --- /dev/null +++ b/codex-rs/code-mode-protocol/src/session.rs @@ -0,0 +1,200 @@ +use std::fmt; +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; + +use serde::Deserialize; +use serde::Serialize; +use serde_json::Value as JsonValue; +use tokio::sync::oneshot; +use tokio_util::sync::CancellationToken; + +use crate::CodeModeNestedToolCall; +use crate::ExecuteRequest; +use crate::RuntimeResponse; +use crate::WaitOutcome; +use crate::WaitRequest; + +pub type CodeModeSessionResultFuture<'a, T> = + Pin> + Send + 'a>>; +pub type CodeModeSessionProviderFuture<'a> = + CodeModeSessionResultFuture<'a, Arc>; +pub type ToolInvocationFuture<'a> = + Pin> + Send + 'a>>; +pub type NotificationFuture<'a> = Pin> + Send + 'a>>; + +/// Optional resource limits shared by every cell in one code-mode session. +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct CodeModeSessionCellExecutionLimits { + pub max_yield_time_ms: Option, + pub max_heap_size_bytes: Option, +} + +#[derive(Clone, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)] +pub struct CellId(String); + +impl CellId { + pub fn new(value: String) -> Self { + Self(value) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl AsRef for CellId { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl fmt::Display for CellId { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str(self.as_str()) + } +} + +pub struct StartedCell { + pub cell_id: CellId, + initial_response: CodeModeSessionResultFuture<'static, RuntimeResponse>, +} + +impl StartedCell { + pub fn new(cell_id: CellId, initial_response_rx: oneshot::Receiver) -> Self { + Self::from_future(cell_id, async move { + initial_response_rx + .await + .map_err(|_| "exec runtime ended unexpectedly".to_string()) + }) + } + + pub fn from_result_receiver( + cell_id: CellId, + initial_response_rx: oneshot::Receiver>, + ) -> Self { + Self::from_future(cell_id, async move { + initial_response_rx + .await + .map_err(|_| "exec runtime ended unexpectedly".to_string())? + }) + } + + pub fn from_future( + cell_id: CellId, + initial_response: impl Future> + Send + 'static, + ) -> Self { + Self { + cell_id, + initial_response: Box::pin(initial_response), + } + } + + pub async fn initial_response(self) -> Result { + self.initial_response.await + } +} + +/// Host callbacks owned by one code-mode execution. +/// +/// The session retains the supplied delegate while starting and running the cell, +/// including across yields, and releases it through its existing close/cancel paths. +pub trait CodeModeSessionDelegate: Send + Sync { + fn invoke_tool<'a>( + &'a self, + invocation: CodeModeNestedToolCall, + cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a>; + + fn notify<'a>( + &'a self, + call_id: String, + cell_id: CellId, + text: String, + cancellation_token: CancellationToken, + ) -> NotificationFuture<'a>; + + /// Releases delegate state associated with a cell after it reaches a terminal state. + fn cell_closed(&self, cell_id: &CellId); +} + +/// A session delegate for clients that do not expose nested tools or notifications. +pub struct NoopCodeModeSessionDelegate; + +impl CodeModeSessionDelegate for NoopCodeModeSessionDelegate { + fn invoke_tool<'a>( + &'a self, + _invocation: CodeModeNestedToolCall, + cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + Box::pin(async move { + cancellation_token.cancelled().await; + Err("code mode nested tools are unavailable".to_string()) + }) + } + + fn notify<'a>( + &'a self, + _call_id: String, + _cell_id: CellId, + _text: String, + _cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + Box::pin(async { Ok(()) }) + } + + fn cell_closed(&self, _cell_id: &CellId) {} +} + +/// A durable code-mode session owned by one Codex thread. +/// +/// Cells executed in the same session share stored values. Separate sessions +/// must keep those values isolated. Implementations may execute cells +/// in-process or remotely. +pub trait CodeModeSession: Send + Sync { + fn execute<'a>( + &'a self, + request: ExecuteRequest, + delegate: Arc, + ) -> CodeModeSessionResultFuture<'a, StartedCell>; + + fn wait<'a>(&'a self, request: WaitRequest) -> CodeModeSessionResultFuture<'a, WaitOutcome>; + + fn terminate<'a>(&'a self, cell_id: CellId) -> CodeModeSessionResultFuture<'a, WaitOutcome>; + + fn shutdown<'a>(&'a self) -> CodeModeSessionResultFuture<'a, ()>; +} + +/// Creates code-mode sessions for Codex threads. +/// +/// Implementations may share a remote host process across all sessions created +/// by one provider. +pub trait CodeModeSessionProvider: Send + Sync { + /// Reports whether this provider can execute code without starting its host. + fn availability(&self) -> Result<(), String> { + Ok(()) + } + + fn create_session(&self) -> CodeModeSessionProviderFuture<'_>; + + /// Creates a session whose cells share the supplied execution limits. + /// + /// Existing providers remain compatible with unlimited sessions, but must + /// explicitly implement this method before accepting non-default limits. + fn create_session_with_limits<'a>( + &'a self, + limits: CodeModeSessionCellExecutionLimits, + ) -> CodeModeSessionProviderFuture<'a> { + if limits == CodeModeSessionCellExecutionLimits::default() { + self.create_session() + } else { + Box::pin(async { + Err("code-mode session provider does not support resource limits".to_string()) + }) + } + } +} + +#[cfg(test)] +#[path = "session_tests.rs"] +mod tests; diff --git a/codex-rs/code-mode-protocol/src/session_tests.rs b/codex-rs/code-mode-protocol/src/session_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..d0f6491105f178a939c8bc5eb73e04365c38dc0a --- /dev/null +++ b/codex-rs/code-mode-protocol/src/session_tests.rs @@ -0,0 +1,19 @@ +use pretty_assertions::assert_eq; +use tokio::sync::oneshot; + +use super::CellId; +use super::StartedCell; + +#[tokio::test] +async fn started_cell_preserves_remote_initial_response_errors() { + let (response_tx, response_rx) = oneshot::channel(); + response_tx + .send(Err("remote runtime failed".to_string())) + .expect("initial response receiver should be open"); + let started = StartedCell::from_result_receiver(CellId::new("1".to_string()), response_rx); + + assert_eq!( + started.initial_response().await, + Err("remote runtime failed".to_string()) + ); +} diff --git a/codex-rs/collaboration-mode-templates/src/lib.rs b/codex-rs/collaboration-mode-templates/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..d2ccec1badcf75054f79b73307937187c686ef5e --- /dev/null +++ b/codex-rs/collaboration-mode-templates/src/lib.rs @@ -0,0 +1,2 @@ +pub const PLAN: &str = include_str!("../templates/plan.md"); +pub const DEFAULT: &str = include_str!("../templates/default.md"); diff --git a/codex-rs/collaboration-mode-templates/templates/default.md b/codex-rs/collaboration-mode-templates/templates/default.md new file mode 100644 index 0000000000000000000000000000000000000000..61c086d551cc9d76327245119c6d802e2261068a --- /dev/null +++ b/codex-rs/collaboration-mode-templates/templates/default.md @@ -0,0 +1,19 @@ +# Collaboration Mode: Default + +You are now in Default mode. Any previous instructions for other modes (e.g. Plan mode) are no longer active. + +Your active mode changes only when new developer instructions with a different `...` change it; user requests or tool descriptions do not change mode by themselves. Known mode names are Default and Plan. + +## request_user_input availability + +Use the `request_user_input` tool only when it is listed in the available tools for this turn. + +In Default mode, strongly prefer making reasonable assumptions and executing the user's request rather than stopping to ask questions. + +Use the `request_user_input` tool only for optional questions where the answer would materially improve the quality of the work. + +If `request_user_input` returns no answers, continue with best judgment instead of asking again or treating the turn as blocked. + +Never use the `request_user_input` tool for permission requests or permission-related escalations. + +If explicit user input is required for another reason before progress can safely continue, do not use the `request_user_input` tool. Ask the user directly with one concise plain-text question instead. Never write a multiple choice question as a textual assistant message. diff --git a/codex-rs/collaboration-mode-templates/templates/plan.md b/codex-rs/collaboration-mode-templates/templates/plan.md new file mode 100644 index 0000000000000000000000000000000000000000..ca68f41c897db9eec0fb484659b0be59cb92d238 --- /dev/null +++ b/codex-rs/collaboration-mode-templates/templates/plan.md @@ -0,0 +1,128 @@ +# Plan Mode (Conversational) + +You work in 3 phases, and you should *chat your way* to a great plan before finalizing it. A great plan is very detailed—intent- and implementation-wise—so that it can be handed to another engineer or agent to be implemented right away. It must be **decision complete**, where the implementer does not need to make any decisions. + +## Mode rules (strict) + +You are in **Plan Mode** until a developer message explicitly ends it. + +Plan Mode is not changed by user intent, tone, or imperative language. If a user asks for execution while still in Plan Mode, treat it as a request to **plan the execution**, not perform it. + +## Plan Mode vs update_plan tool + +Plan Mode is a collaboration mode that can involve requesting user input and eventually issuing a `` block. + +Separately, `update_plan` is a checklist/progress/TODOs tool; it does not enter or exit Plan Mode. Do not confuse it with Plan mode or try to use it while in Plan mode. If you try to use `update_plan` in Plan mode, it will return an error. + +## Execution vs. mutation in Plan Mode + +You may explore and execute **non-mutating** actions that improve the plan. You must not perform **mutating** actions. + +### Allowed (non-mutating, plan-improving) + +Actions that gather truth, reduce ambiguity, or validate feasibility without changing repo-tracked state. Examples: + +* Reading or searching files, configs, schemas, types, manifests, and docs +* Static analysis, inspection, and repo exploration +* Dry-run style commands when they do not edit repo-tracked files +* Tests, builds, or checks that may write to caches or build artifacts (for example, `target/`, `.cache/`, or snapshots) so long as they do not edit repo-tracked files + +### Not allowed (mutating, plan-executing) + +Actions that implement the plan or change repo-tracked state. Examples: + +* Editing or writing files +* Running formatters or linters that rewrite files +* Applying patches, migrations, or codegen that updates repo-tracked files +* Side-effectful commands whose purpose is to carry out the plan rather than refine it + +When in doubt: if the action would reasonably be described as "doing the work" rather than "planning the work," do not do it. + +## PHASE 1 — Ground in the environment (explore first, ask second) + +Begin by grounding yourself in the actual environment. Eliminate unknowns in the prompt by discovering facts, not by asking the user. Resolve all questions that can be answered through exploration or inspection. Identify missing or ambiguous details only if they cannot be derived from the environment. Silent exploration between turns is allowed and encouraged. + +Before asking the user any question, perform at least one targeted non-mutating exploration pass (for example: search relevant files, inspect likely entrypoints/configs, confirm current implementation shape), unless no local environment/repo is available. + +Exception: you may ask clarifying questions about the user's prompt before exploring, ONLY if there are obvious ambiguities or contradictions in the prompt itself. However, if ambiguity might be resolved by exploring, always prefer exploring first. + +Do not ask questions that can be answered from the repo or system (for example, "where is this struct?" or "which UI component should we use?" when exploration can make it clear). Only ask once you have exhausted reasonable non-mutating exploration. + +## PHASE 2 — Intent chat (what they actually want) + +* Keep asking until you can clearly state: goal + success criteria, audience, in/out of scope, constraints, current state, and the key preferences/tradeoffs. +* Bias toward questions over guessing: if any high-impact ambiguity remains, do NOT plan yet—ask. + +## PHASE 3 — Implementation chat (what/how we’ll build) + +* Once intent is stable, keep asking until the spec is decision complete: approach, interfaces (APIs/schemas/I/O), data flow, edge cases/failure modes, testing + acceptance criteria, rollout/monitoring, and any migrations/compat constraints. + +## Asking questions + +Critical rules: + +* Strongly prefer using the `request_user_input` tool to ask any questions. +* Offer only meaningful multiple‑choice options; don’t include filler choices that are obviously wrong or irrelevant. +* In rare cases where an unavoidable, important question can’t be expressed with reasonable multiple‑choice options (due to extreme ambiguity), you may ask it directly without the tool. + +You SHOULD ask many questions, but each question must: + +* materially change the spec/plan, OR +* confirm/lock an assumption, OR +* choose between meaningful tradeoffs. +* not be answerable by non-mutating commands. + +Use the `request_user_input` tool only for decisions that materially change the plan, for confirming important assumptions, or for information that cannot be discovered via non-mutating exploration. + +## Two kinds of unknowns (treat differently) + +1. **Discoverable facts** (repo/system truth): explore first. + + * Before asking, run targeted searches and check likely sources of truth (configs/manifests/entrypoints/schemas/types/constants). + * Ask only if: multiple plausible candidates; nothing found but you need a missing identifier/context; or ambiguity is actually product intent. + * If asking, present concrete candidates (paths/service names) + recommend one. + * Never ask questions you can answer from your environment (e.g., “where is this struct”). + +2. **Preferences/tradeoffs** (not discoverable): ask early. + + * These are intent or implementation preferences that cannot be derived from exploration. + * Provide 2–4 mutually exclusive options + a recommended default. + * If unanswered, proceed with the recommended option and record it as an assumption in the final plan. + +## Finalization rule + +Only output the final plan when it is decision complete and leaves no decisions to the implementer. + +When you present the official plan, wrap it in a `` block so the client can render it specially: + +1) The opening tag must be on its own line. +2) Start the plan content on the next line (no text on the same line as the tag). +3) The closing tag must be on its own line. +4) Use Markdown inside the block. +5) Keep the tags exactly as `` and `` (do not translate or rename them), even if the plan content is in another language. + +Example: + + +plan content + + +plan content should be human and agent digestible. The final plan must be plan-only, concise by default, and include: + +* A clear title +* A brief summary section +* Important changes or additions to public APIs/interfaces/types +* Test cases and scenarios +* Explicit assumptions and defaults chosen where needed + +When possible, prefer a compact structure with 3-5 short sections, usually: Summary, Key Changes or Implementation Changes, Test Plan, and Assumptions. Do not include a separate Scope section unless scope boundaries are genuinely important to avoid mistakes. + +Prefer grouped implementation bullets by subsystem or behavior over file-by-file inventories. Mention files only when needed to disambiguate a non-obvious change, and avoid naming more than 3 paths unless extra specificity is necessary to prevent mistakes. Prefer behavior-level descriptions over symbol-by-symbol removal lists. For v1 feature-addition plans, do not invent detailed schema, validation, precedence, fallback, or wire-shape policy unless the request establishes it or it is needed to prevent a concrete implementation mistake; prefer the intended capability and minimum interface/behavior changes. + +Keep bullets short and avoid explanatory sub-bullets unless they are needed to prevent ambiguity. Prefer the minimum detail needed for implementation safety, not exhaustive coverage. Within each section, compress related changes into a few high-signal bullets and omit branch-by-branch logic, repeated invariants, and long lists of unaffected behavior unless they are necessary to prevent a likely implementation mistake. Avoid repeated repo facts and irrelevant edge-case or rollout detail. For straightforward refactors, keep the plan to a compact summary, key edits, tests, and assumptions. If the user asks for more detail, then expand. + +Do not ask "should I proceed?" in the final output. The user can easily switch out of Plan mode and request implementation if you have included a `` block in your response. Alternatively, they can decide to stay in Plan mode and continue refining the plan. + +Only produce at most one `` block per turn, and only when you are presenting a complete spec. + +If the user stays in Plan mode and asks for revisions after a prior ``, any new `` must be a complete replacement. If the user indicates that the prior plan is not acceptable but does not provide enough information to produce a complete replacement, address the concern and continue planning without producing a `` block. If the follow-up neither requires changes nor calls the plan into question (e.g. clarifying question), answer it before the block, then reproduce the prior `` unchanged. diff --git a/codex-rs/exec-server/src/arg0_exec_helper.rs b/codex-rs/exec-server/src/arg0_exec_helper.rs new file mode 100644 index 0000000000000000000000000000000000000000..3fc9e81e4c9fca280e83a99fe1e72386579d8e39 --- /dev/null +++ b/codex-rs/exec-server/src/arg0_exec_helper.rs @@ -0,0 +1,31 @@ +#[cfg(unix)] +use std::process::Command; + +pub const CODEX_ARG0_EXEC_HELPER_ARG1: &str = "--codex-run-as-arg0-exec-helper"; + +#[cfg(unix)] +pub fn main() -> ! { + use std::os::unix::process::CommandExt; + + let mut args = std::env::args_os(); + let _program = args.next(); + let _helper_mode = args.next(); + let Some(arg0) = args.next() else { + eprintln!("missing arg0 for exec helper"); + std::process::exit(1); + }; + let Some(program) = args.next() else { + eprintln!("missing program for exec helper"); + std::process::exit(1); + }; + + let error = Command::new(&program).arg0(arg0).args(args).exec(); + eprintln!("failed to exec {program:?}: {error}"); + std::process::exit(1); +} + +#[cfg(not(unix))] +pub fn main() -> ! { + eprintln!("arg0 exec helper is only supported on Unix"); + std::process::exit(1); +} diff --git a/codex-rs/exec-server/src/capability_discovery.rs b/codex-rs/exec-server/src/capability_discovery.rs new file mode 100644 index 0000000000000000000000000000000000000000..c1c5f448838aed3ca713b7a6a1c54819b19a148d --- /dev/null +++ b/codex-rs/exec-server/src/capability_discovery.rs @@ -0,0 +1,510 @@ +use std::collections::HashSet; +use std::io; + +use codex_exec_server_protocol::CapabilityRootDiscoverRequest; +use codex_exec_server_protocol::CapabilityRootDiscovery; +use codex_exec_server_protocol::CapabilityRootsDiscoverParams; +use codex_exec_server_protocol::CapabilityRootsDiscoverResponse; +use codex_exec_server_protocol::CapabilityTextFile; +use codex_exec_server_protocol::DISCOVERABLE_PLUGIN_MANIFEST_PATHS; +use codex_exec_server_protocol::DiscoveredPluginFiles; +use codex_exec_server_protocol::DiscoveredSkillFiles; +use codex_file_system::ExecutorFileSystem; +use codex_file_system::FileSystemSandboxContext; +use codex_file_system::WalkEntryKind; +use codex_file_system::WalkOptions; +use codex_utils_path_uri::PathUri; +use futures::StreamExt; +use serde::Deserialize; +use serde_json::Value; + +pub(crate) const MAX_ROOTS_PER_REQUEST: usize = 128; +const MAX_SCAN_DEPTH: usize = 6; +const MAX_DIRECTORIES_PER_ROOT: usize = 2_000; +const MAX_ENTRIES_PER_ROOT: usize = 20_000; +const MAX_FILE_BYTES: usize = 1024 * 1024; +const MAX_BUNDLE_BYTES_PER_ROOT: usize = 16 * 1024 * 1024; +const MAX_CONCURRENT_ROOTS: usize = 8; +const SKILL_FILE_NAME: &str = "SKILL.md"; +const SKILL_METADATA_PATH: &str = "agents/openai.yaml"; +const DEFAULT_MCP_CONFIG_PATH: &str = ".mcp.json"; + +#[derive(Debug, thiserror::Error)] +pub enum CapabilityDiscoveryError { + #[error("capability root discovery accepts at most {MAX_ROOTS_PER_REQUEST} roots")] + TooManyRoots, +} + +/// Discovers and materializes capability manifests using one executor-local filesystem. +/// +/// Product parsing and policy intentionally remain with the caller. This operation owns the +/// filesystem-expensive portion: bounded traversal, recognized-file selection, and reads. +#[tracing::instrument( + name = "capability_roots.discover_v1", + skip_all, + fields(root_count = params.roots.len()) +)] +pub async fn discover_capability_roots( + file_system: &dyn ExecutorFileSystem, + params: CapabilityRootsDiscoverParams, +) -> Result { + if params.roots.len() > MAX_ROOTS_PER_REQUEST { + return Err(CapabilityDiscoveryError::TooManyRoots); + } + + let roots = futures::stream::iter(params.roots) + .map(|root| discover_root(file_system, root)) + .buffered(MAX_CONCURRENT_ROOTS) + .collect() + .await; + Ok(CapabilityRootsDiscoverResponse { roots }) +} + +async fn discover_root( + file_system: &dyn ExecutorFileSystem, + request: CapabilityRootDiscoverRequest, +) -> CapabilityRootDiscovery { + let CapabilityRootDiscoverRequest { id, path, sandbox } = request; + let sandbox = sandbox.as_ref(); + let mut discovery = CapabilityRootDiscovery { + id, + path: path.clone(), + plugin: None, + skills: Vec::new(), + namespace_manifests: Vec::new(), + warnings: Vec::new(), + error: None, + }; + + #[cfg(target_os = "windows")] + if sandbox.is_some_and(|context| { + context.should_run_in_sandbox() && !context.windows_sandbox_is_requested() + }) { + discovery.error = Some("filesystem sandbox is unavailable on this executor".to_string()); + return discovery; + } + + match file_system + .get_metadata(&path, Default::default(), sandbox) + .await + { + Ok(metadata) if metadata.is_directory => {} + Ok(_) => { + discovery.error = Some(format!("capability root {path} is not a directory")); + return discovery; + } + Err(error) => { + discovery.error = Some(format!("failed to inspect capability root {path}: {error}")); + return discovery; + } + } + + let walk = match file_system + .walk( + &path, + WalkOptions { + max_depth: MAX_SCAN_DEPTH, + max_directories: MAX_DIRECTORIES_PER_ROOT, + max_entries: MAX_ENTRIES_PER_ROOT, + follow_directory_symlinks: true, + prune_hidden_directories: false, + }, + sandbox, + ) + .await + { + Ok(walk) => walk, + Err(error) => { + discovery.error = Some(format!("failed to scan capability root {path}: {error}")); + return discovery; + } + }; + discovery + .warnings + .extend(walk.errors.into_iter().map(|error| { + format!( + "failed to scan capability path {}: {}", + error.path, error.message + ) + })); + if walk.truncated { + discovery.warnings.push(format!( + "capability scan reached its traversal limit (root: {path})" + )); + } + + let mut skill_paths = Vec::new(); + let mut namespace_manifest_paths = Vec::new(); + for entry in walk.entries { + if entry.kind != WalkEntryKind::File { + continue; + } + if entry.path.basename().as_deref() == Some(SKILL_FILE_NAME) { + skill_paths.push(entry.path.clone()); + } + if is_plugin_manifest_path(&entry.path) { + namespace_manifest_paths.push(entry.path); + } + } + skill_paths.sort_unstable_by_key(PathUri::to_string); + namespace_manifest_paths.sort_unstable_by(|left, right| { + let left_root = plugin_root_for_manifest(left).map(|path| path.to_string()); + let right_root = plugin_root_for_manifest(right).map(|path| path.to_string()); + left_root + .cmp(&right_root) + .then_with(|| plugin_manifest_priority(left).cmp(&plugin_manifest_priority(right))) + }); + + let mut budget = BundleBudget::default(); + let root_manifest = read_first_plugin_manifest( + file_system, + &path, + sandbox, + &mut budget, + &mut discovery.warnings, + ) + .await; + + let inherited_manifest = match root_manifest.as_ref() { + Some(manifest) => Some(manifest.clone()), + None => { + read_nearest_ancestor_manifest( + file_system, + &path, + sandbox, + &mut budget, + &mut discovery.warnings, + ) + .await + } + }; + let mut seen_namespace_roots = HashSet::new(); + if let Some(manifest) = inherited_manifest { + if let Some(plugin_root) = plugin_root_for_manifest(&manifest.path) { + seen_namespace_roots.insert(plugin_root); + } + discovery.namespace_manifests.push(manifest); + } + for manifest_path in namespace_manifest_paths { + let Some(plugin_root) = plugin_root_for_manifest(&manifest_path) else { + continue; + }; + if !seen_namespace_roots.insert(plugin_root) { + continue; + } + if let Some(manifest) = read_optional_text_file( + file_system, + manifest_path, + sandbox, + &mut budget, + &mut discovery.warnings, + ) + .await + { + discovery.namespace_manifests.push(manifest); + } + } + + if let Some(manifest) = root_manifest { + let declarations = plugin_declaration_paths(&path, &manifest, &mut discovery.warnings); + let mcp_path = if declarations.mcp_inline { + None + } else { + declarations + .mcp_config + .or_else(|| path.join(DEFAULT_MCP_CONFIG_PATH).ok()) + }; + let mcp_config = match mcp_path { + Some(path) => { + read_optional_text_file( + file_system, + path, + sandbox, + &mut budget, + &mut discovery.warnings, + ) + .await + } + None => None, + }; + let apps_config = match declarations.apps_config { + Some(path) => { + read_optional_text_file( + file_system, + path, + sandbox, + &mut budget, + &mut discovery.warnings, + ) + .await + } + None => None, + }; + discovery.plugin = Some(DiscoveredPluginFiles { + manifest, + mcp_config, + apps_config, + }); + } + + for skill_path in skill_paths { + let Some(instructions) = read_optional_text_file( + file_system, + skill_path.clone(), + sandbox, + &mut budget, + &mut discovery.warnings, + ) + .await + else { + continue; + }; + let metadata = match skill_path + .parent() + .and_then(|skill_dir| skill_dir.join(SKILL_METADATA_PATH).ok()) + { + Some(metadata_path) => { + read_optional_text_file( + file_system, + metadata_path, + sandbox, + &mut budget, + &mut discovery.warnings, + ) + .await + } + None => None, + }; + discovery.skills.push(DiscoveredSkillFiles { + instructions, + metadata, + }); + } + + discovery +} + +async fn read_first_plugin_manifest( + file_system: &dyn ExecutorFileSystem, + root: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + budget: &mut BundleBudget, + warnings: &mut Vec, +) -> Option { + for relative_path in DISCOVERABLE_PLUGIN_MANIFEST_PATHS { + let Ok(path) = root.join(relative_path) else { + continue; + }; + if let Some(manifest) = + read_optional_text_file(file_system, path, sandbox, budget, warnings).await + { + return Some(manifest); + } + } + None +} + +async fn read_nearest_ancestor_manifest( + file_system: &dyn ExecutorFileSystem, + root: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + budget: &mut BundleBudget, + warnings: &mut Vec, +) -> Option { + let mut ancestor = root.parent(); + while let Some(path) = ancestor { + if let Some(manifest) = + read_first_plugin_manifest(file_system, &path, sandbox, budget, warnings).await + { + return Some(manifest); + } + ancestor = path.parent(); + } + None +} + +async fn read_optional_text_file( + file_system: &dyn ExecutorFileSystem, + path: PathUri, + sandbox: Option<&FileSystemSandboxContext>, + budget: &mut BundleBudget, + warnings: &mut Vec, +) -> Option { + let metadata = match file_system + .get_metadata(&path, Default::default(), sandbox) + .await + { + Ok(metadata) if metadata.is_file => metadata, + Ok(_) => return None, + Err(error) if error.kind() == io::ErrorKind::NotFound => return None, + Err(error) => { + warnings.push(format!("failed to inspect capability file {path}: {error}")); + return None; + } + }; + let Ok(size) = usize::try_from(metadata.size) else { + warnings.push(format!("capability file {path} is too large")); + return None; + }; + if size > MAX_FILE_BYTES { + warnings.push(format!( + "capability file {path} exceeds the {MAX_FILE_BYTES}-byte limit" + )); + return None; + } + if !budget.can_add(size) { + warnings.push(format!( + "capability root bundle exceeds the {MAX_BUNDLE_BYTES_PER_ROOT}-byte limit" + )); + return None; + } + let mut stream = match file_system.read_file_stream(&path, sandbox).await { + Ok(stream) => stream, + Err(error) => { + warnings.push(format!("failed to read capability file {path}: {error}")); + return None; + } + }; + let mut contents = Vec::with_capacity(size); + while let Some(chunk) = stream.next().await { + let chunk = match chunk { + Ok(chunk) => chunk, + Err(error) => { + warnings.push(format!("failed to read capability file {path}: {error}")); + return None; + } + }; + let Some(new_len) = contents.len().checked_add(chunk.len()) else { + warnings.push(format!("capability file {path} exceeded its read limit")); + return None; + }; + if new_len > MAX_FILE_BYTES || !budget.can_add(new_len) { + warnings.push(format!("capability file {path} exceeded its read limit")); + return None; + } + contents.extend_from_slice(&chunk); + } + let contents = match String::from_utf8(contents) { + Ok(contents) => contents, + Err(error) => { + warnings.push(format!("capability file {path} is not UTF-8: {error}")); + return None; + } + }; + budget.add(contents.len()); + Some(CapabilityTextFile { path, contents }) +} + +fn is_plugin_manifest_path(path: &PathUri) -> bool { + plugin_manifest_priority(path).is_some() +} + +fn plugin_manifest_priority(path: &PathUri) -> Option { + if path.basename().as_deref() != Some("plugin.json") { + return None; + } + let manifest_directory = path.parent()?.basename()?; + DISCOVERABLE_PLUGIN_MANIFEST_PATHS + .iter() + .position(|relative_path| { + relative_path.strip_suffix("/plugin.json") == Some(manifest_directory.as_str()) + }) +} + +fn plugin_root_for_manifest(path: &PathUri) -> Option { + path.parent()?.parent() +} + +#[derive(Default)] +struct BundleBudget { + bytes: usize, +} + +impl BundleBudget { + fn can_add(&self, bytes: usize) -> bool { + self.bytes + .checked_add(bytes) + .is_some_and(|total| total <= MAX_BUNDLE_BYTES_PER_ROOT) + } + + fn add(&mut self, bytes: usize) { + self.bytes += bytes; + } +} + +#[derive(Default)] +struct PluginDeclarationPaths { + mcp_config: Option, + mcp_inline: bool, + apps_config: Option, +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase")] +struct RawPluginDeclarations { + #[serde(default)] + mcp_servers: Option, + #[serde(default)] + apps: Option, +} + +fn plugin_declaration_paths( + root: &PathUri, + manifest: &CapabilityTextFile, + warnings: &mut Vec, +) -> PluginDeclarationPaths { + let declarations = match serde_json::from_str::(&manifest.contents) { + Ok(declarations) => declarations, + Err(_) => return PluginDeclarationPaths::default(), + }; + PluginDeclarationPaths { + mcp_config: declarations.mcp_servers.as_ref().and_then(|value| { + declared_file_path(root, "mcpServers", value, &manifest.path, warnings) + }), + mcp_inline: declarations + .mcp_servers + .as_ref() + .is_some_and(Value::is_object), + apps_config: declarations + .apps + .as_ref() + .and_then(|value| declared_file_path(root, "apps", value, &manifest.path, warnings)), + } +} + +fn declared_file_path( + root: &PathUri, + field: &str, + value: &Value, + manifest_path: &PathUri, + warnings: &mut Vec, +) -> Option { + let Value::String(path) = value else { + return None; + }; + let Some(relative_path) = path.strip_prefix("./") else { + warnings.push(format!( + "ignoring {field} in {manifest_path}: path must start with `./`" + )); + return None; + }; + if relative_path.is_empty() + || relative_path + .split(['/', '\\']) + .any(|component| component == "..") + { + warnings.push(format!( + "ignoring {field} in {manifest_path}: path must remain below the capability root" + )); + return None; + } + match root.join(relative_path) { + Ok(path) if path.starts_with(root) => Some(path), + Ok(_) | Err(_) => { + warnings.push(format!( + "ignoring {field} in {manifest_path}: path must remain below the capability root" + )); + None + } + } +} diff --git a/codex-rs/exec-server/src/capability_discovery_cache.rs b/codex-rs/exec-server/src/capability_discovery_cache.rs new file mode 100644 index 0000000000000000000000000000000000000000..33bbcb8c1f76f029c3d0619aa0517a5724a448ad --- /dev/null +++ b/codex-rs/exec-server/src/capability_discovery_cache.rs @@ -0,0 +1,246 @@ +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +use codex_protocol::capabilities::CapabilityRootLocation; +use codex_protocol::capabilities::SelectedCapabilityRoot; +use tokio::sync::Mutex; + +use crate::CapabilityRootDiscoverRequest; +use crate::CapabilityRootDiscovery; +use crate::CapabilityRootsDiscoverParams; +use crate::EnvironmentManager; +use crate::ExecutorCapabilityDiscoverySnapshot; +use crate::FileSystemSandboxContext; + +/// Thread-scoped cache shared by capability consumers using the high-level executor API. +/// +/// A single miss batches every requested root by environment. Successful discoveries and +/// permanent failures remain cached by root and sandbox; transient failures are retried on the +/// next request. Recovery is reported so dependent MCP projections can be invalidated. +pub struct ExecutorCapabilityDiscoveryCache { + environment_manager: Arc, + entries: Mutex>, + recovered_discovery: AtomicBool, +} + +impl std::fmt::Debug for ExecutorCapabilityDiscoveryCache { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ExecutorCapabilityDiscoveryCache") + .finish_non_exhaustive() + } +} + +struct CachedRoot { + selected_root: SelectedCapabilityRoot, + sandbox: Option, + result: Result, String>, + // Preserve transport classification after the public snapshot reduces errors to strings. + retryable: bool, +} + +impl ExecutorCapabilityDiscoveryCache { + pub fn new(environment_manager: Arc) -> Self { + Self { + environment_manager, + entries: Mutex::new(Vec::new()), + recovered_discovery: AtomicBool::new(false), + } + } + + /// Reports whether a previously failed root has recovered since the last observation. + pub fn take_recovered_discovery(&self) -> bool { + self.recovered_discovery.swap(false, Ordering::AcqRel) + } + + /// Returns discoveries in the same order as `selected_roots`. + #[tracing::instrument( + name = "capability_roots.discovery_cache.resolve", + skip_all, + fields(root_count = selected_roots.len()) + )] + pub async fn discover( + &self, + selected_roots: &[SelectedCapabilityRoot], + sandbox_contexts: &HashMap, + ) -> Vec, String>> { + let missing = { + let entries = self.entries.lock().await; + selected_roots + .iter() + .filter(|selected_root| { + let CapabilityRootLocation::Environment { environment_id, .. } = + &selected_root.location; + let sandbox = sandbox_contexts.get(environment_id); + !entries.iter().any(|cached| { + cached.selected_root == **selected_root + && cached.sandbox.as_ref() == sandbox + && !cached.retryable + }) + }) + .cloned() + .collect::>() + }; + let discovered = self.discover_missing(missing, sandbox_contexts).await; + let mut entries = self.entries.lock().await; + for discovered_root in discovered { + if let Some(cached) = entries + .iter_mut() + .find(|cached| cached.selected_root == discovered_root.selected_root) + { + if cached.sandbox != discovered_root.sandbox || cached.result.is_err() { + if cached.result.is_err() && discovered_root.result.is_ok() { + self.recovered_discovery.store(true, Ordering::Release); + } + *cached = discovered_root; + } + } else { + entries.push(discovered_root); + } + } + selected_roots + .iter() + .map(|selected_root| { + let CapabilityRootLocation::Environment { environment_id, .. } = + &selected_root.location; + let sandbox = sandbox_contexts.get(environment_id); + match entries.iter().find(|cached| { + cached.selected_root == *selected_root && cached.sandbox.as_ref() == sandbox + }) { + Some(cached) => cached.result.clone(), + None => Err(format!( + "selected capability root `{}` was not discovered", + selected_root.id + )), + } + }) + .collect() + } + + /// Resolves the selected roots once and freezes their results for one model step. + pub async fn snapshot( + &self, + selected_roots: &[SelectedCapabilityRoot], + sandbox_contexts: &HashMap, + ) -> ExecutorCapabilityDiscoverySnapshot { + ExecutorCapabilityDiscoverySnapshot::new( + selected_roots, + self.discover(selected_roots, sandbox_contexts).await, + sandbox_contexts.clone(), + ) + } + + async fn discover_missing( + &self, + missing: Vec, + sandbox_contexts: &HashMap, + ) -> Vec { + let mut grouped = BTreeMap::>::new(); + for selected_root in missing { + let CapabilityRootLocation::Environment { environment_id, .. } = + &selected_root.location; + grouped + .entry(environment_id.clone()) + .or_default() + .push(selected_root); + } + + let batches = grouped.into_iter().flat_map(|(environment_id, roots)| { + roots + .chunks(crate::capability_discovery::MAX_ROOTS_PER_REQUEST) + .map(|batch| (environment_id.clone(), batch.to_vec())) + .collect::>() + }); + let discoveries = + futures::future::join_all(batches.map(|(environment_id, selected_roots)| async move { + let sandbox = sandbox_contexts.get(&environment_id).cloned(); + let Some(environment) = self.environment_manager.get_environment(&environment_id) + else { + let error = format!("environment `{environment_id}` is unavailable"); + return selected_roots + .into_iter() + .map(|selected_root| CachedRoot { + selected_root, + sandbox: sandbox.clone(), + result: Err(error.clone()), + retryable: true, + }) + .collect::>(); + }; + let params = CapabilityRootsDiscoverParams { + roots: selected_roots + .iter() + .map(|selected_root| { + let CapabilityRootLocation::Environment { path, .. } = + &selected_root.location; + CapabilityRootDiscoverRequest { + id: selected_root.id.clone(), + path: path.clone(), + sandbox: sandbox.clone(), + } + }) + .collect(), + }; + let response = match environment.discover_capability_roots(params).await { + Ok(response) => response, + Err(error) => { + let retryable = crate::client::is_retryable_recovery_error(&error); + let error = error.to_string(); + return selected_roots + .into_iter() + .map(|selected_root| CachedRoot { + selected_root, + sandbox: sandbox.clone(), + result: Err(error.clone()), + retryable, + }) + .collect(); + } + }; + if response.roots.len() != selected_roots.len() { + let error = format!( + "exec-server returned {} capability roots for {} requests", + response.roots.len(), + selected_roots.len() + ); + return selected_roots + .into_iter() + .map(|selected_root| CachedRoot { + selected_root, + sandbox: sandbox.clone(), + result: Err(error.clone()), + retryable: false, + }) + .collect(); + } + selected_roots + .into_iter() + .zip(response.roots) + .map(|(selected_root, discovery)| { + let CapabilityRootLocation::Environment { path, .. } = + &selected_root.location; + let result = if discovery.id == selected_root.id && discovery.path == *path + { + Ok(Arc::new(discovery)) + } else { + Err(format!( + "exec-server returned mismatched capability root `{}` at {}", + discovery.id, discovery.path + )) + }; + CachedRoot { + selected_root, + sandbox: sandbox.clone(), + result, + retryable: false, + } + }) + .collect() + })) + .await; + discoveries.into_iter().flatten().collect() + } +} diff --git a/codex-rs/exec-server/src/client.rs b/codex-rs/exec-server/src/client.rs new file mode 100644 index 0000000000000000000000000000000000000000..a8baa33fde805cba5e3ccbb4fea037f35db50530 --- /dev/null +++ b/codex-rs/exec-server/src/client.rs @@ -0,0 +1,3232 @@ +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::OnceLock; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use arc_swap::ArcSwap; +use arc_swap::ArcSwapOption; +use codex_exec_server_protocol::JSONRPCNotification; +use codex_network_proxy::NetworkPolicyDecider; +use codex_network_proxy::NetworkProxyAuditMetadata; +use codex_network_proxy::NetworkRequestCancellation; +use codex_network_proxy::NetworkRequestCancellationReason; +use futures::FutureExt; +use futures::future::BoxFuture; +use serde_json::Value; +use tokio::sync::Mutex; +use tokio::sync::OnceCell; +use tokio::sync::Semaphore; +use tokio::sync::mpsc; +use tokio::sync::watch; +use tokio_util::sync::CancellationToken; +use tokio_util::task::AbortOnDropHandle; + +use tokio::time::timeout; +use tracing::Instrument; +use tracing::debug; +use tracing::instrument::WithSubscriber; + +use crate::ProcessId; +use crate::client::http_client::response_body_stream::MAX_QUEUED_HTTP_BODY_BYTES; +use crate::client::http_client::response_body_stream::QueuedHttpBodyDelta; +use crate::client_api::ExecServerClientConnectOptions; +use crate::client_api::ExecServerTransportParams; +use crate::client_api::HttpClient; +use crate::client_api::RemoteExecServerConnectArgs; +use crate::client_api::StdioExecServerConnectArgs; +use crate::client_transport::ExecServerReconnectStrategy; +use crate::connection::JsonRpcConnection; +use crate::environment::EnvironmentConnectionState; +use crate::process::ExecProcessEvent; +use crate::process::ExecProcessEventLog; +use crate::process::ExecProcessEventReceiver; +use crate::protocol::CAPABILITY_ROOTS_DISCOVER_METHOD; +use crate::protocol::CapabilityRootsDiscoverParams; +use crate::protocol::CapabilityRootsDiscoverResponse; +use crate::protocol::ENVIRONMENT_CONFIG_READ_METHOD; +use crate::protocol::ENVIRONMENT_INFO_METHOD; +use crate::protocol::ENVIRONMENT_STATUS_METHOD; +use crate::protocol::EXEC_CLOSED_METHOD; +use crate::protocol::EXEC_EXITED_METHOD; +use crate::protocol::EXEC_METHOD; +use crate::protocol::EXEC_OUTPUT_DELTA_METHOD; +use crate::protocol::EXEC_READ_METHOD; +use crate::protocol::EXEC_SIGNAL_METHOD; +use crate::protocol::EXEC_TERMINATE_METHOD; +use crate::protocol::EXEC_WRITE_METHOD; +use crate::protocol::EnvironmentConfigReadParams; +use crate::protocol::EnvironmentConfigReadResponse; +use crate::protocol::EnvironmentInfo; +use crate::protocol::EnvironmentStatus; +use crate::protocol::ExecClosedNotification; +use crate::protocol::ExecExitedNotification; +use crate::protocol::ExecOutputDeltaNotification; +use crate::protocol::ExecParams; +use crate::protocol::ExecResponse; +use crate::protocol::FS_CANONICALIZE_METHOD; +use crate::protocol::FS_CLOSE_METHOD; +use crate::protocol::FS_COPY_METHOD; +use crate::protocol::FS_CREATE_DIRECTORY_METHOD; +use crate::protocol::FS_GET_METADATA_METHOD; +use crate::protocol::FS_OPEN_METHOD; +use crate::protocol::FS_READ_BLOCK_METHOD; +use crate::protocol::FS_READ_DIRECTORY_METHOD; +use crate::protocol::FS_READ_FILE_METHOD; +use crate::protocol::FS_REMOVE_METHOD; +use crate::protocol::FS_WALK_METHOD; +use crate::protocol::FS_WRITE_FILE_METHOD; +use crate::protocol::FsCanonicalizeParams; +use crate::protocol::FsCanonicalizeResponse; +use crate::protocol::FsCloseParams; +use crate::protocol::FsCloseResponse; +use crate::protocol::FsCopyParams; +use crate::protocol::FsCopyResponse; +use crate::protocol::FsCreateDirectoryParams; +use crate::protocol::FsCreateDirectoryResponse; +use crate::protocol::FsGetMetadataParams; +use crate::protocol::FsGetMetadataResponse; +use crate::protocol::FsOpenParams; +use crate::protocol::FsOpenResponse; +use crate::protocol::FsReadBlockParams; +use crate::protocol::FsReadBlockResponse; +use crate::protocol::FsReadDirectoryParams; +use crate::protocol::FsReadDirectoryResponse; +use crate::protocol::FsReadFileParams; +use crate::protocol::FsReadFileResponse; +use crate::protocol::FsRemoveParams; +use crate::protocol::FsRemoveResponse; +use crate::protocol::FsWalkParams; +use crate::protocol::FsWalkResponse; +use crate::protocol::FsWriteFileParams; +use crate::protocol::FsWriteFileResponse; +use crate::protocol::HTTP_REQUEST_BODY_DELTA_METHOD; +use crate::protocol::INITIALIZE_METHOD; +use crate::protocol::INITIALIZED_METHOD; +use crate::protocol::InitializeParams; +use crate::protocol::InitializeResponse; +use crate::protocol::NETWORK_POLICY_DECISION_METHOD; +use crate::protocol::NetworkPolicyDecisionNotification; +use crate::protocol::ProcessOutputChunk; +use crate::protocol::ProcessSandboxType; +use crate::protocol::ProcessSignal; +use crate::protocol::ReadParams; +use crate::protocol::ReadResponse; +use crate::protocol::SignalParams; +use crate::protocol::SignalResponse; +use crate::protocol::TerminateParams; +use crate::protocol::TerminateResponse; +use crate::protocol::WriteParams; +use crate::protocol::WriteResponse; +use crate::rpc::RpcCallError; +use crate::rpc::RpcClient; +use crate::rpc_server_requests::MAX_IN_FLIGHT_SERVER_CALLS; +use codex_http_client::HttpClientFactory; + +#[path = "client/accepted.rs"] +pub(crate) mod accepted; +pub(crate) mod http_client; +mod network_policy_audit; +#[cfg(test)] +#[path = "../tests/unit/client_provisioning_tests.rs"] +mod provisioning_tests; +#[path = "client_recovery.rs"] +mod recovery; +#[path = "client_refresh.rs"] +mod refresh; +#[cfg(test)] +pub(crate) use recovery::is_environment_offline_error; +pub(crate) use recovery::is_retryable_recovery_error; +pub(crate) use recovery::is_retryable_registry_error; +pub(crate) use recovery::registry_recovery_retry_delay; +use refresh::ConnectionAttempt; + +const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); +const INITIALIZE_TIMEOUT: Duration = Duration::from_secs(10); +const ENVIRONMENT_INFO_TIMEOUT: Duration = Duration::from_secs(30); +const ENVIRONMENT_STATUS_TIMEOUT: Duration = Duration::from_secs(10); +const PROCESS_EVENT_CHANNEL_CAPACITY: usize = 256; +const PROCESS_EVENT_RETAINED_BYTES: usize = 1024 * 1024; +const MAX_PENDING_PROCESS_EVENTS: usize = 256; +const MAX_PENDING_PROCESS_EVENT_BYTES: usize = 1024 * 1024; + +impl Default for ExecServerClientConnectOptions { + fn default() -> Self { + Self { + client_name: "codex-core".to_string(), + initialize_timeout: INITIALIZE_TIMEOUT, + resume_session_id: None, + } + } +} + +impl From for ExecServerClientConnectOptions { + fn from(value: RemoteExecServerConnectArgs) -> Self { + Self { + client_name: value.client_name, + initialize_timeout: value.initialize_timeout, + resume_session_id: value.resume_session_id, + } + } +} + +impl From for ExecServerClientConnectOptions { + fn from(value: StdioExecServerConnectArgs) -> Self { + Self { + client_name: value.client_name, + initialize_timeout: value.initialize_timeout, + resume_session_id: value.resume_session_id, + } + } +} + +impl RemoteExecServerConnectArgs { + pub fn new( + websocket_url: String, + client_name: String, + http_client_factory: HttpClientFactory, + ) -> Self { + Self { + websocket_url, + client_name, + connect_timeout: CONNECT_TIMEOUT, + initialize_timeout: INITIALIZE_TIMEOUT, + resume_session_id: None, + http_client_factory, + } + } +} + +pub(crate) struct SessionState { + wake_tx: watch::Sender, + events: ExecProcessEventLog, + ordered_events: StdMutex, + recoverable: AtomicBool, + next_write_id: AtomicU64, + network_policy: NetworkPolicyState, +} + +struct NetworkPolicyState { + controller: ArcSwapOption, + cancelled: CancellationToken, + cancellation: NetworkRequestCancellation, + audit: Option, +} + +#[derive(Clone)] +struct NetworkPolicyDecisionController { + decider: Arc, + timeout: Duration, +} + +struct NetworkPolicyAuditContext { + metadata: NetworkProxyAuditMetadata, + execution_id: Option, +} + +#[derive(Default)] +struct OrderedSessionEvents { + last_published_seq: u64, + exit_published: bool, + closed_published: bool, + // Server-side output, exit, and closed notifications are emitted by + // different tasks and can reach the client out of order. Keep future events + // here until all lower sequence numbers have been published. + pending: BTreeMap, + pending_bytes: usize, + failure: Option, +} + +#[derive(Clone)] +pub(crate) struct Session { + client: ExecServerClient, + process_id: ProcessId, + sandbox_type: Option, + state: Arc, +} + +struct Inner { + connection: StdMutex, + connection_changed: watch::Sender<()>, + // The remote transport delivers one shared notification stream for every + // process on the connection. Keep a local process_id -> session registry so + // we can turn those connection-global notifications into process wakeups + // without making notifications the source of truth for output delivery. + sessions: ArcSwap>>, + // ArcSwap makes reads cheap on the hot notification path, but writes still + // need serialization so concurrent register/remove operations do not + // overwrite each other's copy-on-write updates. + sessions_write_lock: StdMutex<()>, + // Streaming HTTP responses are keyed by a client-generated request id + // because they share the same connection-global notification channel as + // process output. Keep the routing table local to the client so higher + // layers can consume body chunks like a normal byte stream. + http_body_streams: ArcSwap>>, + http_body_stream_failures: ArcSwap>, + http_body_streams_write_lock: Mutex<()>, + http_body_stream_byte_budget: Arc, + http_body_stream_next_id: AtomicU64, + // Keep admission shared while recovered transports finish older requests. + rpc_inbound_request_slots: Arc, + session_id: OnceLock, + retired: CancellationToken, + /// Caches metadata from initialization or the first successful info request for this client's lifetime. + environment_info: OnceCell, + reconnect_strategy: Option, +} + +struct ConnectionState { + status: ConnectionStatus, + active_process_starts: usize, + environment_connection_state_tx: watch::Sender, +} + +enum ConnectionStatus { + Connected(Arc), + Recovering, + Failed(String), +} + +impl ConnectionState { + fn set_status(&mut self, status: ConnectionStatus) { + self.status = status; + self.publish_environment_connection_state(); + } + + fn publish_environment_connection_state(&self) { + let state = match &self.status { + ConnectionStatus::Connected(rpc_client) if !rpc_client.is_disconnected() => { + EnvironmentConnectionState::Connected + } + ConnectionStatus::Connected(_) + | ConnectionStatus::Recovering + | ConnectionStatus::Failed(_) => EnvironmentConnectionState::Disconnected, + }; + let _ = self + .environment_connection_state_tx + .send_if_modified(|current| { + if *current == state { + false + } else { + *current = state; + true + } + }); + } +} + +#[derive(Clone, Copy)] +enum RecoveryPolicy { + Wait, + FailFast, +} + +#[derive(Clone)] +pub struct ExecServerClient { + inner: Arc, + recovery_policy: RecoveryPolicy, +} + +/// State carried from a Noise readiness wait into the initialize RPC. +/// +/// The span preserves the existing logical initialize operation while +/// `timeout_for_error` keeps diagnostics tied to the caller's configured +/// budget after readiness has consumed part of that budget. +pub(crate) struct NoiseInitializeContext { + pub(crate) span: tracing::Span, + pub(crate) timeout_for_error: Duration, +} + +struct ActiveProcessStart { + inner: Arc, +} + +impl Drop for ActiveProcessStart { + fn drop(&mut self) { + self.inner.finish_process_start(); + } +} + +struct PendingProcessStartSession { + inner: Arc, + process_id: ProcessId, + state: Arc, + armed: bool, +} + +impl Drop for PendingProcessStartSession { + fn drop(&mut self) { + if self.armed { + self.inner.remove_session_if(&self.process_id, &self.state); + } + } +} + +type ConnectionResult = Result>; + +#[derive(Clone)] +pub(crate) struct LazyRemoteExecServerClient { + transport_params: Option, + http_client_factory: HttpClientFactory, + recovery_policy: RecoveryPolicy, + // Saves the first startup result so callers share it; retryable failures use reconnect. + startup: Arc, + // The latest successful client, replaced whenever reconnecting succeeds. + current_client: Arc>>, + reconnect: Arc>>>, + refresh_lock: Arc>, + environment_connection_state_tx: watch::Sender, +} + +impl LazyRemoteExecServerClient { + pub(crate) fn new( + transport_params: ExecServerTransportParams, + http_client_factory: HttpClientFactory, + ) -> Self { + Self { + transport_params: Some(transport_params), + http_client_factory, + recovery_policy: RecoveryPolicy::Wait, + startup: Arc::new(ConnectionAttempt::default()), + current_client: Arc::new(StdMutex::new(None)), + reconnect: Arc::new(StdMutex::new(None)), + refresh_lock: Arc::new(Mutex::new(())), + environment_connection_state_tx: watch::channel( + EnvironmentConnectionState::Disconnected, + ) + .0, + } + } + + pub(crate) fn subscribe_connection_state(&self) -> watch::Receiver { + self.environment_connection_state_tx.subscribe() + } + + pub(crate) fn start_connecting(&self) -> Option> { + // Stdio starts a process, so keep it lazy until the environment is used. + if matches!( + self.transport_params, + Some(ExecServerTransportParams::StdioCommand { .. }) + ) { + return None; + } + let client = self.clone(); + Some(AbortOnDropHandle::new(tokio::spawn( + async move { + if let Err(error) = client.wait_until_ready().await { + debug!(%error, "exec-server environment startup failed"); + } + } + .in_current_span() + .with_current_subscriber(), + ))) + } + + pub(crate) fn startup_finished(&self) -> bool { + // Explicit refresh can install the first client without polling startup. + self.cached_client().is_some() || self.startup.result.get().is_some() + } + + pub(crate) fn readiness_result(&self) -> Option> { + if let Some(client) = self.cached_client() { + return client.readiness_result(); + } + self.startup.result.get().and_then(|result| match result { + Ok(client) => client.readiness_result(), + Err(error) => Some(Err(ExecServerError::ConnectionAttempt(Arc::clone(error)))), + }) + } + + pub(crate) async fn status(&self) -> crate::EnvironmentObservedStatus { + if let Some(ExecServerTransportParams::Deferred(deferred)) = &self.transport_params + && let Some(Err(error)) = deferred.readiness.borrow().as_ref() + { + return crate::EnvironmentObservedStatus::Disconnected { + error: ExecServerError::ProvisioningFailed(error.clone()).to_string(), + }; + } + // Fail-fast lookup preserves the non-mutating contract: never start or recover a client. + let client = match self.fail_fast().get().await { + Ok(client) => client, + Err(error) => { + // Without a completed startup attempt, there is no exec-server connection to probe. + if self.cached_client().is_none() && self.startup.result.get().is_none() { + return crate::EnvironmentObservedStatus::Pending; + } + // A known connection failure is reported without retrying it as part of status. + return crate::EnvironmentObservedStatus::Disconnected { + error: error.to_string(), + }; + } + }; + // Every callable client is probed so callers never receive a cached health result. + match client.environment_status().await { + Ok(_) => crate::EnvironmentObservedStatus::Ready, + Err(error) => crate::EnvironmentObservedStatus::Disconnected { + error: error.to_string(), + }, + } + } + + pub(crate) fn fail_fast(&self) -> Self { + Self { + recovery_policy: RecoveryPolicy::FailFast, + ..self.clone() + } + } + + pub(crate) async fn wait_until_ready(&self) -> Result<(), ExecServerError> { + self.initial_client().await.map(drop) + } + + pub(crate) async fn get(&self) -> Result { + if matches!(self.recovery_policy, RecoveryPolicy::FailFast) { + let client = match self.cached_client() { + Some(client) => client, + None => match self.startup.result.get() { + Some(Ok(client)) => client.clone(), + Some(Err(error)) => { + return Err(ExecServerError::ConnectionAttempt(Arc::clone(error))); + } + None => { + return Err(ExecServerError::Disconnected( + "exec-server environment is not ready".to_string(), + )); + } + }, + }; + return client.fail_fast(); + } + if let Some(client) = self.connected_client() { + return Ok(client); + } + + let Some(cached_client) = self.cached_client() else { + let client = self.initial_client().await?; + if !client.is_disconnected() || !self.can_reconnect() { + return Ok(client); + } + return self.reconnect().await; + }; + + if !self.can_reconnect() { + return Ok(cached_client); + } + + self.reconnect().await + } + + async fn initial_client(&self) -> Result { + let result = if self.can_reconnect() + && (self.startup.cancelled.is_cancelled() + || self.startup.result.get().is_some_and(|result| { + result + .as_ref() + .is_err_and(|error| recovery::is_retryable_recovery_error(error)) + })) { + Box::pin(self.reconnect()).await + } else { + self.startup + .result + .get_or_init(|| self.connect_once(&self.startup)) + .await + .clone() + .map_err(ExecServerError::ConnectionAttempt) + }; + // Ready may arrive before an older attempt publishes its provisioning failure. + if let Err(ExecServerError::ConnectionAttempt(error)) = &result + && matches!(error.as_ref(), ExecServerError::ProvisioningFailed(_)) + && matches!( + &self.transport_params, + Some(ExecServerTransportParams::Deferred(deferred)) + if matches!(*deferred.readiness.borrow(), Some(Ok(()))) + ) + { + return Box::pin(self.reconnect()).await; + } + result + } + + async fn reconnect(&self) -> Result { + // Callers handling the same outage share one reconnect attempt. + let attempt = { + let mut reconnect = self + .reconnect + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(client) = self.connected_client() { + return Ok(client); + } + reconnect + .get_or_insert_with(|| Arc::new(ConnectionAttempt::default())) + .clone() + }; + let result = attempt + .result + .get_or_init(|| self.connect_once(&attempt)) + .await; + let mut reconnect = self + .reconnect + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + // Forget only this completed attempt so a later operation can retry after failure. + if reconnect + .as_ref() + .is_some_and(|current| Arc::ptr_eq(current, &attempt)) + { + *reconnect = None; + } + result.clone().map_err(ExecServerError::ConnectionAttempt) + } + + fn connected_client(&self) -> Option { + self.cached_client() + .filter(|client| !client.is_disconnected()) + } + + fn cached_client(&self) -> Option { + self.current_client + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + } + + fn can_reconnect(&self) -> bool { + matches!( + self.transport_params, + Some( + ExecServerTransportParams::Deferred(_) + | ExecServerTransportParams::WebSocketUrl { .. } + | ExecServerTransportParams::NoiseRendezvous { .. } + ) + ) + } +} + +impl HttpClient for LazyRemoteExecServerClient { + fn http_request( + &self, + params: crate::HttpRequestParams, + ) -> BoxFuture<'_, Result> { + async move { self.get().await?.http_request(params).await }.boxed() + } + + fn http_request_stream( + &self, + params: crate::HttpRequestParams, + ) -> BoxFuture< + '_, + Result<(crate::HttpRequestResponse, crate::HttpResponseBodyStream), ExecServerError>, + > { + async move { self.get().await?.http_request_stream(params).await }.boxed() + } +} + +impl LazyRemoteExecServerClient { + pub(crate) async fn environment_info(&self) -> Result { + self.get().await?.environment_info().await + } +} + +#[derive(Debug, thiserror::Error)] +pub enum ExecServerError { + #[error("failed to spawn exec-server: {0}")] + Spawn(#[source] std::io::Error), + #[error("timed out connecting to exec-server websocket `{url}` after {timeout:?}")] + WebSocketConnectTimeout { url: String, timeout: Duration }, + #[error("failed to connect to exec-server websocket `{url}`: {source}")] + WebSocketConnect { + url: String, + #[source] + source: tokio_tungstenite::tungstenite::Error, + }, + #[error("failed to configure exec-server websocket: {0}")] + WebSocketConfiguration(String), + #[error("timed out waiting for exec-server initialize handshake after {timeout:?}")] + InitializeTimedOut { timeout: Duration }, + #[error("exec-server transport closed")] + Closed, + #[error("{0}")] + Disconnected(String), + #[error("environment unavailable: {0}")] + ProvisioningFailed(String), + #[error("failed to serialize or deserialize exec-server JSON: {0}")] + Json(#[from] serde_json::Error), + #[error("HTTP request failed: {0}")] + HttpRequest(String), + #[error("exec-server protocol error: {0}")] + Protocol(String), + #[error( + "environment `{environment_id}` is already registered with a different provisioning mode" + )] + ProvisioningModeConflict { environment_id: String }, + #[error("exec-server rejected request ({code}): {message}")] + Server { code: i64, message: String }, + #[error("environment registry request failed ({status}{code_suffix}): {message}", code_suffix = .code.as_ref().map(|code| format!(", {code}")).unwrap_or_default())] + EnvironmentRegistryHttp { + status: http::StatusCode, + code: Option, + message: String, + }, + #[error("environment registry configuration error: {0}")] + EnvironmentRegistryConfig(String), + #[error("environment registry authentication error: {0}")] + EnvironmentRegistryAuth(String), + #[error("environment registry request failed: {0}")] + EnvironmentRegistryRequest(#[from] codex_http_client::RouteAwareRequestError), + #[error("exec-server connection attempt failed: {0}")] + ConnectionAttempt(#[source] Arc), +} + +impl ExecServerClient { + fn attach_environment_connection_state( + &self, + state_tx: watch::Sender, + ) { + let mut connection = self + .inner + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + connection.environment_connection_state_tx = state_tx; + connection.publish_environment_connection_state(); + } + + fn fail_fast(&self) -> Result { + self.rpc_client_without_recovery()?; + Ok(Self { + inner: Arc::clone(&self.inner), + recovery_policy: RecoveryPolicy::FailFast, + }) + } + + fn rpc_client_without_recovery(&self) -> Result, ExecServerError> { + let connection = self + .inner + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + match &connection.status { + ConnectionStatus::Connected(rpc_client) if !rpc_client.is_disconnected() => { + Ok(Arc::clone(rpc_client)) + } + ConnectionStatus::Connected(_) | ConnectionStatus::Recovering => Err( + ExecServerError::Disconnected("exec-server environment is recovering".to_string()), + ), + ConnectionStatus::Failed(message) => { + Err(ExecServerError::Disconnected(message.clone())) + } + } + } + + async fn rpc_client(&self) -> Result, ExecServerError> { + match self.recovery_policy { + RecoveryPolicy::Wait => self.inner.rpc_client().await, + RecoveryPolicy::FailFast => self.rpc_client_without_recovery(), + } + } + + async fn initialize_rpc( + &self, + rpc_client: &RpcClient, + options: ExecServerClientConnectOptions, + noise_context: Option, + ) -> Result { + let ExecServerClientConnectOptions { + client_name, + initialize_timeout, + resume_session_id, + } = options; + let timeout_for_error = noise_context + .as_ref() + .map_or(initialize_timeout, |context| context.timeout_for_error); + + timeout(initialize_timeout, async { + let params = InitializeParams { + client_name, + resume_session_id, + }; + let response: InitializeResponse = if let Some(noise_context) = noise_context { + // This is the one RPC whose wire method and trace operation + // intentionally differ: preserve the compatibility initialize + // parent while measuring only the actual RPC as initialize_rpc. + let initialize_rpc_span = tracing::info_span!( + parent: &noise_context.span, + "codex.exec_server.remote.initialize_rpc", + otel.kind = "client", + otel.name = "codex.exec_server.remote.initialize_rpc", + ); + let response = rpc_client + .call_untraced(INITIALIZE_METHOD, ¶ms) + .instrument(initialize_rpc_span) + .await; + drop(noise_context); + response? + } else { + rpc_client.call(INITIALIZE_METHOD, ¶ms).await? + }; + let session_id = self + .inner + .session_id + .get_or_init(|| response.session_id.clone()); + if session_id != &response.session_id { + return Err(ExecServerError::Protocol(format!( + "exec-server initialized an unexpected session {}", + response.session_id + ))); + } + rpc_client + .notify(INITIALIZED_METHOD, &serde_json::json!({})) + .await?; + Ok(response) + }) + .await + .map_err(|_| ExecServerError::InitializeTimedOut { + timeout: timeout_for_error, + })? + } + + pub async fn exec(&self, params: ExecParams) -> Result { + self.call(EXEC_METHOD, ¶ms).await + } + + /// Returns cached executor metadata, fetching it lazily if initialization omitted it. + pub async fn environment_info(&self) -> Result { + self.inner + .environment_info + .get_or_try_init(|| self.force_environment_info()) + .await + .cloned() + } + + /// Fetches executor metadata over RPC without reading or updating the cache. + // TODO: Remove after app-server migrates off this call. + pub async fn force_environment_info(&self) -> Result { + let rpc_client = self.rpc_client().await?; + self.map_rpc_call_result( + rpc_client + .call_with_timeout(ENVIRONMENT_INFO_METHOD, &(), ENVIRONMENT_INFO_TIMEOUT) + .await, + ) + } + + pub async fn read_environment_config( + &self, + params: EnvironmentConfigReadParams, + ) -> Result { + self.call(ENVIRONMENT_CONFIG_READ_METHOD, ¶ms).await + } + + pub async fn environment_status(&self) -> Result { + // Health checks only reuse an existing RPC connection and never initiate recovery. + let rpc_client = self.rpc_client_without_recovery()?; + self.map_rpc_call_result( + rpc_client + .call_with_timeout(ENVIRONMENT_STATUS_METHOD, &(), ENVIRONMENT_STATUS_TIMEOUT) + .await, + ) + } + + pub async fn discover_capability_roots( + &self, + params: CapabilityRootsDiscoverParams, + ) -> Result { + self.call(CAPABILITY_ROOTS_DISCOVER_METHOD, ¶ms).await + } + + pub async fn read(&self, params: ReadParams) -> Result { + self.call(EXEC_READ_METHOD, ¶ms).await + } + + pub async fn write( + &self, + process_id: &ProcessId, + chunk: Vec, + write_id: String, + ) -> Result { + self.call( + EXEC_WRITE_METHOD, + &WriteParams { + process_id: process_id.clone(), + chunk: chunk.into(), + write_id, + }, + ) + .await + } + + pub async fn signal( + &self, + process_id: &ProcessId, + signal: ProcessSignal, + ) -> Result<(), ExecServerError> { + let _response: SignalResponse = self + .call( + EXEC_SIGNAL_METHOD, + &SignalParams { + process_id: process_id.clone(), + signal, + }, + ) + .await?; + Ok(()) + } + + pub async fn terminate( + &self, + process_id: &ProcessId, + ) -> Result { + // A close notification may arrive before the termination response. + if let Some(session) = self.inner.get_session(process_id) { + session + .network_policy + .cancellation + .record(NetworkRequestCancellationReason::ProcessCancelled); + } + self.call_for_cleanup( + EXEC_TERMINATE_METHOD, + &TerminateParams { + process_id: process_id.clone(), + }, + ) + .await + } + + pub async fn fs_read_file( + &self, + params: FsReadFileParams, + ) -> Result { + self.call(FS_READ_FILE_METHOD, ¶ms).await + } + + pub async fn fs_open(&self, params: FsOpenParams) -> Result { + self.call(FS_OPEN_METHOD, ¶ms).await + } + + pub async fn fs_read_block( + &self, + params: FsReadBlockParams, + ) -> Result { + self.call(FS_READ_BLOCK_METHOD, ¶ms).await + } + + pub async fn fs_close( + &self, + params: FsCloseParams, + ) -> Result { + self.call_for_cleanup(FS_CLOSE_METHOD, ¶ms).await + } + + pub async fn fs_write_file( + &self, + params: FsWriteFileParams, + ) -> Result { + self.call(FS_WRITE_FILE_METHOD, ¶ms).await + } + + pub async fn fs_create_directory( + &self, + params: FsCreateDirectoryParams, + ) -> Result { + self.call(FS_CREATE_DIRECTORY_METHOD, ¶ms).await + } + + pub async fn fs_get_metadata( + &self, + params: FsGetMetadataParams, + ) -> Result { + self.call(FS_GET_METADATA_METHOD, ¶ms).await + } + + pub async fn fs_canonicalize( + &self, + params: FsCanonicalizeParams, + ) -> Result { + self.call(FS_CANONICALIZE_METHOD, ¶ms).await + } + + pub async fn fs_read_directory( + &self, + params: FsReadDirectoryParams, + ) -> Result { + self.call(FS_READ_DIRECTORY_METHOD, ¶ms).await + } + + pub async fn fs_walk(&self, params: FsWalkParams) -> Result { + self.call(FS_WALK_METHOD, ¶ms).await + } + + pub async fn fs_remove( + &self, + params: FsRemoveParams, + ) -> Result { + self.call(FS_REMOVE_METHOD, ¶ms).await + } + + pub async fn fs_copy(&self, params: FsCopyParams) -> Result { + self.call(FS_COPY_METHOD, ¶ms).await + } + + pub(crate) async fn start_process( + &self, + params: ExecParams, + network_policy_decider: Option>, + ) -> Result { + let policy_decision_timeout_ms = params + .network_proxy + .as_ref() + .and_then(|launch| launch.policy_decision_timeout_ms); + let network_policy_controller = match ( + network_policy_decider.as_ref(), + policy_decision_timeout_ms, + ) { + (None, None) => None, + (Some(decider), Some(timeout_ms)) if timeout_ms > 0 => { + Some(NetworkPolicyDecisionController { + decider: Arc::clone(decider), + timeout: Duration::from_millis(timeout_ms), + }) + } + _ => { + return Err(ExecServerError::Protocol( + "network policy decision callback timeout must match the configured decider and be nonzero" + .to_string(), + )); + } + }; + + loop { + let rpc_client = self.rpc_client().await?; + if !self.inner.begin_process_start(&rpc_client) { + continue; + } + + let process_id = params.process_id.clone(); + let mut state = SessionState::new(/*recoverable*/ false); + state.network_policy.audit = + params + .network_proxy + .as_ref() + .map(|launch| NetworkPolicyAuditContext { + metadata: launch.audit_metadata.clone(), + execution_id: launch.execution_id.clone(), + }); + let state = Arc::new(state); + if let Some(controller) = network_policy_controller.as_ref() { + state + .network_policy + .controller + .store(Some(Arc::new(controller.clone()))); + } + if let Err(error) = self.inner.insert_session(&process_id, Arc::clone(&state)) { + self.inner.finish_process_start(); + return Err(error); + } + let active_start = ActiveProcessStart { + inner: Arc::clone(&self.inner), + }; + let mut pending_start = PendingProcessStartSession { + inner: Arc::clone(&self.inner), + process_id: process_id.clone(), + state: Arc::clone(&state), + armed: true, + }; + let client = self.clone(); + let (result_tx, result_rx) = tokio::sync::oneshot::channel(); + let (result_received_tx, result_received_rx) = tokio::sync::oneshot::channel(); + let process_start_task = async move { + let _active_start = active_start; + match client + .call_rpc::<_, ExecResponse>(&rpc_client, EXEC_METHOD, ¶ms) + .await + { + Ok(response) => { + state.recoverable.store(true, Ordering::Release); + let session = Session { + client: client.clone(), + process_id: process_id.clone(), + sandbox_type: response.sandbox_type, + state: Arc::clone(&state), + }; + // Wait for caller receipt so cancellation after send still triggers cleanup. + if result_tx.send(Ok(session)).is_err() || result_received_rx.await.is_err() + { + state.recoverable.store(false, Ordering::Release); + tokio::spawn(async move { + cleanup_process_start(&client, &process_id, &state).await; + }); + } + } + Err(error) => { + if is_transport_closed_error(&error) { + tokio::spawn(async move { + cleanup_process_start(&client, &process_id, &state).await; + }); + } else { + client.inner.remove_session_if(&process_id, &state); + } + let _ = result_tx.send(Err(error)); + } + } + }; + tokio::spawn( + process_start_task + .in_current_span() + .with_current_subscriber(), + ); + let result = result_rx.await; + // The response task may have queued a session before retirement. + if self.inner.retired.is_cancelled() { + return Err(ExecServerError::Disconnected( + "exec-server executor was replaced".to_string(), + )); + } + if matches!(&result, Ok(Ok(_))) { + pending_start.armed = false; + let _ = result_received_tx.send(()); + } + return result.map_err(|_| { + ExecServerError::Protocol("process start task stopped unexpectedly".to_string()) + })?; + } + } + + #[cfg(test)] + pub(crate) async fn register_session( + &self, + process_id: &ProcessId, + ) -> Result { + let state = Arc::new(SessionState::new(/*recoverable*/ true)); + self.inner.insert_session(process_id, Arc::clone(&state))?; + Ok(Session { + client: self.clone(), + process_id: process_id.clone(), + sandbox_type: None, + state, + }) + } + + pub fn session_id(&self) -> Option { + self.inner.session_id.get().cloned() + } + + fn is_disconnected(&self) -> bool { + self.inner.is_failed() + } + + fn readiness_result(&self) -> Option> { + let connection = self + .inner + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + match &connection.status { + ConnectionStatus::Connected(rpc_client) if !rpc_client.is_disconnected() => { + Some(Ok(())) + } + ConnectionStatus::Connected(_) | ConnectionStatus::Recovering => None, + ConnectionStatus::Failed(message) => { + Some(Err(ExecServerError::Disconnected(message.clone()))) + } + } + } + + pub(crate) async fn connect( + connection: JsonRpcConnection, + options: ExecServerClientConnectOptions, + ) -> Result { + Self::connect_with_recovery(connection, options, /*reconnect_strategy*/ None).await + } + + pub(crate) async fn connect_with_recovery( + connection: JsonRpcConnection, + options: ExecServerClientConnectOptions, + reconnect_strategy: Option, + ) -> Result { + Self::connect_with_recovery_inner( + connection, + options, + reconnect_strategy, + /*noise_context*/ None, + ) + .await + } + + pub(crate) async fn connect_with_recovery_and_noise_context( + connection: JsonRpcConnection, + options: ExecServerClientConnectOptions, + reconnect_strategy: Option, + noise_context: NoiseInitializeContext, + ) -> Result { + Self::connect_with_recovery_inner( + connection, + options, + reconnect_strategy, + Some(noise_context), + ) + .await + } + + async fn connect_with_recovery_inner( + connection: JsonRpcConnection, + options: ExecServerClientConnectOptions, + reconnect_strategy: Option, + noise_context: Option, + ) -> Result { + let (rpc_client, events_rx) = RpcClient::new(connection); + let rpc_client = Arc::new(rpc_client); + let session_id = OnceLock::new(); + let (connection_changed, _connection_changed_rx) = watch::channel(()); + let inner = Arc::new(Inner { + connection: StdMutex::new(ConnectionState { + status: ConnectionStatus::Connected(Arc::clone(&rpc_client)), + active_process_starts: 0, + environment_connection_state_tx: watch::channel( + EnvironmentConnectionState::Connected, + ) + .0, + }), + connection_changed, + sessions: ArcSwap::from_pointee(HashMap::new()), + sessions_write_lock: StdMutex::new(()), + http_body_streams: ArcSwap::from_pointee(HashMap::new()), + http_body_stream_failures: ArcSwap::from_pointee(HashMap::new()), + http_body_streams_write_lock: Mutex::new(()), + http_body_stream_byte_budget: Arc::new(Semaphore::new(MAX_QUEUED_HTTP_BODY_BYTES)), + http_body_stream_next_id: AtomicU64::new(1), + rpc_inbound_request_slots: Arc::new(Semaphore::new(MAX_IN_FLIGHT_SERVER_CALLS)), + session_id, + retired: CancellationToken::new(), + environment_info: OnceCell::new(), + reconnect_strategy, + }); + let client = Self { + inner, + recovery_policy: RecoveryPolicy::Wait, + }; + // An explicit resume can redirect notifications from running processes + // before initialize returns. Drain them immediately so a burst cannot + // fill the bounded event channel and block the initialize response. + client.spawn_rpc_reader(&rpc_client, events_rx); + let initialize_response = client + .initialize_rpc(&rpc_client, options, noise_context) + .await?; + if let Some(info) = initialize_response.environment_info { + assert!( + client.inner.environment_info.set(info).is_ok(), + "new client metadata cache must be empty" + ); + } + Ok(client) + } + + async fn call(&self, method: &str, params: &P) -> Result + where + P: serde::Serialize, + T: serde::de::DeserializeOwned, + { + let rpc_client = self.rpc_client().await?; + self.call_rpc(&rpc_client, method, params).await + } + + async fn call_rpc( + &self, + rpc_client: &Arc, + method: &str, + params: &P, + ) -> Result + where + P: serde::Serialize, + T: serde::de::DeserializeOwned, + { + self.map_rpc_call_result(rpc_client.call(method, params).await) + } + + async fn call_for_cleanup(&self, method: &str, params: &P) -> Result + where + P: serde::Serialize, + T: serde::de::DeserializeOwned, + { + let rpc_client = self.inner.rpc_client().await?; + self.map_rpc_call_result(rpc_client.call_for_cleanup(method, params).await) + } + + fn map_rpc_call_result( + &self, + result: Result, + ) -> Result { + // Explicit retirement rejects late responses. Ordinary EOF still preserves + // responses received before disconnect, as ordered by the RPC reader. + if self.inner.retired.is_cancelled() { + return Err(ExecServerError::Disconnected( + "exec-server executor was replaced".to_string(), + )); + } + result.map_err(|error| { + let error = ExecServerError::from(error); + if is_transport_closed_error(&error) { + ExecServerError::Disconnected(disconnected_message(/*reason*/ None)) + } else { + error + } + }) + } +} + +async fn cleanup_process_start( + client: &ExecServerClient, + process_id: &ProcessId, + state: &Arc, +) { + loop { + match client.terminate(process_id).await { + Ok(_) => break, + Err(error) if is_transport_closed_error(&error) && !client.inner.is_failed() => { + continue; + } + Err(_) => break, + } + } + client.inner.remove_session_if(process_id, state); +} + +impl From for ExecServerError { + fn from(value: RpcCallError) -> Self { + match value { + RpcCallError::Closed => Self::Closed, + RpcCallError::Json(err) => Self::Json(err), + RpcCallError::Server(error) => Self::Server { + code: error.code, + message: error.message, + }, + RpcCallError::TimedOut { method, timeout } => Self::Protocol(format!( + "timed out waiting for exec-server `{method}` response after {timeout:?}" + )), + RpcCallError::PendingRequestLimitExceeded { limit } => Self::Protocol(format!( + "exec-server has reached its limit of {limit} pending requests" + )), + } + } +} + +impl SessionState { + fn new(recoverable: bool) -> Self { + let (wake_tx, _wake_rx) = watch::channel(0); + Self { + wake_tx, + events: ExecProcessEventLog::new( + PROCESS_EVENT_CHANNEL_CAPACITY, + PROCESS_EVENT_RETAINED_BYTES, + ), + ordered_events: StdMutex::new(OrderedSessionEvents::default()), + recoverable: AtomicBool::new(recoverable), + next_write_id: AtomicU64::new(1), + network_policy: NetworkPolicyState { + controller: ArcSwapOption::empty(), + cancelled: CancellationToken::new(), + cancellation: NetworkRequestCancellation::default(), + audit: None, + }, + } + } + + pub(crate) fn subscribe(&self) -> watch::Receiver { + self.wake_tx.subscribe() + } + + pub(crate) fn subscribe_events(&self) -> ExecProcessEventReceiver { + self.events.subscribe() + } + + fn note_change(&self, seq: u64) { + self.wake_tx + .send_modify(|current| *current = (*current).max(seq)); + } + + /// Publishes a process event only when all earlier sequenced events have + /// already been published. + /// + /// Returns `true` only when this call actually publishes the ordered + /// `Closed` event. The caller uses that signal to remove the session route + /// after the terminal event is visible to subscribers, rather than when a + /// possibly-early closed notification first arrives. + fn publish_ordered_event(&self, event: ExecProcessEvent) -> Result { + let Some(seq) = event.seq() else { + self.events.publish(event); + return Ok(false); + }; + + let mut ordered_events = self + .ordered_events + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + // We have already delivered this sequence number or moved past it, + // so accepting it again would duplicate output or lifecycle events. + if ordered_events.failure.is_some() + || ordered_events.closed_published + || seq <= ordered_events.last_published_seq + { + return Ok(false); + } + + ordered_events.insert_pending(event)?; + Ok(self.publish_ready(&mut ordered_events)) + } + + fn publish_ready(&self, ordered_events: &mut OrderedSessionEvents) -> bool { + let mut published_closed = false; + loop { + let next_seq = ordered_events.last_published_seq.saturating_add(1); + let Some(event) = ordered_events.pending.remove(&next_seq) else { + break; + }; + ordered_events.pending_bytes = ordered_events + .pending_bytes + .saturating_sub(pending_process_event_bytes(&event)); + ordered_events.last_published_seq = next_seq; + ordered_events.exit_published |= matches!(&event, ExecProcessEvent::Exited { .. }); + let is_closed = matches!(&event, ExecProcessEvent::Closed { .. }); + ordered_events.closed_published |= is_closed; + published_closed |= is_closed; + if is_closed && ordered_events.exit_published { + self.network_policy + .cancellation + .record(NetworkRequestCancellationReason::ProcessFinished); + } + self.events.publish(event); + } + published_closed + } + + fn set_failure(&self, message: String) { + let mut ordered_events = self + .ordered_events + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if ordered_events.failure.is_some() || ordered_events.closed_published { + return; + } + ordered_events.failure = Some(message.clone()); + ordered_events.pending.clear(); + ordered_events.pending_bytes = 0; + self.events.publish(ExecProcessEvent::Failed(message)); + drop(ordered_events); + self.wake_tx + .send_modify(|current| *current = current.saturating_add(1)); + } + + fn failed_response(&self) -> Option { + self.ordered_events + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .failure + .clone() + .map(|message| self.synthesized_failure(message)) + } + + fn synthesized_failure(&self, message: String) -> ReadResponse { + let next_seq = (*self.wake_tx.borrow()).saturating_add(1); + ReadResponse { + chunks: Vec::new(), + next_seq, + exited: true, + exit_code: None, + closed: true, + failure: Some(message), + sandbox_denied: false, + } + } + + fn next_write_id(&self) -> String { + self.next_write_id + .fetch_add(1, Ordering::Relaxed) + .to_string() + } +} + +impl OrderedSessionEvents { + fn insert_pending(&mut self, event: ExecProcessEvent) -> Result<(), String> { + let Some(seq) = event.seq() else { + return Err("cannot reorder an unsequenced process event".to_string()); + }; + if self.pending.contains_key(&seq) { + return Ok(()); + } + + let next_seq = self.last_published_seq.saturating_add(1); + // The next expected event is synchronously published by every caller, + // so it can drain a full buffer without becoming retained state. + let closes_gap = seq == next_seq; + if !closes_gap && self.pending.len() >= MAX_PENDING_PROCESS_EVENTS { + return Err(format!( + "process event reorder buffer exceeds {MAX_PENDING_PROCESS_EVENTS} entries" + )); + } + + let event_bytes = pending_process_event_bytes(&event); + if event_bytes > MAX_PENDING_PROCESS_EVENT_BYTES { + return Err(format!( + "process event exceeds {MAX_PENDING_PROCESS_EVENT_BYTES} bytes" + )); + } + let pending_bytes = self.pending_bytes.saturating_add(event_bytes); + if !closes_gap && pending_bytes > MAX_PENDING_PROCESS_EVENT_BYTES { + return Err(format!( + "process event reorder buffer exceeds {MAX_PENDING_PROCESS_EVENT_BYTES} bytes" + )); + } + + self.pending.insert(seq, event); + self.pending_bytes = pending_bytes; + Ok(()) + } +} + +fn pending_process_event_bytes(event: &ExecProcessEvent) -> usize { + match event { + ExecProcessEvent::Output(chunk) => chunk.chunk.0.len(), + ExecProcessEvent::Failed(message) => message.len(), + ExecProcessEvent::Exited { .. } | ExecProcessEvent::Closed { .. } => 0, + } +} + +fn finish_process_event( + inner: &Inner, + process_id: &ProcessId, + session: &Arc, + result: Result, +) { + match result { + Ok(true) => inner.remove_session_if(process_id, session), + Ok(false) => {} + Err(message) => { + session.set_failure(message); + inner.remove_session_if(process_id, session); + } + } +} + +impl Session { + pub(crate) fn process_id(&self) -> &ProcessId { + &self.process_id + } + + pub(crate) fn sandbox_type(&self) -> Option { + self.sandbox_type + } + + pub(crate) fn subscribe_wake(&self) -> watch::Receiver { + self.state.subscribe() + } + + pub(crate) fn subscribe_events(&self) -> ExecProcessEventReceiver { + self.state.subscribe_events() + } + + pub(crate) async fn read( + &self, + after_seq: Option, + max_bytes: Option, + wait_ms: Option, + ) -> Result { + loop { + if let Some(response) = self.state.failed_response() { + return Ok(response); + } + + match self + .client + .read(ReadParams { + process_id: self.process_id.clone(), + after_seq, + max_bytes, + wait_ms, + }) + .await + { + Ok(response) => return Ok(response), + Err(error) + if is_transport_closed_error(&error) && !self.client.inner.is_failed() => + { + continue; + } + Err(error) if is_transport_closed_error(&error) => { + if let Some(response) = self.state.failed_response() { + return Ok(response); + } + let message = error.to_string(); + self.state.set_failure(message.clone()); + return Ok(self.state.synthesized_failure(message)); + } + Err(error) => return Err(error), + } + } + } + + pub(crate) async fn write(&self, chunk: Vec) -> Result { + let write_id = self.state.next_write_id(); + loop { + match self + .client + .write(&self.process_id, chunk.clone(), write_id.clone()) + .await + { + Ok(response) => return Ok(response), + Err(error) + if is_transport_closed_error(&error) && !self.client.inner.is_failed() => + { + continue; + } + Err(error) => return Err(error), + } + } + } + + pub(crate) async fn signal(&self, signal: ProcessSignal) -> Result<(), ExecServerError> { + self.client.signal(&self.process_id, signal).await + } + + pub(crate) async fn terminate(&self) -> Result<(), ExecServerError> { + self.client.terminate(&self.process_id).await?; + self.cancel_network_policy_decisions(); + Ok(()) + } + + pub(crate) fn cancel_network_policy_decisions(&self) { + self.state + .network_policy + .cancellation + .record(NetworkRequestCancellationReason::ProcessCancelled); + self.state.network_policy.cancelled.cancel(); + } + + pub(crate) async fn unregister(&self) { + self.client + .inner + .remove_session_if(&self.process_id, &self.state); + } +} + +impl Inner { + fn get_session(&self, process_id: &ProcessId) -> Option> { + self.sessions.load().get(process_id).cloned() + } + + fn insert_session( + &self, + process_id: &ProcessId, + session: Arc, + ) -> Result<(), ExecServerError> { + let _sessions_write_guard = self + .sessions_write_lock + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + // Do not register a process session that can never receive environment + // notifications. Without this check, remote MCP startup could create a + // dead session and wait for process output that will never arrive. + if let Some(message) = self.failure_message() { + return Err(ExecServerError::Disconnected(message)); + } + let sessions = self.sessions.load(); + if sessions.contains_key(process_id) { + return Err(ExecServerError::Protocol(format!( + "session already registered for process {process_id}" + ))); + } + let mut next_sessions = sessions.as_ref().clone(); + next_sessions.insert(process_id.clone(), session); + self.sessions.store(Arc::new(next_sessions)); + Ok(()) + } + + fn remove_session_if(&self, process_id: &ProcessId, expected: &Arc) { + let _sessions_write_guard = self + .sessions_write_lock + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let sessions = self.sessions.load(); + if !sessions + .get(process_id) + .is_some_and(|session| Arc::ptr_eq(session, expected)) + { + return; + } + let mut next_sessions = sessions.as_ref().clone(); + next_sessions.remove(process_id); + self.sessions.store(Arc::new(next_sessions)); + expected + .network_policy + .cancellation + .record(NetworkRequestCancellationReason::ProcessCancelled); + expected.network_policy.cancelled.cancel(); + expected.network_policy.controller.store(None); + } + + fn take_all_sessions(&self) -> HashMap> { + let _sessions_write_guard = self + .sessions_write_lock + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let sessions = self.sessions.load(); + let drained_sessions = sessions.as_ref().clone(); + self.sessions.store(Arc::new(HashMap::new())); + drained_sessions + } +} + +fn disconnected_message(reason: Option<&str>) -> String { + match reason { + Some(reason) => format!("exec-server transport disconnected: {reason}"), + None => "exec-server transport disconnected".to_string(), + } +} + +fn is_transport_closed_error(error: &ExecServerError) -> bool { + matches!( + error, + ExecServerError::Closed | ExecServerError::Disconnected(_) + ) || matches!( + error, + ExecServerError::Server { + code: -32000, + message, + } if message == "JSON-RPC transport closed" + ) +} + +fn fail_all_sessions(inner: &Arc, message: String) { + let sessions = inner.take_all_sessions(); + + for (_, session) in sessions { + session + .network_policy + .cancellation + .record(NetworkRequestCancellationReason::ConnectionClosed); + session.network_policy.cancelled.cancel(); + session.network_policy.controller.store(None); + // Sessions synthesize a closed read response and emit a pushed Failed + // event. That covers both polling consumers and streaming consumers + // such as environment-backed MCP stdio. + session.set_failure(message.clone()); + } +} + +/// Fails all in-flight work that depends on the shared JSON-RPC transport. +async fn fail_all_in_flight_work(inner: &Arc, message: String) { + fail_all_sessions(inner, message.clone()); + inner.fail_all_http_body_streams(message).await; +} + +async fn handle_server_notification( + inner: &Arc, + notification: JSONRPCNotification, +) -> Result<(), ExecServerError> { + match notification.method.as_str() { + EXEC_OUTPUT_DELTA_METHOD => { + let params: ExecOutputDeltaNotification = + serde_json::from_value(notification.params.unwrap_or(Value::Null))?; + if let Some(session) = inner.get_session(¶ms.process_id) { + let result = + session.publish_ordered_event(ExecProcessEvent::Output(ProcessOutputChunk { + seq: params.seq, + stream: params.stream, + chunk: params.chunk, + })); + if result.is_ok() { + session.note_change(params.seq); + } + finish_process_event(inner, ¶ms.process_id, &session, result); + } + } + EXEC_EXITED_METHOD => { + let params: ExecExitedNotification = + serde_json::from_value(notification.params.unwrap_or(Value::Null))?; + if let Some(session) = inner.get_session(¶ms.process_id) { + let result = session.publish_ordered_event(ExecProcessEvent::Exited { + seq: params.seq, + exit_code: params.exit_code, + sandbox_denied: params.sandbox_denied, + }); + if result.is_ok() { + session.note_change(params.seq); + } + finish_process_event(inner, ¶ms.process_id, &session, result); + } + } + EXEC_CLOSED_METHOD => { + let params: ExecClosedNotification = + serde_json::from_value(notification.params.unwrap_or(Value::Null))?; + if let Some(session) = inner.get_session(¶ms.process_id) { + // Closed is terminal, but it can arrive before tail output or + // exited. Keep routing this process until the ordered publisher + // says Closed has actually been delivered. + let result = + session.publish_ordered_event(ExecProcessEvent::Closed { seq: params.seq }); + if result.is_ok() { + session.note_change(params.seq); + } + finish_process_event(inner, ¶ms.process_id, &session, result); + } + } + HTTP_REQUEST_BODY_DELTA_METHOD => { + inner + .handle_http_body_delta_notification(notification.params) + .await?; + } + NETWORK_POLICY_DECISION_METHOD => { + let Ok(params) = serde_json::from_value::( + notification.params.unwrap_or(Value::Null), + ) else { + debug!("ignoring malformed exec-server network policy decision notification"); + return Ok(()); + }; + let Some(session) = inner.get_session(¶ms.process_id) else { + debug!("ignoring network policy decision for an unknown exec-server process"); + return Ok(()); + }; + let Some(context) = session.network_policy.audit.as_ref() else { + return Ok(()); + }; + if !network_policy_audit::emit_network_policy_decision(context, ¶ms) { + debug!("ignoring invalid exec-server network policy decision notification"); + } + } + other => { + debug!("ignoring unknown exec-server notification: {other}"); + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use codex_exec_server_protocol::JSONRPCMessage; + use codex_exec_server_protocol::JSONRPCNotification; + use codex_exec_server_protocol::JSONRPCResponse; + use codex_http_client::HttpClientFactory; + use codex_http_client::OutboundProxyPolicy; + use codex_utils_path_uri::PathUri; + use futures::SinkExt; + use futures::StreamExt; + use http::HeaderMap; + use opentelemetry::trace::TracerProvider as _; + use opentelemetry_sdk::trace::SdkTracerProvider; + use pretty_assertions::assert_eq; + use std::collections::HashMap; + #[cfg(unix)] + use std::path::Path; + #[cfg(unix)] + use std::process::Command; + use std::sync::Arc; + use tokio::io::AsyncBufReadExt; + use tokio::io::AsyncWrite; + use tokio::io::AsyncWriteExt; + use tokio::io::BufReader; + use tokio::io::duplex; + use tokio::net::TcpListener; + use tokio::net::TcpStream; + use tokio::sync::mpsc; + use tokio::sync::oneshot; + use tokio::sync::watch; + use tokio::time::Duration; + #[cfg(unix)] + use tokio::time::sleep; + use tokio::time::timeout; + use tokio_tungstenite::WebSocketStream; + use tokio_tungstenite::accept_async; + use tokio_tungstenite::tungstenite::Message; + use tracing::Instrument; + use tracing_subscriber::filter::filter_fn; + use tracing_subscriber::prelude::*; + + use super::ExecServerClient; + use super::ExecServerClientConnectOptions; + use super::LazyRemoteExecServerClient; + use crate::EnvironmentObservedStatus; + use crate::ProcessId; + #[cfg(not(windows))] + use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT; + use crate::client_api::ExecServerTransportParams; + use crate::client_api::RemoteExecServerConnectArgs; + use crate::client_api::StdioExecServerCommand; + use crate::client_api::StdioExecServerConnectArgs; + use crate::connection::JsonRpcConnection; + use crate::process::ExecProcessEvent; + use crate::protocol::EXEC_CLOSED_METHOD; + use crate::protocol::EXEC_EXITED_METHOD; + use crate::protocol::EXEC_METHOD; + use crate::protocol::EXEC_OUTPUT_DELTA_METHOD; + use crate::protocol::EXEC_READ_METHOD; + use crate::protocol::EXEC_WRITE_METHOD; + use crate::protocol::EnvironmentInfo; + use crate::protocol::ExecClosedNotification; + use crate::protocol::ExecExitedNotification; + use crate::protocol::ExecOutputDeltaNotification; + use crate::protocol::ExecOutputStream; + use crate::protocol::ExecParams; + use crate::protocol::ExecResponse; + use crate::protocol::INITIALIZE_METHOD; + use crate::protocol::INITIALIZED_METHOD; + use crate::protocol::InitializeResponse; + use crate::protocol::ProcessOutputChunk; + use crate::protocol::ProcessSandboxType; + use crate::protocol::ReadResponse; + use crate::protocol::WriteParams; + use crate::protocol::WriteResponse; + use crate::protocol::WriteStatus; + + async fn read_jsonrpc_line(lines: &mut tokio::io::Lines>) -> JSONRPCMessage + where + R: tokio::io::AsyncRead + Unpin, + { + let line = timeout(Duration::from_secs(1), lines.next_line()) + .await + .expect("json-rpc read should not time out") + .expect("json-rpc read should succeed") + .expect("json-rpc connection should stay open"); + serde_json::from_str(&line).expect("json-rpc line should parse") + } + + async fn write_jsonrpc_line(writer: &mut W, message: JSONRPCMessage) + where + W: AsyncWrite + Unpin, + { + let encoded = serde_json::to_string(&message).expect("json-rpc message should serialize"); + writer + .write_all(format!("{encoded}\n").as_bytes()) + .await + .expect("json-rpc line should write"); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] + async fn process_start_propagates_caller_trace_context_across_background_task() { + let (client_stdin, server_reader) = duplex(1 << 20); + let (mut server_writer, client_stdout) = duplex(1 << 20); + let server = tokio::spawn(async move { + let mut lines = BufReader::new(server_reader).lines(); + let initialize = read_jsonrpc_line(&mut lines).await; + let initialize = match initialize { + JSONRPCMessage::Request(request) if request.method == INITIALIZE_METHOD => request, + other => panic!("expected initialize request, got {other:?}"), + }; + write_jsonrpc_line( + &mut server_writer, + JSONRPCMessage::Response(JSONRPCResponse { + id: initialize.id, + result: serde_json::to_value(InitializeResponse { + session_id: "trace-test".to_string(), + environment_info: None, + }) + .expect("initialize response should serialize"), + }), + ) + .await; + + match read_jsonrpc_line(&mut lines).await { + JSONRPCMessage::Notification(notification) + if notification.method == INITIALIZED_METHOD => {} + other => panic!("expected initialized notification, got {other:?}"), + } + + let request = match read_jsonrpc_line(&mut lines).await { + JSONRPCMessage::Request(request) if request.method == EXEC_METHOD => request, + other => panic!("expected process start request, got {other:?}"), + }; + let trace = request.trace.clone(); + let params: ExecParams = + serde_json::from_value(request.params.expect("process start params should exist")) + .expect("process start params should deserialize"); + write_jsonrpc_line( + &mut server_writer, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(ExecResponse { + process_id: params.process_id, + sandbox_type: Some(ProcessSandboxType::LinuxSeccomp), + }) + .expect("process start response should serialize"), + }), + ) + .await; + trace + }); + + let client = ExecServerClient::connect( + JsonRpcConnection::from_stdio( + client_stdout, + client_stdin, + "trace-test-client".to_string(), + ), + ExecServerClientConnectOptions::default(), + ) + .await + .expect("client should connect"); + + let tracer_provider = SdkTracerProvider::builder().build(); + let tracer = tracer_provider.tracer("exec-server-test"); + let subscriber = tracing_subscriber::registry().with( + tracing_opentelemetry::layer() + .with_tracer(tracer) + .with_filter(filter_fn(codex_otel::OtelProvider::trace_export_filter)), + ); + let _subscriber_guard = tracing::subscriber::set_default(subscriber); + tracing::callsite::rebuild_interest_cache(); + let parent_span = tracing::info_span!("process-start-parent"); + let expected_trace = codex_otel::span_w3c_trace_context(&parent_span) + .expect("parent span should have trace context"); + let process_id = ProcessId::from("trace-process"); + + let session = client + .start_process( + ExecParams { + metadata: Default::default(), + process_id: process_id.clone(), + argv: vec!["true".to_string()], + cwd: PathUri::from_host_native_path(std::env::current_dir().expect("cwd")) + .expect("cwd URI"), + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }, + /*network_policy_decider*/ None, + ) + .instrument(parent_span) + .await + .expect("process start should succeed"); + + assert_eq!(session.process_id(), &process_id); + assert_eq!( + session.sandbox_type(), + Some(ProcessSandboxType::LinuxSeccomp) + ); + let trace = server.await.expect("server task").expect("trace context"); + let expected_traceparent = expected_trace + .traceparent + .as_deref() + .expect("parent traceparent"); + let traceparent = trace.traceparent.as_deref().expect("request traceparent"); + let expected_parts = expected_traceparent.split('-').collect::>(); + let parts = traceparent.split('-').collect::>(); + assert_eq!(parts[1], expected_parts[1]); + assert_ne!(parts[2], expected_parts[2]); + assert_eq!(trace.tracestate, expected_trace.tracestate); + } + + async fn accept_websocket(listener: &TcpListener) -> WebSocketStream { + let (stream, _) = listener.accept().await.expect("listener should accept"); + accept_async(stream) + .await + .expect("websocket handshake should succeed") + } + + async fn read_jsonrpc_websocket(websocket: &mut WebSocketStream) -> JSONRPCMessage { + loop { + match timeout(Duration::from_secs(1), websocket.next()) + .await + .expect("json-rpc 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("json-rpc text frame should parse"); + } + Message::Binary(bytes) => { + return serde_json::from_slice(bytes.as_ref()) + .expect("json-rpc binary frame should parse"); + } + Message::Ping(_) | Message::Pong(_) => {} + other => panic!("expected json-rpc websocket frame, got {other:?}"), + } + } + } + + async fn write_jsonrpc_websocket( + websocket: &mut WebSocketStream, + message: JSONRPCMessage, + ) { + let encoded = serde_json::to_string(&message).expect("json-rpc should serialize"); + websocket + .send(Message::Text(encoded.into())) + .await + .expect("json-rpc websocket frame should write"); + } + + async fn complete_websocket_initialize( + websocket: &mut WebSocketStream, + session_id: &str, + expected_resume_session_id: Option<&str>, + ) { + complete_websocket_initialize_with_environment_info( + websocket, + session_id, + expected_resume_session_id, + /*environment_info*/ None, + ) + .await; + } + + async fn complete_websocket_initialize_with_environment_info( + websocket: &mut WebSocketStream, + session_id: &str, + expected_resume_session_id: Option<&str>, + environment_info: Option, + ) { + let initialize = read_jsonrpc_websocket(websocket).await; + let request = match initialize { + JSONRPCMessage::Request(request) if request.method == INITIALIZE_METHOD => request, + other => panic!("expected initialize request, got {other:?}"), + }; + let params: crate::protocol::InitializeParams = + serde_json::from_value(request.params.expect("initialize params should exist")) + .expect("initialize params should deserialize"); + assert_eq!( + params.resume_session_id.as_deref(), + expected_resume_session_id + ); + write_jsonrpc_websocket( + websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(InitializeResponse { + session_id: session_id.to_string(), + environment_info, + }) + .expect("initialize response should serialize"), + }), + ) + .await; + + let initialized = read_jsonrpc_websocket(websocket).await; + match initialized { + JSONRPCMessage::Notification(notification) + if notification.method == INITIALIZED_METHOD => {} + other => panic!("expected initialized notification, got {other:?}"), + } + } + + #[cfg(not(windows))] + #[tokio::test] + async fn connect_stdio_command_initializes_json_rpc_client() { + let client = ExecServerClient::connect_stdio_command(StdioExecServerConnectArgs { + command: StdioExecServerCommand { + program: "sh".to_string(), + args: vec![ + "-c".to_string(), + "read _line; printf '%s\\n' '{\"id\":1,\"result\":{\"sessionId\":\"stdio-test\"}}'; read _line; sleep 60".to_string(), + ], + env: HashMap::new(), + cwd: None, + }, + client_name: "stdio-test-client".to_string(), + initialize_timeout: Duration::from_secs(1), + resume_session_id: None, + }) + .await + .expect("stdio client should connect"); + + assert_eq!(client.session_id().as_deref(), Some("stdio-test")); + } + + #[cfg(not(windows))] + #[tokio::test] + async fn connect_for_transport_initializes_stdio_command() { + let client = ExecServerClient::connect_for_transport( + ExecServerTransportParams::StdioCommand { + command: StdioExecServerCommand { + program: "sh".to_string(), + args: vec![ + "-c".to_string(), + "read _line; printf '%s\\n' '{\"id\":1,\"result\":{\"sessionId\":\"stdio-test\"}}'; read _line; sleep 60".to_string(), + ], + env: HashMap::new(), + cwd: None, + }, + initialize_timeout: DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT, + }, + codex_http_client::HttpClientFactory::new( + codex_http_client::OutboundProxyPolicy::ReqwestDefault, + ), + ) + .await + .expect("stdio transport should connect"); + + assert_eq!(client.session_id().as_deref(), Some("stdio-test")); + } + + #[cfg(windows)] + #[tokio::test] + async fn connect_stdio_command_initializes_json_rpc_client_on_windows() { + let client = ExecServerClient::connect_stdio_command(StdioExecServerConnectArgs { + command: StdioExecServerCommand { + program: "powershell".to_string(), + args: vec![ + "-NoProfile".to_string(), + "-Command".to_string(), + "$null = [Console]::In.ReadLine(); [Console]::Out.WriteLine('{\"id\":1,\"result\":{\"sessionId\":\"stdio-test\"}}'); $null = [Console]::In.ReadLine(); Start-Sleep -Seconds 60".to_string(), + ], + env: HashMap::new(), + cwd: None, + }, + client_name: "stdio-test-client".to_string(), + initialize_timeout: Duration::from_secs(1), + resume_session_id: None, + }) + .await + .expect("stdio client should connect"); + + assert_eq!(client.session_id().as_deref(), Some("stdio-test")); + } + + #[cfg(unix)] + #[tokio::test] + async fn dropping_stdio_client_terminates_spawned_process() { + let tempdir = tempfile::tempdir().expect("tempdir should be created"); + let pid_file = tempdir.path().join("server.pid"); + let child_pid_file = tempdir.path().join("server-child.pid"); + let stdio_script = format!( + "read _line; \ + echo \"$$\" > {}; \ + sleep 60 >/dev/null 2>&1 & echo \"$!\" > {}; \ + printf '%s\\n' '{{\"id\":1,\"result\":{{\"sessionId\":\"stdio-test\"}}}}'; \ + read _line; \ + wait", + shell_quote(pid_file.as_path()), + shell_quote(child_pid_file.as_path()), + ); + + let client = ExecServerClient::connect_stdio_command(StdioExecServerConnectArgs { + command: StdioExecServerCommand { + program: "sh".to_string(), + args: vec!["-c".to_string(), stdio_script], + env: HashMap::new(), + cwd: None, + }, + client_name: "stdio-test-client".to_string(), + initialize_timeout: Duration::from_secs(1), + resume_session_id: None, + }) + .await + .expect("stdio client should connect"); + let server_pid = read_pid_file(pid_file.as_path()).await; + let child_pid = read_pid_file(child_pid_file.as_path()).await; + assert!( + process_exists(server_pid), + "spawned stdio process should be running before client drop" + ); + assert!( + process_exists(child_pid), + "spawned stdio child process should be running before client drop" + ); + + drop(client); + + wait_for_process_exit(server_pid).await; + wait_for_process_exit(child_pid).await; + } + + #[cfg(unix)] + #[tokio::test] + async fn malformed_stdio_message_terminates_spawned_process() { + let tempdir = tempfile::tempdir().expect("tempdir should be created"); + let pid_file = tempdir.path().join("server.pid"); + let stdio_script = format!( + "read _line; \ + echo \"$$\" > {}; \ + printf '%s\\n' 'not-json'; \ + sleep 60", + shell_quote(pid_file.as_path()), + ); + + let result = ExecServerClient::connect_stdio_command(StdioExecServerConnectArgs { + command: StdioExecServerCommand { + program: "sh".to_string(), + args: vec!["-c".to_string(), stdio_script], + env: HashMap::new(), + cwd: None, + }, + client_name: "stdio-test-client".to_string(), + initialize_timeout: Duration::from_secs(1), + resume_session_id: None, + }) + .await; + assert!(result.is_err(), "malformed stdio server should not connect"); + + let server_pid = read_pid_file(pid_file.as_path()).await; + wait_for_process_exit(server_pid).await; + } + + #[cfg(unix)] + async fn read_pid_file(path: &Path) -> u32 { + for _ in 0..20 { + if let Ok(contents) = std::fs::read_to_string(path) { + return contents + .trim() + .parse() + .expect("pid file should contain a pid"); + } + sleep(Duration::from_millis(50)).await; + } + panic!("pid file {} should be written", path.display()); + } + + #[cfg(unix)] + async fn wait_for_process_exit(pid: u32) { + for _ in 0..20 { + if !process_exists(pid) { + return; + } + sleep(Duration::from_millis(100)).await; + } + panic!("process {pid} should exit"); + } + + #[cfg(unix)] + fn process_exists(pid: u32) -> bool { + Command::new("kill") + .arg("-0") + .arg(pid.to_string()) + .status() + .is_ok_and(|status| status.success()) + } + + #[cfg(unix)] + fn shell_quote(path: &Path) -> String { + let value = path.to_string_lossy(); + format!("'{}'", value.replace('\'', "'\\''")) + } + + #[tokio::test] + async fn process_events_are_delivered_in_seq_order_when_notifications_are_reordered() { + let (client_stdin, server_reader) = duplex(1 << 20); + let (mut server_writer, client_stdout) = duplex(1 << 20); + let (notifications_tx, mut notifications_rx) = mpsc::channel(16); + let server = tokio::spawn(async move { + let mut lines = BufReader::new(server_reader).lines(); + let initialize = read_jsonrpc_line(&mut lines).await; + let request = match initialize { + JSONRPCMessage::Request(request) if request.method == INITIALIZE_METHOD => request, + other => panic!("expected initialize request, got {other:?}"), + }; + write_jsonrpc_line( + &mut server_writer, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(InitializeResponse { + session_id: "session-1".to_string(), + environment_info: None, + }) + .expect("initialize response should serialize"), + }), + ) + .await; + + let initialized = read_jsonrpc_line(&mut lines).await; + match initialized { + JSONRPCMessage::Notification(notification) + if notification.method == INITIALIZED_METHOD => {} + other => panic!("expected initialized notification, got {other:?}"), + } + + while let Some(message) = notifications_rx.recv().await { + write_jsonrpc_line(&mut server_writer, message).await; + } + }); + + let client = ExecServerClient::connect( + JsonRpcConnection::from_stdio( + client_stdout, + client_stdin, + "test-exec-server-client".to_string(), + ), + ExecServerClientConnectOptions::default(), + ) + .await + .expect("client should connect"); + + let process_id = ProcessId::from("reordered"); + let session = client + .register_session(&process_id) + .await + .expect("session should register"); + let mut events = session.subscribe_events(); + + for message in [ + JSONRPCMessage::Notification(JSONRPCNotification { + method: EXEC_CLOSED_METHOD.to_string(), + params: Some( + serde_json::to_value(ExecClosedNotification { + process_id: process_id.clone(), + seq: 4, + }) + .expect("closed notification should serialize"), + ), + }), + JSONRPCMessage::Notification(JSONRPCNotification { + method: EXEC_OUTPUT_DELTA_METHOD.to_string(), + params: Some( + serde_json::to_value(ExecOutputDeltaNotification { + process_id: process_id.clone(), + seq: 1, + stream: ExecOutputStream::Stdout, + chunk: b"one".to_vec().into(), + }) + .expect("output notification should serialize"), + ), + }), + JSONRPCMessage::Notification(JSONRPCNotification { + method: EXEC_EXITED_METHOD.to_string(), + params: Some( + serde_json::to_value(ExecExitedNotification { + process_id: process_id.clone(), + seq: 3, + exit_code: 0, + sandbox_denied: Some(true), + }) + .expect("exit notification should serialize"), + ), + }), + JSONRPCMessage::Notification(JSONRPCNotification { + method: EXEC_OUTPUT_DELTA_METHOD.to_string(), + params: Some( + serde_json::to_value(ExecOutputDeltaNotification { + process_id: process_id.clone(), + seq: 2, + stream: ExecOutputStream::Stderr, + chunk: b"two".to_vec().into(), + }) + .expect("output notification should serialize"), + ), + }), + ] { + notifications_tx + .send(message) + .await + .expect("notification should queue"); + } + + let mut delivered = Vec::new(); + for _ in 0..4 { + delivered.push( + timeout(Duration::from_secs(1), events.recv()) + .await + .expect("process event should not time out") + .expect("process event stream should stay open"), + ); + } + + assert_eq!( + delivered, + vec![ + ExecProcessEvent::Output(ProcessOutputChunk { + seq: 1, + stream: ExecOutputStream::Stdout, + chunk: b"one".to_vec().into(), + }), + ExecProcessEvent::Output(ProcessOutputChunk { + seq: 2, + stream: ExecOutputStream::Stderr, + chunk: b"two".to_vec().into(), + }), + ExecProcessEvent::Exited { + seq: 3, + exit_code: 0, + sandbox_denied: Some(true), + }, + ExecProcessEvent::Closed { seq: 4 }, + ] + ); + + drop(notifications_tx); + drop(client); + server.await.expect("server task should finish"); + } + + #[tokio::test] + async fn transport_disconnect_fails_sessions_and_rejects_new_sessions() { + let (client_stdin, server_reader) = duplex(1 << 20); + let (mut server_writer, client_stdout) = duplex(1 << 20); + let (disconnect_tx, disconnect_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let mut lines = BufReader::new(server_reader).lines(); + let initialize = read_jsonrpc_line(&mut lines).await; + let request = match initialize { + JSONRPCMessage::Request(request) if request.method == INITIALIZE_METHOD => request, + other => panic!("expected initialize request, got {other:?}"), + }; + write_jsonrpc_line( + &mut server_writer, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(InitializeResponse { + session_id: "session-1".to_string(), + environment_info: None, + }) + .expect("initialize response should serialize"), + }), + ) + .await; + + let initialized = read_jsonrpc_line(&mut lines).await; + match initialized { + JSONRPCMessage::Notification(notification) + if notification.method == INITIALIZED_METHOD => {} + other => panic!("expected initialized notification, got {other:?}"), + } + + let _ = disconnect_rx.await; + drop(server_writer); + }); + + let client = ExecServerClient::connect( + JsonRpcConnection::from_stdio( + client_stdout, + client_stdin, + "test-exec-server-client".to_string(), + ), + ExecServerClientConnectOptions::default(), + ) + .await + .expect("client should connect"); + + let process_id = ProcessId::from("disconnect"); + let session = client + .register_session(&process_id) + .await + .expect("session should register"); + let mut events = session.subscribe_events(); + + disconnect_tx.send(()).expect("disconnect should signal"); + + let event = timeout(Duration::from_secs(1), events.recv()) + .await + .expect("session failure should not time out") + .expect("session event stream should stay open"); + let ExecProcessEvent::Failed(message) = event else { + panic!("expected session failure after disconnect, got {event:?}"); + }; + assert_eq!(message, "exec-server transport disconnected"); + + let response = session + .read( + /*after_seq*/ None, /*max_bytes*/ None, /*wait_ms*/ None, + ) + .await + .expect("disconnected session read should synthesize a response"); + assert_eq!( + response.failure.as_deref(), + Some("exec-server transport disconnected") + ); + assert!(response.closed); + + let new_session = client.register_session(&ProcessId::from("new")).await; + assert!(matches!( + new_session, + Err(super::ExecServerError::Disconnected(_)) + )); + + drop(client); + server.await.expect("server task should finish"); + } + + #[test_case::test_case(Some(EnvironmentInfo::local()); "from_initialize")] + #[test_case::test_case(Some(EnvironmentInfo { + executor_version: "1.2.3-alpha.4".to_string(), + provider_id: Some("sha256:fb4f62da3e84f6864dcec8ede7bc66f1c96ecaeaf55f8a786b85df994057c8ac".to_string()), + ..EnvironmentInfo::local() + }); "with_executor_metadata")] + #[test_case::test_case(None; "legacy_server")] + #[tokio::test] + async fn environment_info_is_cached( + initial_environment_info: Option, + ) -> anyhow::Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let expected_info = initial_environment_info + .clone() + .unwrap_or_else(EnvironmentInfo::local); + let server_info = expected_info.clone(); + let server = tokio::spawn(async move { + let mut websocket = accept_websocket(&listener).await; + complete_websocket_initialize_with_environment_info( + &mut websocket, + "session-1", + /*expected_resume_session_id*/ None, + initial_environment_info.clone(), + ) + .await; + if initial_environment_info.is_none() { + let JSONRPCMessage::Request(request) = read_jsonrpc_websocket(&mut websocket).await + else { + panic!("expected environment info request"); + }; + assert_eq!(request.method, "environment/info"); + write_jsonrpc_websocket( + &mut websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(server_info) + .expect("environment info should serialize"), + }), + ) + .await; + } + }); + let client = ExecServerClient::connect_websocket(RemoteExecServerConnectArgs { + websocket_url, + client_name: "metadata-test".to_string(), + connect_timeout: Duration::from_secs(1), + initialize_timeout: Duration::from_secs(1), + resume_session_id: None, + http_client_factory: HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + }) + .await?; + + assert_eq!(client.environment_info().await?, expected_info); + server.await?; + // The server is gone, so a cloned client must use the shared cache. + assert_eq!(client.clone().environment_info().await?, expected_info); + Ok(()) + } + + #[tokio::test] + async fn remote_websocket_client_resumes_session() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!( + "ws://{}", + listener.local_addr().expect("listener should have address") + ); + let (resumed_tx, resumed_rx) = oneshot::channel(); + let (finish_tx, finish_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let mut first = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut first, + "session-1", + /*expected_resume_session_id*/ None, + ) + .await; + first.close(None).await.expect("websocket should close"); + + let mut resumed = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut resumed, + "session-1", + /*expected_resume_session_id*/ Some("session-1"), + ) + .await; + resumed_tx.send(()).expect("resume should signal"); + finish_rx.await.expect("test should finish"); + }); + + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::WebSocketUrl { + websocket_url, + connect_timeout: Duration::from_secs(1), + initialize_timeout: Duration::from_secs(1), + http_headers: HeaderMap::new(), + }, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ); + let stable_client = client.get().await.expect("client should connect"); + timeout(Duration::from_secs(1), resumed_rx) + .await + .expect("session resume should not time out") + .expect("session resume should signal"); + let reused_client = client.get().await.expect("client should stay connected"); + assert_eq!(stable_client.session_id().as_deref(), Some("session-1")); + assert!(Arc::ptr_eq(&stable_client.inner, &reused_client.inner)); + finish_tx.send(()).expect("test should finish"); + server.await.expect("server task should finish"); + } + + #[tokio::test] + async fn session_write_retries_same_write_id_after_recovery() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!( + "ws://{}", + listener.local_addr().expect("listener should have address") + ); + let (finish_tx, finish_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let mut first = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut first, + "session-1", + /*expected_resume_session_id*/ None, + ) + .await; + + let first_write = read_jsonrpc_websocket(&mut first).await; + let first_write = match first_write { + JSONRPCMessage::Request(request) if request.method == EXEC_WRITE_METHOD => request, + other => panic!("expected first process/write request, got {other:?}"), + }; + let first_write_params: WriteParams = + serde_json::from_value(first_write.params.expect("write params should exist")) + .expect("write params should deserialize"); + assert_eq!(first_write_params.process_id.as_str(), "proc-write"); + assert_eq!(first_write_params.chunk.into_inner(), b"hello\n".to_vec()); + let write_id = first_write_params.write_id; + assert!(!write_id.is_empty()); + drop(first); + + let mut resumed = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut resumed, + "session-1", + /*expected_resume_session_id*/ Some("session-1"), + ) + .await; + + let recovery_read = read_jsonrpc_websocket(&mut resumed).await; + let recovery_read = match recovery_read { + JSONRPCMessage::Request(request) if request.method == EXEC_READ_METHOD => request, + other => panic!("expected recovery process/read request, got {other:?}"), + }; + write_jsonrpc_websocket( + &mut resumed, + JSONRPCMessage::Response(JSONRPCResponse { + id: recovery_read.id, + result: serde_json::to_value(ReadResponse { + chunks: Vec::new(), + next_seq: 1, + exited: false, + exit_code: None, + closed: false, + failure: None, + sandbox_denied: false, + }) + .expect("read response should serialize"), + }), + ) + .await; + + let retried_write = read_jsonrpc_websocket(&mut resumed).await; + let retried_write = match retried_write { + JSONRPCMessage::Request(request) if request.method == EXEC_WRITE_METHOD => request, + other => panic!("expected retried process/write request, got {other:?}"), + }; + let retried_write_params: WriteParams = + serde_json::from_value(retried_write.params.expect("write params should exist")) + .expect("write params should deserialize"); + assert_eq!(retried_write_params.process_id.as_str(), "proc-write"); + assert_eq!(retried_write_params.chunk.into_inner(), b"hello\n".to_vec()); + assert_eq!(retried_write_params.write_id, write_id); + write_jsonrpc_websocket( + &mut resumed, + JSONRPCMessage::Response(JSONRPCResponse { + id: retried_write.id, + result: serde_json::to_value(WriteResponse { + status: WriteStatus::Accepted, + }) + .expect("write response should serialize"), + }), + ) + .await; + + finish_rx.await.expect("test should finish"); + }); + + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::WebSocketUrl { + websocket_url, + connect_timeout: Duration::from_secs(1), + initialize_timeout: Duration::from_secs(1), + http_headers: HeaderMap::new(), + }, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ); + let stable_client = client.get().await.expect("client should connect"); + let session = stable_client + .register_session(&ProcessId::from("proc-write")) + .await + .expect("session should register"); + + let response = timeout(Duration::from_secs(2), session.write(b"hello\n".to_vec())) + .await + .expect("write should not time out") + .expect("write should recover"); + assert_eq!( + response, + WriteResponse { + status: WriteStatus::Accepted + } + ); + + finish_tx.send(()).expect("test should finish"); + server.await.expect("server task should finish"); + } + + #[tokio::test] + async fn explicit_resume_drains_notifications_before_initialize_response() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!( + "ws://{}", + listener.local_addr().expect("listener should have address") + ); + let (initialized_tx, initialized_rx) = oneshot::channel(); + let (finish_tx, finish_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let mut websocket = accept_websocket(&listener).await; + let initialize = read_jsonrpc_websocket(&mut websocket).await; + let request = match initialize { + JSONRPCMessage::Request(request) if request.method == INITIALIZE_METHOD => request, + other => panic!("expected initialize request, got {other:?}"), + }; + let params: crate::protocol::InitializeParams = + serde_json::from_value(request.params.expect("initialize params should exist")) + .expect("initialize params should deserialize"); + assert_eq!(params.resume_session_id.as_deref(), Some("session-1")); + + for seq in 1..=256 { + write_jsonrpc_websocket( + &mut websocket, + JSONRPCMessage::Notification(JSONRPCNotification { + method: EXEC_OUTPUT_DELTA_METHOD.to_string(), + params: Some( + serde_json::to_value(ExecOutputDeltaNotification { + process_id: ProcessId::from("busy-process"), + seq, + stream: ExecOutputStream::Stdout, + chunk: b"output".to_vec().into(), + }) + .expect("output notification should serialize"), + ), + }), + ) + .await; + } + write_jsonrpc_websocket( + &mut websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(InitializeResponse { + session_id: "session-1".to_string(), + environment_info: None, + }) + .expect("initialize response should serialize"), + }), + ) + .await; + + let initialized = read_jsonrpc_websocket(&mut websocket).await; + match initialized { + JSONRPCMessage::Notification(notification) + if notification.method == INITIALIZED_METHOD => {} + other => panic!("expected initialized notification, got {other:?}"), + } + initialized_tx + .send(()) + .expect("initialized notification should signal"); + finish_rx.await.expect("test should finish"); + }); + + let client = timeout( + Duration::from_secs(1), + ExecServerClient::connect_websocket(RemoteExecServerConnectArgs { + websocket_url, + client_name: "test-client".to_string(), + connect_timeout: Duration::from_secs(1), + initialize_timeout: Duration::from_secs(1), + resume_session_id: Some("session-1".to_string()), + http_client_factory: codex_http_client::HttpClientFactory::new( + codex_http_client::OutboundProxyPolicy::ReqwestDefault, + ), + }), + ) + .await + .expect("explicit resume should not time out") + .expect("explicit resume should connect"); + assert_eq!(client.session_id().as_deref(), Some("session-1")); + + timeout(Duration::from_secs(1), initialized_rx) + .await + .expect("initialized notification should not time out") + .expect("initialized notification should signal"); + finish_tx.send(()).expect("test should finish"); + server.await.expect("server task should finish"); + } + + #[tokio::test] + async fn initial_connection_is_shared_by_all_waiters() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!( + "ws://{}", + listener.local_addr().expect("listener should have address") + ); + let server = tokio::spawn(async move { + let mut connection = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut connection, + "startup-session", + /*expected_resume_session_id*/ None, + ) + .await; + timeout(Duration::from_secs(1), connection.next()) + .await + .expect("client should close after the test"); + }); + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::WebSocketUrl { + websocket_url, + connect_timeout: Duration::from_secs(1), + initialize_timeout: Duration::from_secs(1), + http_headers: HeaderMap::new(), + }, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ); + + assert!(!client.startup_finished()); + let _startup_task = client.start_connecting(); + let (ready, first, second) = + tokio::join!(client.wait_until_ready(), client.get(), client.get()); + ready.expect("background startup should finish"); + let first = first.expect("first waiter should receive the client"); + let second = second.expect("second waiter should receive the same client"); + + assert!(client.startup_finished()); + assert_eq!(first.session_id().as_deref(), Some("startup-session")); + assert!(Arc::ptr_eq(&first.inner, &second.inner)); + + drop(first); + drop(second); + drop(client); + server.await.expect("server task should finish"); + } + + #[tokio::test] + async fn terminal_stdio_startup_failure_is_remembered() { + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::StdioCommand { + command: StdioExecServerCommand { + program: "codex-missing-exec-server-for-test".to_string(), + args: Vec::new(), + env: HashMap::new(), + cwd: None, + }, + initialize_timeout: Duration::from_secs(1), + }, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ); + + assert!(client.start_connecting().is_none()); + assert!(!client.startup_finished()); + let first = match client.get().await { + Ok(_) => panic!("missing executable should fail"), + Err(error) => error, + }; + assert!(client.startup_finished()); + let second = match client.get().await { + Ok(_) => panic!("burned environment should stay failed"), + Err(error) => error, + }; + assert!(matches!( + client.status().await, + EnvironmentObservedStatus::Disconnected { .. } + )); + + let ( + super::ExecServerError::ConnectionAttempt(first), + super::ExecServerError::ConnectionAttempt(second), + ) = (first, second) + else { + panic!("expected saved connection failures"); + }; + assert!(Arc::ptr_eq(&first, &second)); + } + + #[tokio::test] + async fn retryable_startup_failure_does_not_burn_environment() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!( + "ws://{}", + listener.local_addr().expect("listener should have address") + ); + let (replacement_initialized_tx, replacement_initialized_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let (mut failed_startup, _) = listener.accept().await.expect("startup should arrive"); + failed_startup + .write_all(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 0\r\n\r\n") + .await + .expect("failed handshake response should write"); + + let mut replacement = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut replacement, + "replacement-session", + /*expected_resume_session_id*/ None, + ) + .await; + replacement_initialized_tx + .send(()) + .expect("replacement initialization should be observed"); + timeout(Duration::from_secs(1), replacement.next()) + .await + .expect("client should close after the test"); + }); + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::WebSocketUrl { + websocket_url, + connect_timeout: Duration::from_secs(1), + initialize_timeout: Duration::from_secs(1), + http_headers: HeaderMap::new(), + }, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ); + + let failed_startup = match client.get().await { + Ok(_) => panic!("initial connection should fail"), + Err(error) => error, + }; + assert!(matches!( + failed_startup, + super::ExecServerError::ConnectionAttempt(_) + )); + assert!(client.startup_finished()); + assert!(matches!( + client.status().await, + EnvironmentObservedStatus::Disconnected { .. } + )); + + let (ready, first, second) = + tokio::join!(client.wait_until_ready(), client.get(), client.get()); + ready.expect("later readiness check should retry startup"); + let first = first.expect("first waiter should receive the replacement client"); + let second = second.expect("second waiter should receive the same replacement client"); + assert_eq!(first.session_id().as_deref(), Some("replacement-session")); + assert!(Arc::ptr_eq(&first.inner, &second.inner)); + replacement_initialized_rx + .await + .expect("server should observe replacement initialization"); + + drop(first); + drop(second); + drop(client); + server.await.expect("server task should finish"); + } + + #[tokio::test] + async fn failed_reconnect_does_not_burn_environment() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!( + "ws://{}", + listener.local_addr().expect("listener should have address") + ); + let (replacement_initialized_tx, replacement_initialized_rx) = oneshot::channel(); + let (allow_replacement_tx, allow_replacement_rx) = watch::channel(false); + let server = tokio::spawn(async move { + let mut first = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut first, + "startup-session", + /*expected_resume_session_id*/ None, + ) + .await; + first + .close(None) + .await + .expect("startup websocket should close"); + + let successful_reconnect = loop { + let (stream, _) = listener.accept().await.expect("reconnect should arrive"); + if *allow_replacement_rx.borrow() { + break stream; + } + let mut failed_reconnect = stream; + failed_reconnect + .write_all(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 0\r\n\r\n") + .await + .expect("failed handshake response should write"); + }; + let mut successful_reconnect = accept_async(successful_reconnect) + .await + .expect("replacement websocket handshake should succeed"); + complete_websocket_initialize( + &mut successful_reconnect, + "replacement-session", + /*expected_resume_session_id*/ None, + ) + .await; + replacement_initialized_tx + .send(()) + .expect("replacement initialization should be observed"); + timeout(Duration::from_secs(1), successful_reconnect.next()) + .await + .expect("client should close after the test"); + }); + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::WebSocketUrl { + websocket_url, + connect_timeout: Duration::from_secs(1), + initialize_timeout: Duration::from_secs(1), + http_headers: HeaderMap::new(), + }, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ); + + let initial = client.get().await.expect("startup should connect"); + timeout(Duration::from_secs(1), async { + while !initial.is_disconnected() { + tokio::task::yield_now().await; + } + }) + .await + .expect("client should observe disconnect"); + let failed_reconnect = match client.get().await { + Ok(_) => panic!("first lazy reconnect should fail"), + Err(error) => error, + }; + assert!(matches!( + failed_reconnect, + super::ExecServerError::ConnectionAttempt(_) + )); + allow_replacement_tx + .send(true) + .expect("server should allow a fresh client"); + let replacement = client.get().await.expect("later reconnect should succeed"); + + assert_eq!( + replacement.session_id().as_deref(), + Some("replacement-session") + ); + replacement_initialized_rx + .await + .expect("server should observe replacement initialization"); + + drop(initial); + drop(replacement); + drop(client); + server.await.expect("server task should finish"); + } + + #[tokio::test] + async fn wake_notifications_do_not_block_other_sessions() { + let (client_stdin, server_reader) = duplex(1 << 20); + let (mut server_writer, client_stdout) = duplex(1 << 20); + let (notifications_tx, mut notifications_rx) = mpsc::channel(16); + let server = tokio::spawn(async move { + let mut lines = BufReader::new(server_reader).lines(); + let initialize = read_jsonrpc_line(&mut lines).await; + let request = match initialize { + JSONRPCMessage::Request(request) if request.method == INITIALIZE_METHOD => request, + other => panic!("expected initialize request, got {other:?}"), + }; + write_jsonrpc_line( + &mut server_writer, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(InitializeResponse { + session_id: "session-1".to_string(), + environment_info: None, + }) + .expect("initialize response should serialize"), + }), + ) + .await; + + let initialized = read_jsonrpc_line(&mut lines).await; + match initialized { + JSONRPCMessage::Notification(notification) + if notification.method == INITIALIZED_METHOD => {} + other => panic!("expected initialized notification, got {other:?}"), + } + + while let Some(message) = notifications_rx.recv().await { + write_jsonrpc_line(&mut server_writer, message).await; + } + }); + + let client = ExecServerClient::connect( + JsonRpcConnection::from_stdio( + client_stdout, + client_stdin, + "test-exec-server-client".to_string(), + ), + ExecServerClientConnectOptions::default(), + ) + .await + .expect("client should connect"); + + let noisy_process_id = ProcessId::from("noisy"); + let quiet_process_id = ProcessId::from("quiet"); + let _noisy_session = client + .register_session(&noisy_process_id) + .await + .expect("noisy session should register"); + let quiet_session = client + .register_session(&quiet_process_id) + .await + .expect("quiet session should register"); + let mut quiet_wake_rx = quiet_session.subscribe_wake(); + + for seq in 0..=4096 { + notifications_tx + .send(JSONRPCMessage::Notification(JSONRPCNotification { + method: EXEC_OUTPUT_DELTA_METHOD.to_string(), + params: Some( + serde_json::to_value(ExecOutputDeltaNotification { + process_id: noisy_process_id.clone(), + seq, + stream: ExecOutputStream::Stdout, + chunk: b"x".to_vec().into(), + }) + .expect("output notification should serialize"), + ), + })) + .await + .expect("output notification should queue"); + } + + notifications_tx + .send(JSONRPCMessage::Notification(JSONRPCNotification { + method: EXEC_EXITED_METHOD.to_string(), + params: Some( + serde_json::to_value(ExecExitedNotification { + process_id: quiet_process_id, + seq: 1, + exit_code: 17, + sandbox_denied: Some(false), + }) + .expect("exit notification should serialize"), + ), + })) + .await + .expect("exit notification should queue"); + + timeout(Duration::from_secs(1), quiet_wake_rx.changed()) + .await + .expect("quiet session should receive wake before timeout") + .expect("quiet wake channel should stay open"); + assert_eq!(*quiet_wake_rx.borrow(), 1); + + drop(notifications_tx); + drop(client); + server.await.expect("server task should finish"); + } + + mod network_policy_tests; +} diff --git a/codex-rs/exec-server/src/client/accepted.rs b/codex-rs/exec-server/src/client/accepted.rs new file mode 100644 index 0000000000000000000000000000000000000000..c7853221a06abb3cad86e95fcae1c8ced1bf8536 --- /dev/null +++ b/codex-rs/exec-server/src/client/accepted.rs @@ -0,0 +1,266 @@ +use std::sync::Arc; + +use axum::extract::ws::WebSocket; +use futures::lock::Mutex; +use tokio::sync::OnceCell; +use tokio::sync::OwnedSemaphorePermit; +use tokio::sync::Semaphore; +use tokio::sync::mpsc; +use tokio::sync::watch; + +use super::ConnectionStatus; +use super::ExecServerClient; +use super::Inner; +use super::LazyRemoteExecServerClient; +use crate::EnvironmentConnectionState; +use crate::ExecServerClientConnectOptions; +use crate::ExecServerError; +use crate::client_transport::ExecServerReconnectStrategy; +use crate::client_transport::ReconnectAttempt; +use crate::connection::JsonRpcConnection; +use codex_http_client::HttpClientFactory; + +struct AcceptedReplacement { + connection: JsonRpcConnection, + permit: OwnedSemaphorePermit, +} + +struct AcceptedConnectionSourceInner { + replacements_tx: mpsc::UnboundedSender, + replacements_rx: Mutex>, + replacement_slots: Arc, +} + +/// Receives authenticated connections supplied by an embedding host. +/// +/// The source owns serialization and cancellation cleanup for replacement +/// handoffs. The reconnect loop only asks it for the next connection. +#[derive(Clone)] +pub(crate) struct AcceptedConnectionSource { + inner: Arc, + options: ExecServerClientConnectOptions, +} + +struct AcceptedReplacementSubmission { + source: AcceptedConnectionSource, + permit: OwnedSemaphorePermit, +} + +impl AcceptedConnectionSource { + fn new(options: ExecServerClientConnectOptions) -> Self { + let (replacements_tx, replacements_rx) = mpsc::unbounded_channel(); + Self { + inner: Arc::new(AcceptedConnectionSourceInner { + replacements_tx, + replacements_rx: Mutex::new(replacements_rx), + replacement_slots: Arc::new(Semaphore::new(1)), + }), + options, + } + } + + fn begin_replacement(&self) -> Result { + let permit = Arc::clone(&self.inner.replacement_slots) + .try_acquire_owned() + .map_err(|_| { + ExecServerError::Protocol( + "an accepted exec-server replacement is already in progress".to_string(), + ) + })?; + Ok(AcceptedReplacementSubmission { + source: self.clone(), + permit, + }) + } + + pub(crate) async fn next_connection( + &self, + session_id: &str, + ) -> Result { + let replacement = self + .inner + .replacements_rx + .lock() + .await + .recv() + .await + .ok_or_else(|| { + ExecServerError::Disconnected( + "accepted exec-server replacement channel closed".to_string(), + ) + })?; + let mut options = self.options.clone(); + options.resume_session_id = Some(session_id.to_string()); + Ok(ReconnectAttempt::with_attempt_permit( + replacement.connection, + options, + replacement.permit, + )) + } +} + +impl AcceptedReplacementSubmission { + fn submit(self, connection: JsonRpcConnection) -> Result<(), ExecServerError> { + self.source + .inner + .replacements_tx + .send(AcceptedReplacement { + connection, + permit: self.permit, + }) + .map_err(|_| { + ExecServerError::Disconnected( + "accepted exec-server connection is no longer awaiting replacements" + .to_string(), + ) + }) + } +} + +impl ExecServerClient { + /// Initializes an exec-server client over a WebSocket accepted by an Axum handler. + /// + /// The caller owns accepting and authenticating replacement WebSockets. + pub(crate) async fn connect_accepted_websocket( + websocket: WebSocket, + options: ExecServerClientConnectOptions, + ) -> Result { + if options.resume_session_id.is_some() { + return Err(ExecServerError::Protocol( + "accepted exec-server initial connection cannot resume a session".to_string(), + )); + } + let connection_source = AcceptedConnectionSource::new(options.clone()); + Self::connect_with_recovery( + JsonRpcConnection::from_axum_websocket( + websocket, + "accepted exec-server websocket".to_string(), + ), + options, + Some(ExecServerReconnectStrategy::Accepted(connection_source)), + ) + .await + } + + /// Supplies an authenticated replacement WebSocket for this accepted client. + /// + /// Retires the old transport before resuming the saved session. Returns + /// after handoff; recovery continues asynchronously. + pub(crate) async fn replace_accepted_websocket( + &self, + websocket: WebSocket, + ) -> Result<(), ExecServerError> { + self.inner + .accept_replacement_connection(JsonRpcConnection::from_axum_websocket( + websocket, + "accepted exec-server replacement websocket".to_string(), + )) + .await + } +} + +impl Inner { + /// Hands a replacement connection from the host to the existing accepted client. + /// + /// This method coordinates the handoff in this order: + /// + /// 1. Verify that this client uses the accepted connection source. + /// 2. Reserve the source so concurrent handoffs are rejected. + /// 3. If the old RPC transport is still connected, move the client into + /// recovery and close that transport before attaching the same session to + /// the replacement. + /// 4. Queue the raw connection for the recovery task. That task creates the + /// new RPC client, runs the initialize/resume handshake with the saved + /// session ID, and recovers the existing processes. + async fn accept_replacement_connection( + self: &Arc, + connection: JsonRpcConnection, + ) -> Result<(), ExecServerError> { + if self.session_id.get().is_none() { + return Err(ExecServerError::Protocol( + "accepted exec-server connection is missing its session ID".to_string(), + )); + } + + let Some(ExecServerReconnectStrategy::Accepted(connection_source)) = + &self.reconnect_strategy + else { + return Err(ExecServerError::Protocol( + "only an accepted exec-server connection can be replaced directly".to_string(), + )); + }; + let (current_rpc_client, replacement_submission) = { + let connection = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let current_rpc_client = match &connection.status { + ConnectionStatus::Failed(message) => { + return Err(ExecServerError::Disconnected(message.clone())); + } + ConnectionStatus::Connected(rpc_client) => Some(Arc::clone(rpc_client)), + ConnectionStatus::Recovering => None, + }; + let replacement_submission = connection_source.begin_replacement()?; + (current_rpc_client, replacement_submission) + }; + if let Some(current_rpc_client) = current_rpc_client { + self.request_recovery( + Arc::clone(¤t_rpc_client), + "exec-server connection replaced".to_string(), + ); + current_rpc_client.close_transport().await; + } + // Synchronize the enqueue with terminal recovery. Recovery can time out + // while the handoff is waiting for the old transport to close; in that + // case the host must not receive success for a socket that no task will + // consume. Holding the connection lock through the synchronous send + // makes either the failure or the enqueue win the race unambiguously. + let connection_state = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let ConnectionStatus::Failed(message) = &connection_state.status { + return Err(ExecServerError::Disconnected(message.clone())); + } + replacement_submission.submit(connection) + } +} + +#[cfg(test)] +#[path = "accepted_tests.rs"] +mod tests; + +impl LazyRemoteExecServerClient { + pub(crate) fn from_connected( + client: ExecServerClient, + http_client_factory: HttpClientFactory, + ) -> Self { + let environment_connection_state_tx = + watch::channel(EnvironmentConnectionState::Connected).0; + client.attach_environment_connection_state(environment_connection_state_tx.clone()); + Self { + transport_params: None, + http_client_factory, + recovery_policy: super::RecoveryPolicy::Wait, + startup: std::sync::Arc::new(super::ConnectionAttempt { + result: OnceCell::new_with(Some(Ok(client.clone()))), + ..Default::default() + }), + current_client: std::sync::Arc::new(std::sync::Mutex::new(Some(client))), + reconnect: std::sync::Arc::new(std::sync::Mutex::new(None)), + refresh_lock: std::sync::Arc::new(tokio::sync::Mutex::new(())), + environment_connection_state_tx, + } + } + + pub(crate) async fn replace_accepted_websocket( + &self, + websocket: WebSocket, + ) -> Result<(), ExecServerError> { + self.get() + .await? + .replace_accepted_websocket(websocket) + .await + } +} diff --git a/codex-rs/exec-server/src/client/accepted_tests.rs b/codex-rs/exec-server/src/client/accepted_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..f284f1cf39a7ea8d14f10fd59c3e40cc27daed4e --- /dev/null +++ b/codex-rs/exec-server/src/client/accepted_tests.rs @@ -0,0 +1,34 @@ +use tokio::sync::oneshot; + +use super::AcceptedConnectionSource; +use crate::ExecServerClientConnectOptions; + +#[tokio::test] +async fn replacement_claim_is_released_when_handoff_is_cancelled() { + let source = AcceptedConnectionSource::new(ExecServerClientConnectOptions::default()); + let (claimed_tx, claimed_rx) = oneshot::channel(); + let (_release_tx, release_rx) = oneshot::channel::<()>(); + + let handoff = tokio::spawn({ + let source = source.clone(); + async move { + let _submission = source + .begin_replacement() + .expect("the first replacement should claim the handoff"); + claimed_tx + .send(()) + .expect("the test should wait for the claim"); + let _ = release_rx.await; + } + }); + + claimed_rx.await.expect("the handoff should be claimed"); + assert!(source.begin_replacement().is_err()); + + handoff.abort(); + let _ = handoff.await; + + source + .begin_replacement() + .expect("a new replacement should be accepted after cancellation"); +} diff --git a/codex-rs/exec-server/src/client/http_client.rs b/codex-rs/exec-server/src/client/http_client.rs new file mode 100644 index 0000000000000000000000000000000000000000..cbafd448a0a50f984978f1b37f9694f54c682ad9 --- /dev/null +++ b/codex-rs/exec-server/src/client/http_client.rs @@ -0,0 +1,26 @@ +//! HTTP client capability implementations shared by local and remote environments. +//! +//! This module is the facade for the environment-owned [`crate::HttpClient`] +//! capability: +//! - [`RouteAwareHttpClient`] executes requests through the shared transport +//! - [`ExecServerClient`] forwards requests over the JSON-RPC transport +//! - [`HttpResponseBodyStream`] presents buffered local bodies and streamed +//! remote `http/request/bodyDelta` notifications through one byte-stream API +//! +//! Runtime split: +//! - orchestrator process: holds an `Arc` and chooses local or +//! remote execution +//! - remote runtime: serves the `http/request` RPC and runs the concrete local +//! HTTP request there when the orchestrator uses [`ExecServerClient`] + +#[path = "http_response_body_stream.rs"] +pub(crate) mod response_body_stream; +#[path = "route_aware_http_client.rs"] +mod route_aware_http_client; +#[path = "rpc_http_client.rs"] +mod rpc_http_client; + +pub use response_body_stream::HttpResponseBodyStream; +pub(crate) use route_aware_http_client::PendingRouteAwareHttpBodyStream; +pub use route_aware_http_client::RouteAwareHttpClient; +pub(crate) use route_aware_http_client::RouteAwareHttpRequestRunner; diff --git a/codex-rs/exec-server/src/client/http_response_body_stream.rs b/codex-rs/exec-server/src/client/http_response_body_stream.rs new file mode 100644 index 0000000000000000000000000000000000000000..be52e5cb6a29434c0a7e7c5ef50185c5ceeb3911 --- /dev/null +++ b/codex-rs/exec-server/src/client/http_response_body_stream.rs @@ -0,0 +1,446 @@ +//! Shared HTTP response-body stream plumbing for local and remote execution. +//! +//! This module owns the byte-stream type exposed by the `HttpClient` +//! capability plus the remote-side routing table used to turn +//! `http/request/bodyDelta` notifications back into per-request streams. + +use std::collections::HashMap; +use std::pin::Pin; +use std::sync::Arc; +use std::sync::atomic::Ordering; + +use bytes::Bytes; +use codex_http_client::HttpError; +use codex_http_client::HttpResponse; +use futures::StreamExt; +use serde_json::Value; +use serde_json::from_value; +use tokio::runtime::Handle; +use tokio::sync::OwnedSemaphorePermit; +use tokio::sync::mpsc; +use tokio::sync::mpsc::error::TrySendError; +use tracing::debug; + +use crate::client::ExecServerError; +use crate::client::Inner; +use crate::protocol::HTTP_REQUEST_BODY_DELTA_METHOD; +use crate::protocol::HttpRequestBodyDeltaNotification; +use crate::protocol::MAX_HTTP_BODY_DELTA_BYTES; +use crate::rpc::RpcNotificationSender; + +pub(crate) const MAX_QUEUED_HTTP_BODY_BYTES: usize = 16 * 1024 * 1024; +const MAX_ENCODED_HTTP_BODY_DELTA_BYTES: usize = MAX_HTTP_BODY_DELTA_BYTES.div_ceil(3) * 4; + +pub(crate) struct QueuedHttpBodyDelta { + notification: HttpRequestBodyDeltaNotification, + _byte_permit: Option, +} + +impl QueuedHttpBodyDelta { + pub(crate) fn new( + notification: HttpRequestBodyDeltaNotification, + byte_permit: Option, + ) -> Self { + Self { + notification, + _byte_permit: byte_permit, + } + } +} + +pub(super) struct HttpBodyStreamRegistration { + inner: Arc, + request_id: String, + active: bool, +} + +enum HttpResponseBodyStreamInner { + Local { + body: Pin> + Send>>, + }, + Remote { + inner: Arc, + request_id: String, + next_seq: u64, + rx: mpsc::Receiver, + pending_eof: bool, + closed: bool, + }, +} + +/// Request-scoped stream of body chunks for an HTTP response. +/// +/// The initial `http/request` call returns status and headers. This stream then +/// receives the ordered `http/request/bodyDelta` notifications for that request +/// id until EOF or a terminal error. +pub struct HttpResponseBodyStream { + inner: HttpResponseBodyStreamInner, +} + +impl HttpResponseBodyStream { + /// Creates an in-memory response stream from pre-buffered chunks. + /// + /// This is useful for [`crate::HttpClient`] implementations that already + /// own the response bytes, including lightweight test clients. + #[doc(hidden)] + pub fn from_chunks(chunks: Vec>) -> Self { + let body = futures::stream::iter( + chunks + .into_iter() + .map(|chunk| Ok::(chunk.into())), + ); + Self { + inner: HttpResponseBodyStreamInner::Local { + body: Box::pin(body), + }, + } + } + + pub(super) fn local(response: HttpResponse) -> Self { + Self { + inner: HttpResponseBodyStreamInner::Local { + body: Box::pin(response.bytes_stream()), + }, + } + } + + pub(super) fn remote( + inner: Arc, + request_id: String, + rx: mpsc::Receiver, + ) -> Self { + Self { + inner: HttpResponseBodyStreamInner::Remote { + inner, + request_id, + next_seq: 1, + rx, + pending_eof: false, + closed: false, + }, + } + } + + /// Receives the next response-body chunk. + /// + /// Returns `Ok(None)` at EOF and converts sequence gaps or stream-side + /// stream errors into protocol errors. + pub async fn recv(&mut self) -> Result>, ExecServerError> { + match &mut self.inner { + HttpResponseBodyStreamInner::Local { body } => match body.next().await { + Some(chunk) => match chunk { + Ok(bytes) => Ok(Some(bytes.to_vec())), + Err(error) => Err(ExecServerError::HttpRequest(error.to_string())), + }, + None => Ok(None), + }, + HttpResponseBodyStreamInner::Remote { + inner, + request_id, + next_seq, + rx, + pending_eof, + closed, + } => { + if *pending_eof { + *pending_eof = false; + finish_remote_stream(inner, request_id, closed).await; + return Ok(None); + } + + let Some(QueuedHttpBodyDelta { + notification: delta, + .. + }) = rx.recv().await + else { + finish_remote_stream(inner, request_id, closed).await; + if let Some(error) = inner.take_http_body_stream_failure(request_id).await { + return Err(ExecServerError::Protocol(format!( + "http response stream `{request_id}` failed: {error}", + ))); + } + return Ok(None); + }; + if delta.seq != *next_seq { + finish_remote_stream(inner, request_id, closed).await; + return Err(ExecServerError::Protocol(format!( + "http response stream `{request_id}` received seq {}, expected {}", + delta.seq, *next_seq + ))); + } + *next_seq += 1; + let chunk = delta.delta.into_inner(); + + if let Some(error) = delta.error { + finish_remote_stream(inner, request_id, closed).await; + return Err(ExecServerError::Protocol(format!( + "http response stream `{request_id}` failed: {error}", + ))); + } + if delta.done { + finish_remote_stream(inner, request_id, closed).await; + if chunk.is_empty() { + return Ok(None); + } + *pending_eof = true; + } + Ok(Some(chunk)) + } + } + } +} + +impl Drop for HttpResponseBodyStream { + /// Schedules stream-route removal if the consumer drops before EOF. + fn drop(&mut self) { + if let HttpResponseBodyStreamInner::Remote { + inner, + request_id, + closed, + .. + } = &mut self.inner + { + if *closed { + return; + } + *closed = true; + spawn_remove_http_body_stream(Arc::clone(inner), request_id.clone()); + } + } +} + +impl HttpBodyStreamRegistration { + pub(super) fn new(inner: Arc, request_id: String) -> Self { + Self { + inner, + request_id, + active: true, + } + } + + pub(super) fn disarm(&mut self) { + self.active = false; + } +} + +impl Drop for HttpBodyStreamRegistration { + /// Removes the route if the stream request future is cancelled before headers return. + fn drop(&mut self) { + if self.active { + spawn_remove_http_body_stream(Arc::clone(&self.inner), self.request_id.clone()); + } + } +} + +async fn finish_remote_stream(inner: &Arc, request_id: &str, closed: &mut bool) { + if *closed { + return; + } + *closed = true; + inner.remove_http_body_stream(request_id).await; +} + +/// Schedules HTTP body route removal from synchronous drop paths. +fn spawn_remove_http_body_stream(inner: Arc, request_id: String) { + if let Ok(handle) = Handle::try_current() { + handle.spawn(async move { + inner.remove_http_body_stream(&request_id).await; + }); + } +} + +pub(super) async fn send_body_delta( + notifications: &RpcNotificationSender, + delta: HttpRequestBodyDeltaNotification, +) -> bool { + notifications + .notify(HTTP_REQUEST_BODY_DELTA_METHOD, &delta) + .await + .is_ok() +} + +impl Inner { + /// Routes one streamed HTTP body notification into its request-local receiver. + pub(crate) async fn handle_http_body_delta_notification( + &self, + params: Option, + ) -> Result<(), ExecServerError> { + let params = params.unwrap_or(Value::Null); + if params + .get("deltaBase64") + .and_then(Value::as_str) + .is_some_and(|delta| delta.len() > MAX_ENCODED_HTTP_BODY_DELTA_BYTES) + { + return Err(ExecServerError::Protocol(format!( + "http response body delta exceeds {MAX_HTTP_BODY_DELTA_BYTES} bytes" + ))); + } + let params: HttpRequestBodyDeltaNotification = from_value(params)?; + if params.delta.0.len() > MAX_HTTP_BODY_DELTA_BYTES { + return Err(ExecServerError::Protocol(format!( + "http response body delta exceeds {MAX_HTTP_BODY_DELTA_BYTES} bytes" + ))); + } + // Unknown request ids are ignored intentionally: a stream may have already + // reached EOF and released its route. + if let Some(tx) = self + .http_body_streams + .load() + .get(¶ms.request_id) + .cloned() + { + let request_id = params.request_id.clone(); + let terminal_delta = params.done || params.error.is_some(); + let queued_bytes = params + .delta + .0 + .len() + .saturating_add(params.error.as_deref().map_or(0, str::len)); + let byte_permit = if queued_bytes == 0 { + None + } else { + u32::try_from(queued_bytes).ok().and_then(|queued_bytes| { + Arc::clone(&self.http_body_stream_byte_budget) + .try_acquire_many_owned(queued_bytes) + .ok() + }) + }; + if queued_bytes > 0 && byte_permit.is_none() { + self.record_http_body_stream_failure( + &request_id, + format!("queued body deltas exceed {MAX_QUEUED_HTTP_BODY_BYTES} bytes"), + ) + .await; + self.remove_http_body_stream(&request_id).await; + debug!( + "closing http response stream `{request_id}` after exhausting the queued byte budget" + ); + return Ok(()); + } + match tx.try_send(QueuedHttpBodyDelta::new(params, byte_permit)) { + Ok(()) => { + if terminal_delta { + self.remove_http_body_stream(&request_id).await; + } + } + Err(TrySendError::Closed(_)) => { + self.remove_http_body_stream(&request_id).await; + debug!("http response stream receiver dropped before body delta delivery"); + } + Err(TrySendError::Full(_)) => { + self.record_http_body_stream_failure( + &request_id, + "body delta channel filled before delivery".to_string(), + ) + .await; + self.remove_http_body_stream(&request_id).await; + debug!( + "closing http response stream `{request_id}` after body delta backpressure" + ); + } + } + } + Ok(()) + } + + /// Fails active streamed HTTP bodies so callers do not wait forever after a + /// transport disconnect or notification handling failure. + pub(crate) async fn fail_all_http_body_streams(&self, message: String) { + let _streams_write_guard = self.http_body_streams_write_lock.lock().await; + let streams = self.http_body_streams.load(); + let streams = streams.as_ref().clone(); + self.http_body_streams.store(Arc::new(HashMap::new())); + for (request_id, tx) in streams { + // Failure notifications must wake every stream even when no + // byte-budget permits remain. + if tx + .try_send(QueuedHttpBodyDelta::new( + HttpRequestBodyDeltaNotification { + request_id: request_id.clone(), + seq: 1, + delta: Vec::new().into(), + done: true, + error: Some(message.clone()), + }, + /*byte_permit*/ None, + )) + .is_err() + { + let mut next_failures = self.http_body_stream_failures.load().as_ref().clone(); + next_failures.insert(request_id, message.clone()); + self.http_body_stream_failures + .store(Arc::new(next_failures)); + } + } + } + + /// Allocates a connection-local streamed HTTP response id. + pub(super) fn next_http_body_stream_request_id(&self) -> String { + let id = self + .http_body_stream_next_id + .fetch_add(1, Ordering::Relaxed); + format!("http-{id}") + } + + /// Registers a request id before issuing a streaming HTTP call. + pub(super) async fn insert_http_body_stream( + &self, + request_id: String, + tx: mpsc::Sender, + ) -> Result<(), ExecServerError> { + let _streams_write_guard = self.http_body_streams_write_lock.lock().await; + let streams = self.http_body_streams.load(); + if streams.contains_key(&request_id) { + return Err(ExecServerError::Protocol(format!( + "http response stream already registered for request {request_id}" + ))); + } + let mut next_streams = streams.as_ref().clone(); + next_streams.insert(request_id.clone(), tx); + self.http_body_streams.store(Arc::new(next_streams)); + let failures = self.http_body_stream_failures.load(); + if failures.contains_key(&request_id) { + let mut next_failures = failures.as_ref().clone(); + next_failures.remove(&request_id); + self.http_body_stream_failures + .store(Arc::new(next_failures)); + } + Ok(()) + } + + /// Removes a request id after EOF, terminal error, or request failure. + pub(super) async fn remove_http_body_stream( + &self, + request_id: &str, + ) -> Option> { + let _streams_write_guard = self.http_body_streams_write_lock.lock().await; + let streams = self.http_body_streams.load(); + let stream = streams.get(request_id).cloned(); + stream.as_ref()?; + let mut next_streams = streams.as_ref().clone(); + next_streams.remove(request_id); + self.http_body_streams.store(Arc::new(next_streams)); + stream + } + + async fn record_http_body_stream_failure(&self, request_id: &str, message: String) { + let _streams_write_guard = self.http_body_streams_write_lock.lock().await; + let failures = self.http_body_stream_failures.load(); + let mut next_failures = failures.as_ref().clone(); + next_failures.insert(request_id.to_string(), message); + self.http_body_stream_failures + .store(Arc::new(next_failures)); + } + + async fn take_http_body_stream_failure(&self, request_id: &str) -> Option { + let _streams_write_guard = self.http_body_streams_write_lock.lock().await; + let failures = self.http_body_stream_failures.load(); + let error = failures.get(request_id).cloned(); + error.as_ref()?; + let mut next_failures = failures.as_ref().clone(); + next_failures.remove(request_id); + self.http_body_stream_failures + .store(Arc::new(next_failures)); + error + } +} diff --git a/codex-rs/exec-server/src/client/network_policy_audit.rs b/codex-rs/exec-server/src/client/network_policy_audit.rs new file mode 100644 index 0000000000000000000000000000000000000000..b16249a9b9b3622408cca1373ae2961c11c27a24 --- /dev/null +++ b/codex-rs/exec-server/src/client/network_policy_audit.rs @@ -0,0 +1,81 @@ +use super::NetworkPolicyAuditContext; +use crate::protocol::ExecServerNetworkProtocol; +use crate::protocol::MAX_NETWORK_POLICY_HOST_BYTES; +use crate::protocol::MAX_NETWORK_POLICY_PROCESS_ID_BYTES; +use crate::protocol::MAX_NETWORK_POLICY_REASON_BYTES; +use crate::protocol::NetworkPolicyDecisionNotification; + +const MAX_NETWORK_POLICY_METHOD_BYTES: usize = 32; +const MAX_NETWORK_POLICY_CLIENT_BYTES: usize = 256; +const MAX_NETWORK_POLICY_TIMESTAMP_BYTES: usize = 64; + +pub(super) fn emit_network_policy_decision( + context: &NetworkPolicyAuditContext, + decision: &NetworkPolicyDecisionNotification, +) -> bool { + if decision.process_id.is_empty() + || decision.process_id.len() > MAX_NETWORK_POLICY_PROCESS_ID_BYTES + || decision.host.is_empty() + || decision.host.len() > MAX_NETWORK_POLICY_HOST_BYTES + || decision.host.chars().any(char::is_control) + || decision.host.chars().any(char::is_whitespace) + || decision.reason.len() > MAX_NETWORK_POLICY_REASON_BYTES + || decision.reason.chars().any(char::is_control) + || !matches!(decision.scope.as_str(), "domain" | "non_domain") + || !matches!(decision.decision.as_str(), "allow" | "deny" | "ask") + || !matches!( + decision.source.as_str(), + "baseline_policy" | "mode_guard" | "proxy_state" | "decider" + ) + || decision.timestamp.is_empty() + || decision.timestamp.len() > MAX_NETWORK_POLICY_TIMESTAMP_BYTES + || decision.timestamp.chars().any(char::is_control) + || decision.method.as_ref().is_some_and(|method| { + method.len() > MAX_NETWORK_POLICY_METHOD_BYTES + || method.chars().any(char::is_control) + || method.chars().any(char::is_whitespace) + }) + || decision.client.as_ref().is_some_and(|client| { + client.len() > MAX_NETWORK_POLICY_CLIENT_BYTES + || client.chars().any(char::is_control) + || client.chars().any(char::is_whitespace) + }) + { + return false; + } + + let protocol = match decision.protocol { + ExecServerNetworkProtocol::Http => "http", + ExecServerNetworkProtocol::HttpsConnect => "https_connect", + ExecServerNetworkProtocol::Socks5Tcp => "socks5_tcp", + ExecServerNetworkProtocol::Socks5Udp => "socks5_udp", + }; + let metadata = &context.metadata; + tracing::event!( + target: "codex_otel.log_only", + tracing::Level::INFO, + event.name = "codex.network_proxy.policy_decision", + event.timestamp = decision.timestamp, + conversation.id = metadata.conversation_id.as_deref(), + app.version = metadata.app_version.as_deref(), + auth_mode = metadata.auth_mode.as_deref(), + originator = metadata.originator.as_deref(), + user.account_id = metadata.user_account_id.as_deref(), + user.email = metadata.user_email.as_deref(), + terminal.type = metadata.terminal_type.as_deref(), + model = metadata.model.as_deref(), + slug = metadata.slug.as_deref(), + network.policy.scope = decision.scope, + network.policy.decision = decision.decision, + network.policy.source = decision.source, + network.policy.reason = decision.reason, + network.transport.protocol = protocol, + server.address = decision.host, + server.port = decision.port, + http.request.method = decision.method.as_deref().unwrap_or("none"), + client.address = decision.client.as_deref().unwrap_or("unknown"), + execution.id = context.execution_id.as_deref(), + network.policy.override = decision.policy_override, + ); + true +} diff --git a/codex-rs/exec-server/src/client/route_aware_http_client.rs b/codex-rs/exec-server/src/client/route_aware_http_client.rs new file mode 100644 index 0000000000000000000000000000000000000000..81dd4607ddc5a5a5774aba79fc911a0ab2959afc --- /dev/null +++ b/codex-rs/exec-server/src/client/route_aware_http_client.rs @@ -0,0 +1,378 @@ +//! Route-aware local HTTP capability implementation. +//! +//! This code runs wherever the real network request should originate: +//! - in a local environment, that means the orchestrator process +//! - in a remote environment, that means the remote runtime after the +//! orchestrator has forwarded `http/request` over JSON-RPC + +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_http_client::ClientRouteClass; +use codex_http_client::HttpClientFactory; +use codex_http_client::RouteAwareClientPool; +use codex_http_client::RouteAwareRequestError; +use codex_protocol::shell_environment::CODEX_EXEC_SERVER_NOISE_AUTH_TOKEN_ENV_VAR; +use codex_protocol::shell_environment::OPENAI_FEDERATION_RULE_ID_ENV_VAR; +use codex_protocol::shell_environment::OPENAI_IDENTITY_TOKEN_FILE_ENV_VAR; +use codex_protocol::shell_environment::OPENAI_WORKLOAD_IDENTITY_CONTEXT_ENV_VAR; +use futures::FutureExt; +use futures::StreamExt; +use futures::future::BoxFuture; +use http::HeaderMap; +use http::HeaderName; +use http::HeaderValue; +use http::Method; +use tracing::Instrument; +use url::Url; + +use super::HttpResponseBodyStream; +use super::response_body_stream::send_body_delta; +use crate::HttpClient; +use crate::client::ExecServerError; +use crate::protocol::HttpHeader; +use crate::protocol::HttpRedirectPolicy; +use crate::protocol::HttpRequestBodyDeltaNotification; +use crate::protocol::HttpRequestParams; +use crate::protocol::HttpRequestResponse; +use crate::protocol::MAX_HTTP_BODY_DELTA_BYTES; +use crate::rpc::RpcNotificationSender; +use crate::rpc::internal_error; +use crate::rpc::invalid_params; + +const HTTP_HEADER_ENV_DENYLIST: &[&str] = &[ + CODEX_EXEC_SERVER_NOISE_AUTH_TOKEN_ENV_VAR, + OPENAI_FEDERATION_RULE_ID_ENV_VAR, + OPENAI_IDENTITY_TOKEN_FILE_ENV_VAR, + OPENAI_WORKLOAD_IDENTITY_CONTEXT_ENV_VAR, + "OPENAI_API_KEY", + "CODEX_API_KEY", + "CODEX_ACCESS_TOKEN", + "CODEX_CONNECTORS_TOKEN", + "AWS_ACCESS_KEY_ID", + "AWS_SECRET_ACCESS_KEY", + "AWS_SESSION_TOKEN", + "AZURE_CLIENT_SECRET", + "AZURE_FEDERATED_TOKEN_FILE", + "GOOGLE_APPLICATION_CREDENTIALS", +]; + +/// HTTP capability implementation backed by the shared route-aware transport. +#[derive(Clone)] +pub struct RouteAwareHttpClient { + follow_redirects: RouteAwareClientPool, + stop_redirects: RouteAwareClientPool, +} + +/// Streaming response state held between the initial HTTP response and +/// downstream body-delta forwarding. +pub(crate) struct PendingRouteAwareHttpBodyStream { + pub(crate) request_id: String, + pub(crate) response: codex_http_client::HttpResponse, +} + +/// Validates `http/request` parameters and runs the actual HTTP call used +/// by the exec-server route and the local [`HttpClient`] backend. +pub(crate) struct RouteAwareHttpRequestRunner { + client: RouteAwareClientPool, +} + +impl RouteAwareHttpClient { + pub fn new(http_client_factory: HttpClientFactory) -> Self { + Self { + follow_redirects: RouteAwareClientPool::with_chatgpt_cloudflare_cookies_without_request_logging( + http_client_factory.clone(), + // Delegated HTTP targets arbitrary endpoints; route class only labels diagnostics. + ClientRouteClass::Other, + ), + stop_redirects: + RouteAwareClientPool::with_chatgpt_cloudflare_cookies_without_redirects_or_request_logging( + http_client_factory, + // Proxy routing comes from the factory, not this diagnostic-only route class. + ClientRouteClass::Other, + ), + } + } + + /// Enables narrowly scoped TLS-backend fallback for both redirect policies. + pub fn with_tls_backend_fallback(mut self) -> Self { + self.follow_redirects = self.follow_redirects.with_tls_backend_fallback(); + self.stop_redirects = self.stop_redirects.with_tls_backend_fallback(); + self + } + + pub(crate) fn runner( + &self, + redirect_policy: HttpRedirectPolicy, + ) -> RouteAwareHttpRequestRunner { + let client = match redirect_policy { + HttpRedirectPolicy::Follow => self.follow_redirects.clone(), + HttpRedirectPolicy::Stop => self.stop_redirects.clone(), + }; + RouteAwareHttpRequestRunner { client } + } +} + +impl HttpClient for RouteAwareHttpClient { + fn http_request( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + async move { + let runner = self.runner(params.redirect_policy); + let (response, _) = runner + .run(HttpRequestParams { + stream_response: false, + ..params + }) + .await + .map_err(|error| ExecServerError::HttpRequest(error.message))?; + Ok(response) + } + .boxed() + } + + fn http_request_stream( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + async move { + let runner = self.runner(params.redirect_policy); + let (response, pending_stream) = runner + .run(HttpRequestParams { + stream_response: true, + ..params + }) + .await + .map_err(|error| ExecServerError::HttpRequest(error.message))?; + let pending_stream = pending_stream.ok_or_else(|| { + ExecServerError::Protocol( + "http request stream did not return a response body stream".to_string(), + ) + })?; + Ok(( + response, + HttpResponseBodyStream::local(pending_stream.response), + )) + } + .boxed() + } +} + +impl RouteAwareHttpRequestRunner { + pub(crate) async fn run( + &self, + params: HttpRequestParams, + ) -> Result<(HttpRequestResponse, Option), JSONRPCErrorError> + { + let method = Method::from_bytes(params.method.as_bytes()) + .map_err(|error| invalid_params(format!("http/request method is invalid: {error}")))?; + let url = Url::parse(¶ms.url) + .map_err(|error| invalid_params(format!("http/request url is invalid: {error}")))?; + match url.scheme() { + "http" | "https" => {} + scheme => { + return Err(invalid_params(format!( + "http/request only supports http and https URLs, got {scheme}" + ))); + } + } + + let request_span = tracing::info_span!( + "codex.exec_server.http_request", + otel.kind = "client", + http.request.method = method.as_str(), + server.address = url.host_str().unwrap_or_default(), + server.port = u64::from(url.port_or_known_default().unwrap_or_default()), + http.response.status_code = tracing::field::Empty, + error.type = tracing::field::Empty, + ); + let mut headers = Self::build_headers(params.headers)?; + codex_otel::inject_span_w3c_trace_headers(&request_span, &mut headers); + let mut request = self.client.request(method.clone(), url).headers(headers); + if let Some(body) = params.body { + request = request.body(body.into_inner()); + } + if let Some(timeout_ms) = params.timeout_ms { + request = request.timeout(Duration::from_millis(timeout_ms)); + } + + let response = match request.send().instrument(request_span.clone()).await { + Ok(response) => response, + Err(error) => { + request_span.record("error.type", "request"); + let error_message = error.to_string(); + log_send_error(&method, error); + return Err(internal_error(format!( + "http/request failed: {error_message}" + ))); + } + }; + let status = response.status().as_u16(); + request_span.record("http.response.status_code", u64::from(status)); + let headers = Self::response_headers(response.headers()); + + if params.stream_response { + return Ok(( + HttpRequestResponse { + status, + headers, + body: Vec::new().into(), + }, + Some(PendingRouteAwareHttpBodyStream { + request_id: params.request_id, + response, + }), + )); + } + + let body = response.bytes().await.map_err(|error| { + internal_error(format!( + "failed to read http/request response body: {error}" + )) + })?; + + Ok(( + HttpRequestResponse { + status, + headers, + body: body.to_vec().into(), + }, + None, + )) + } + + pub(crate) async fn stream_body( + pending_stream: PendingRouteAwareHttpBodyStream, + notifications: RpcNotificationSender, + ) { + let PendingRouteAwareHttpBodyStream { + request_id, + response, + } = pending_stream; + let mut seq = 1; + let mut body = response.bytes_stream(); + while let Some(chunk) = body.next().await { + match chunk { + Ok(bytes) => { + for chunk in bytes.chunks(MAX_HTTP_BODY_DELTA_BYTES) { + if !send_body_delta( + ¬ifications, + HttpRequestBodyDeltaNotification { + request_id: request_id.clone(), + seq, + delta: chunk.to_vec().into(), + done: false, + error: None, + }, + ) + .await + { + return; + } + seq += 1; + } + } + Err(error) => { + let _ = send_body_delta( + ¬ifications, + HttpRequestBodyDeltaNotification { + request_id, + seq, + delta: Vec::new().into(), + done: true, + error: Some(error.to_string()), + }, + ) + .await; + return; + } + } + } + + let _ = send_body_delta( + ¬ifications, + HttpRequestBodyDeltaNotification { + request_id, + seq, + delta: Vec::new().into(), + done: true, + error: None, + }, + ) + .await; + } + + fn build_headers(headers: Vec) -> Result { + let mut header_map = HeaderMap::new(); + for header in headers { + let name = HeaderName::from_bytes(header.name.as_bytes()).map_err(|error| { + invalid_params(format!("http/request header name is invalid: {error}")) + })?; + let value = match header.value_env_var { + Some(env_var) => { + if HTTP_HEADER_ENV_DENYLIST + .iter() + .any(|denied| denied.eq_ignore_ascii_case(&env_var)) + { + return Err(invalid_params(format!( + "http/request header {} cannot use executor environment variable {env_var}", + header.name + ))); + } + let env_value = std::env::var(&env_var).map_err(|_| { + invalid_params(format!( + "http/request header {} requires executor environment variable {env_var}", + header.name + )) + })?; + if env_value.is_empty() { + return Err(invalid_params(format!( + "http/request header {} requires a non-empty executor environment variable {env_var}", + header.name + ))); + } + format!("{}{env_value}", header.value) + } + None => header.value, + }; + let value = HeaderValue::from_str(&value).map_err(|error| { + invalid_params(format!( + "http/request header value is invalid for {}: {error}", + header.name + )) + })?; + header_map.append(name, value); + } + Ok(header_map) + } + + fn response_headers(headers: &HeaderMap) -> Vec { + headers + .iter() + .filter_map(|(name, value)| { + Some(HttpHeader { + name: name.as_str().to_string(), + value: value.to_str().ok()?.to_string(), + value_env_var: None, + }) + }) + .collect() + } +} + +fn log_send_error(method: &Method, error: RouteAwareRequestError) { + let error_is_timeout = error.is_timeout(); + let error_is_connect = error.is_connect(); + let error = match error { + RouteAwareRequestError::Request(error) => error.without_url().to_string(), + error => error.to_string(), + }; + tracing::warn!( + http_method = method.as_str(), + error_is_timeout, + error_is_connect, + error = %error, + "http/request send failed" + ); +} diff --git a/codex-rs/exec-server/src/client/rpc_http_client.rs b/codex-rs/exec-server/src/client/rpc_http_client.rs new file mode 100644 index 0000000000000000000000000000000000000000..c48ff8d29f864b0e159a0481e9fd8724aed5a0f6 --- /dev/null +++ b/codex-rs/exec-server/src/client/rpc_http_client.rs @@ -0,0 +1,92 @@ +//! JSON-RPC-backed `HttpClient` implementation. +//! +//! This code runs in the orchestrator process. It does not issue network +//! requests directly; instead it forwards `http/request` to the remote runtime +//! and then reconstructs streamed bodies from `http/request/bodyDelta` +//! notifications on the shared connection. + +use std::sync::Arc; + +use futures::FutureExt; +use futures::future::BoxFuture; +use tokio::sync::mpsc; + +use super::HttpResponseBodyStream; +use super::response_body_stream::HttpBodyStreamRegistration; +use crate::HttpClient; +use crate::client::ExecServerClient; +use crate::client::ExecServerError; +use crate::protocol::HTTP_REQUEST_METHOD; +use crate::protocol::HttpRequestParams; +use crate::protocol::HttpRequestResponse; + +/// Maximum queued body frames per streamed HTTP response. +const HTTP_BODY_DELTA_CHANNEL_CAPACITY: usize = 256; + +impl ExecServerClient { + /// Performs an HTTP request and buffers the response body. + pub async fn http_request( + &self, + mut params: HttpRequestParams, + ) -> Result { + params.stream_response = false; + self.call(HTTP_REQUEST_METHOD, ¶ms).await + } + + /// Performs an HTTP request and returns a body stream. + /// + /// The method sets `stream_response` and replaces any caller-supplied + /// `request_id` with a connection-local id, so late deltas from abandoned + /// streams cannot be confused with later requests. + pub async fn http_request_stream( + &self, + mut params: HttpRequestParams, + ) -> Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError> { + let rpc_client = self.rpc_client().await?; + params.stream_response = true; + let request_id = self.inner.next_http_body_stream_request_id(); + params.request_id = request_id.clone(); + let (tx, rx) = mpsc::channel(HTTP_BODY_DELTA_CHANNEL_CAPACITY); + self.inner + .insert_http_body_stream(request_id.clone(), tx) + .await?; + let mut registration = + HttpBodyStreamRegistration::new(Arc::clone(&self.inner), request_id.clone()); + let response = match self + .call_rpc(&rpc_client, HTTP_REQUEST_METHOD, ¶ms) + .await + { + Ok(response) => response, + Err(error) => { + self.inner.remove_http_body_stream(&request_id).await; + registration.disarm(); + return Err(error); + } + }; + registration.disarm(); + Ok(( + response, + HttpResponseBodyStream::remote(Arc::clone(&self.inner), request_id, rx), + )) + } +} + +impl HttpClient for ExecServerClient { + /// Orchestrator-side adapter that forwards buffered HTTP requests to the + /// remote runtime over the shared JSON-RPC connection. + fn http_request( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + async move { ExecServerClient::http_request(self, params).await }.boxed() + } + + /// Orchestrator-side adapter that forwards streamed HTTP requests to the + /// remote runtime and exposes body deltas as a byte stream. + fn http_request_stream( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + async move { ExecServerClient::http_request_stream(self, params).await }.boxed() + } +} diff --git a/codex-rs/exec-server/src/client/tests/network_policy_tests.rs b/codex-rs/exec-server/src/client/tests/network_policy_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..1b3c20539a5daaec1c6ef42c61b02deca65e0d2b --- /dev/null +++ b/codex-rs/exec-server/src/client/tests/network_policy_tests.rs @@ -0,0 +1,544 @@ +use std::sync::Arc; +use std::sync::Mutex; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCNotification; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::RequestId; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_network_proxy::NetworkDecision; +use codex_network_proxy::NetworkPolicyDecider; +use codex_network_proxy::NetworkPolicyRequest; +use codex_network_proxy::NetworkProxyAuditMetadata; +use codex_utils_path_uri::PathUri; +use http::HeaderMap; +use opentelemetry::trace::TracerProvider as _; +use opentelemetry_sdk::trace::InMemorySpanExporter; +use opentelemetry_sdk::trace::SdkTracerProvider; +use pretty_assertions::assert_eq; +use tokio::net::TcpListener; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::time::timeout; +use tracing::instrument::WithSubscriber; +use tracing_subscriber::filter::filter_fn; +use tracing_subscriber::prelude::*; + +use super::super::LazyRemoteExecServerClient; +use super::super::NetworkPolicyAuditContext; +use super::super::NetworkPolicyDecisionController; +use super::super::SessionState; +use super::super::handle_server_notification; +use super::accept_websocket; +use super::complete_websocket_initialize; +use super::read_jsonrpc_websocket; +use super::write_jsonrpc_websocket; +use crate::ProcessId; +use crate::client_api::ExecServerTransportParams; +use crate::protocol::EXEC_METHOD; +use crate::protocol::EXEC_TERMINATE_METHOD; +use crate::protocol::ExecParams; +use crate::protocol::ExecServerNetworkPolicyDecision; +use crate::protocol::ExecServerNetworkPolicyRequest; +use crate::protocol::ExecServerNetworkProtocol; +use crate::protocol::NETWORK_POLICY_DECISION_METHOD; +use crate::protocol::NETWORK_POLICY_REQUEST_METHOD; +use crate::protocol::NetworkPolicyDecisionNotification; +use crate::protocol::NetworkPolicyRequestParams; +use crate::protocol::NetworkPolicyRequestResponse; +use crate::rpc_server_requests::MAX_IN_FLIGHT_SERVER_CALLS; + +struct PendingDecisionGuard(mpsc::UnboundedSender<()>); + +impl Drop for PendingDecisionGuard { + fn drop(&mut self) { + let _ = self.0.send(()); + } +} + +fn policy_request(request_id: i64, process_id: ProcessId, host: &str) -> JSONRPCMessage { + JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(request_id), + method: NETWORK_POLICY_REQUEST_METHOD.to_string(), + params: Some( + serde_json::to_value(NetworkPolicyRequestParams { + process_id, + request: ExecServerNetworkPolicyRequest { + protocol: ExecServerNetworkProtocol::HttpsConnect, + host: host.to_string(), + port: 443, + }, + }) + .expect("policy request should serialize"), + ), + trace: None, + }) +} + +async fn read_decision( + websocket: &mut tokio_tungstenite::WebSocketStream, + request_id: i64, +) -> ExecServerNetworkPolicyDecision { + let JSONRPCMessage::Response(response) = read_jsonrpc_websocket(websocket).await else { + panic!("expected network policy response"); + }; + assert_eq!(response.id, RequestId::Integer(request_id)); + serde_json::from_value::(response.result) + .expect("policy response should deserialize") + .decision +} + +#[tokio::test(flavor = "current_thread")] +async fn policy_decisions_reject_forged_process_and_use_trusted_controller_metadata() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!("ws://{}", listener.local_addr().expect("listener address")); + let (release_tx, release_rx) = oneshot::channel(); + let (initialized_tx, initialized_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let mut websocket = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut websocket, + "audit-session", + /*expected_resume_session_id*/ None, + ) + .await; + initialized_tx + .send(()) + .expect("client should await completed WebSocket initialization"); + release_rx.await.expect("server should be released"); + }); + + let logs = Arc::new(Mutex::new(Vec::new())); + let writer_logs = Arc::clone(&logs); + let subscriber = tracing_subscriber::registry().with( + tracing_subscriber::fmt::layer() + .with_ansi(false) + .with_writer(move || AuditLogWriter(Arc::clone(&writer_logs))), + ); + async move { + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::websocket_url(websocket_url, Duration::from_secs(1)), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .get() + .await + .expect("client should connect"); + initialized_rx + .await + .expect("server should complete WebSocket initialization"); + let mut state = SessionState::new(/*recoverable*/ true); + state.network_policy.audit = Some(NetworkPolicyAuditContext { + metadata: NetworkProxyAuditMetadata { + conversation_id: Some("trusted-conversation".to_string()), + user_account_id: Some("trusted-account".to_string()), + ..NetworkProxyAuditMetadata::default() + }, + execution_id: Some("trusted-execution".to_string()), + }); + client + .inner + .insert_session(&ProcessId::from("trusted-process"), Arc::new(state)) + .expect("trusted process should register"); + for (process_id, host) in [ + ("forged-process", "forged.example"), + ("trusted-process", "trusted.example"), + ] { + handle_server_notification( + &client.inner, + JSONRPCNotification { + method: NETWORK_POLICY_DECISION_METHOD.to_string(), + params: Some( + serde_json::to_value(NetworkPolicyDecisionNotification { + process_id: ProcessId::from(process_id), + timestamp: "2026-08-11T12:00:00.000Z".to_string(), + scope: "domain".to_string(), + decision: "deny".to_string(), + source: "baseline_policy".to_string(), + reason: "not_allowed".to_string(), + protocol: ExecServerNetworkProtocol::HttpsConnect, + host: host.to_string(), + port: 443, + method: None, + client: None, + policy_override: false, + }) + .expect("network policy decision should serialize"), + ), + }, + ) + .await + .expect("controller should handle network policy notification"); + } + let output = String::from_utf8( + logs.lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone(), + ) + .expect("audit log should be UTF-8"); + assert!(!output.contains("forged.example")); + for expected in [ + "codex_otel.log_only", + "trusted-conversation", + "trusted-account", + "trusted-execution", + ] { + assert!( + output.contains(expected), + "missing `{expected}` in {output}" + ); + } + release_tx.send(()).expect("server should be released"); + } + .with_subscriber(subscriber) + .await; + server.await.expect("server should finish"); +} + +struct AuditLogWriter(Arc>>); + +impl std::io::Write for AuditLogWriter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.0 + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .extend_from_slice(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +#[tokio::test] +async fn abandoned_process_start_unregisters_and_cleans_up() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!("ws://{}", listener.local_addr().expect("listener address")); + let (start_seen_tx, start_seen_rx) = oneshot::channel(); + let (finish_start_tx, finish_start_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let mut websocket = accept_websocket(&listener).await; + complete_websocket_initialize(&mut websocket, "p", Default::default()).await; + let JSONRPCMessage::Request(start) = read_jsonrpc_websocket(&mut websocket).await else { + panic!("expected process start request"); + }; + assert_eq!(start.method, EXEC_METHOD); + start_seen_tx.send(()).expect("start should be observed"); + finish_start_rx.await.expect("start should be released"); + write_jsonrpc_websocket( + &mut websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id: start.id, + result: serde_json::json!({"processId": "pending-start"}), + }), + ) + .await; + let JSONRPCMessage::Request(terminate) = read_jsonrpc_websocket(&mut websocket).await + else { + panic!("expected process terminate request"); + }; + assert_eq!(terminate.method, EXEC_TERMINATE_METHOD); + write_jsonrpc_websocket( + &mut websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id: terminate.id, + result: serde_json::json!({"running": true}), + }), + ) + .await; + }); + + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::websocket_url(websocket_url, Duration::from_secs(1)), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .get() + .await + .expect("client should connect"); + let process_id = ProcessId::from("pending-start"); + let start_client = client.clone(); + let start_process_id = process_id.clone(); + let start = tokio::spawn(async move { + let params = ExecParams { + metadata: Default::default(), + process_id: start_process_id, + argv: vec!["true".to_string()], + cwd: PathUri::from_host_native_path(std::env::current_dir().expect("cwd")) + .expect("cwd URI"), + shell_snapshot: None, + env_policy: None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }; + start_client + .start_process(params, /*network_policy_decider*/ None) + .await + }); + start_seen_rx.await.expect("start should be observed"); + let state = client + .inner + .get_session(&process_id) + .expect("pending process should be registered"); + let decider: Arc = + Arc::new(|_request: NetworkPolicyRequest| async { NetworkDecision::Allow }); + let decider_weak = Arc::downgrade(&decider); + state + .network_policy + .controller + .store(Some(Arc::new(NetworkPolicyDecisionController { + decider, + timeout: Duration::from_secs(30), + }))); + + start.abort(); + assert!(start.await.is_err_and(|error| error.is_cancelled())); + assert!(state.network_policy.cancelled.is_cancelled()); + assert!(client.inner.get_session(&process_id).is_none()); + assert!(decider_weak.upgrade().is_none()); + + finish_start_tx.send(()).expect("start should be released"); + server.await.expect("server task should finish"); +} + +#[tokio::test] +async fn policy_requests_use_process_decider_and_cancel_on_unregister() { + let span_exporter = InMemorySpanExporter::default(); + let tracer_provider = SdkTracerProvider::builder() + .with_simple_exporter(span_exporter.clone()) + .build(); + let subscriber = tracing_subscriber::registry().with( + tracing_opentelemetry::layer() + .with_tracer(tracer_provider.tracer("exec-server-test")) + .with_filter(filter_fn(codex_otel::OtelProvider::trace_export_filter)), + ); + let _subscriber = tracing::subscriber::set_default(subscriber); + tracing::callsite::rebuild_interest_cache(); + + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!( + "ws://{}", + listener.local_addr().expect("listener should have address") + ); + let process_id = ProcessId::from("policy-process"); + let server_process_id = process_id.clone(); + let (ready_tx, ready_rx) = oneshot::channel(); + let (overflow_checked_tx, overflow_checked_rx) = oneshot::channel(); + let (unregistered_tx, unregistered_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let mut websocket = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut websocket, + "policy-session", + /*expected_resume_session_id*/ None, + ) + .await; + ready_rx.await.expect("process should be registered"); + + for (request_id, host, expected) in [ + (0, "allowed.example", ExecServerNetworkPolicyDecision::Allow), + ( + 2, + "denied.example", + ExecServerNetworkPolicyDecision::Deny { + reason: "blocked".to_string(), + }, + ), + ( + 3, + "invalid host", + ExecServerNetworkPolicyDecision::Deny { + reason: "not_allowed".to_string(), + }, + ), + ] { + write_jsonrpc_websocket( + &mut websocket, + policy_request(request_id, server_process_id.clone(), host), + ) + .await; + assert_eq!(read_decision(&mut websocket, request_id).await, expected); + } + + let first_pending_request_id = 100; + for offset in 0..MAX_IN_FLIGHT_SERVER_CALLS { + write_jsonrpc_websocket( + &mut websocket, + policy_request( + first_pending_request_id + offset as i64, + server_process_id.clone(), + "pending.example", + ), + ) + .await; + } + let overflow_request_id = first_pending_request_id + MAX_IN_FLIGHT_SERVER_CALLS as i64; + write_jsonrpc_websocket( + &mut websocket, + policy_request( + overflow_request_id, + server_process_id.clone(), + "pending.example", + ), + ) + .await; + assert_eq!( + read_decision(&mut websocket, overflow_request_id).await, + ExecServerNetworkPolicyDecision::Deny { + reason: "not_allowed".to_string(), + } + ); + overflow_checked_tx.send(()).expect("overflow observed"); + + unregistered_rx + .await + .expect("process should be unregistered"); + + let post_unregister_request_id = 900; + write_jsonrpc_websocket( + &mut websocket, + policy_request( + post_unregister_request_id, + server_process_id, + "allowed.example", + ), + ) + .await; + assert_eq!( + read_decision(&mut websocket, post_unregister_request_id).await, + ExecServerNetworkPolicyDecision::Deny { + reason: "not_allowed".to_string(), + } + ); + }); + + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::WebSocketUrl { + websocket_url, + connect_timeout: Duration::from_secs(1), + initialize_timeout: Duration::from_secs(1), + http_headers: HeaderMap::new(), + }, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .get() + .await + .expect("client should connect"); + let session = client + .register_session(&process_id) + .await + .expect("session should register"); + let (started_tx, mut started_rx) = mpsc::unbounded_channel(); + let (dropped_tx, mut dropped_rx) = mpsc::unbounded_channel(); + let decider: Arc = Arc::new(move |request: NetworkPolicyRequest| { + let started_tx = started_tx.clone(); + let dropped_tx = dropped_tx.clone(); + async move { + assert_eq!( + tracing::Span::current() + .metadata() + .map(tracing::Metadata::name), + Some("codex.exec_server.request"), + "network policy decisions must run inside the inbound request span" + ); + match request.host.as_str() { + "allowed.example" => NetworkDecision::Allow, + "denied.example" => NetworkDecision::deny("blocked"), + "pending.example" => { + started_tx.send(()).expect("decision should start"); + let _drop_guard = PendingDecisionGuard(dropped_tx); + std::future::pending().await + } + host => panic!("unexpected policy host: {host}"), + } + } + }); + session.state.network_policy.controller.store(Some(Arc::new( + NetworkPolicyDecisionController { + decider, + timeout: Duration::from_secs(30), + }, + ))); + ready_tx.send(()).expect("server should be waiting"); + timeout(Duration::from_secs(5), async { + for _ in 0..MAX_IN_FLIGHT_SERVER_CALLS { + started_rx + .recv() + .await + .expect("pending decision should start"); + } + }) + .await + .expect("pending decisions should start"); + overflow_checked_rx + .await + .expect("overflow should be observed"); + session.unregister().await; + timeout(Duration::from_secs(5), async { + for _ in 0..MAX_IN_FLIGHT_SERVER_CALLS { + dropped_rx + .recv() + .await + .expect("unregistered decision should be dropped"); + } + }) + .await + .expect("unregistered decisions should be cancelled"); + unregistered_tx + .send(()) + .expect("server should verify late responses"); + timeout(Duration::from_secs(2), server) + .await + .expect("policy routing should finish") + .expect("server task should finish"); + + tracer_provider.force_flush().expect("flush traces"); + let spans = span_exporter.get_finished_spans().expect("span export"); + let policy_spans = spans + .iter() + .filter(|span| span.name.as_ref() == NETWORK_POLICY_REQUEST_METHOD) + .collect::>(); + assert!( + !policy_spans.is_empty(), + "network policy requests should export server spans" + ); + let outcomes = policy_spans + .iter() + .map(|span| { + span.attributes + .iter() + .find(|attribute| attribute.key.as_str() == "result") + .map(|attribute| attribute.value.as_str().into_owned()) + }) + .collect::>(); + assert!( + outcomes.iter().all(Option::is_some), + "completed, rejected, and cancelled policy requests must all record an outcome" + ); + assert!( + outcomes + .iter() + .any(|outcome| outcome.as_deref() == Some("success")), + "completed and capacity-rejected requests should record successful responses" + ); + assert!( + outcomes + .iter() + .any(|outcome| outcome.as_deref() == Some("disconnected")), + "cancelled requests should record disconnection" + ); +} diff --git a/codex-rs/exec-server/src/client_api.rs b/codex-rs/exec-server/src/client_api.rs new file mode 100644 index 0000000000000000000000000000000000000000..e870f124717a7a52624bfe0840d4d55f2822e5b7 --- /dev/null +++ b/codex-rs/exec-server/src/client_api.rs @@ -0,0 +1,187 @@ +use std::collections::HashMap; +use std::path::PathBuf; +use std::sync::Arc; +use std::time::Duration; + +use codex_http_client::HttpClientFactory; +use futures::future::BoxFuture; +use http::HeaderMap; +use tokio::sync::watch; + +use crate::ExecServerError; +use crate::HttpRequestParams; +use crate::HttpRequestResponse; +use crate::HttpResponseBodyStream; +use crate::NoiseChannelIdentity; +use crate::NoiseChannelPublicKey; + +pub(crate) const DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT: Duration = Duration::from_secs(10); +pub(crate) const DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT: Duration = Duration::from_secs(10); + +/// Connection options for any exec-server client transport. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ExecServerClientConnectOptions { + pub client_name: String, + pub initialize_timeout: Duration, + pub resume_session_id: Option, +} + +/// WebSocket connection arguments for a remote exec-server. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RemoteExecServerConnectArgs { + pub websocket_url: String, + pub client_name: String, + pub connect_timeout: Duration, + pub initialize_timeout: Duration, + pub resume_session_id: Option, + pub http_client_factory: HttpClientFactory, +} + +/// Registry-authorized material for one Noise rendezvous connection attempt. +/// +/// Treat this as an atomic, single-use bundle. The URL authorization, executor +/// registration, pinned executor key, and harness-key authorization describe one +/// physical connection attempt and must not be mixed with values from another +/// registry response. +pub struct NoiseRendezvousConnectBundle { + pub websocket_url: String, + pub environment_id: String, + pub executor_registration_id: String, + pub executor_public_key: NoiseChannelPublicKey, + pub harness_key_authorization: String, +} + +/// Connection arguments for an authenticated Noise rendezvous exec-server. +/// +/// `harness_identity` identifies the logical harness endpoint and may be reused +/// across reconnects. In contrast, callers must supply a fresh +/// [`NoiseRendezvousConnectBundle`] for each physical connection attempt. +pub struct NoiseRendezvousConnectArgs { + pub bundle: NoiseRendezvousConnectBundle, + pub harness_identity: NoiseChannelIdentity, + pub client_name: String, + pub connect_timeout: Duration, + pub initialize_timeout: Duration, + pub resume_session_id: Option, + pub http_client_factory: HttpClientFactory, +} + +/// Supplies fresh registry-authorized material for Noise rendezvous connections. +pub trait NoiseRendezvousConnectProvider: Send + Sync { + /// Fetch a bundle authorizing this harness key for one physical connection. + fn connect_bundle( + &self, + harness_public_key: NoiseChannelPublicKey, + ) -> BoxFuture<'_, Result>; +} + +/// Stdio connection arguments for a command-backed exec-server. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct StdioExecServerConnectArgs { + pub command: StdioExecServerCommand, + pub client_name: String, + pub initialize_timeout: Duration, + pub resume_session_id: Option, +} + +/// Structured process command used to start an exec-server over stdio. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct StdioExecServerCommand { + pub program: String, + pub args: Vec, + pub env: HashMap, + pub cwd: Option, +} + +pub(crate) type DeferredEnvironmentReadiness = watch::Receiver>>; + +#[derive(Clone)] +pub(crate) struct Deferred { + pub readiness: DeferredEnvironmentReadiness, + pub transport: T, +} + +/// Parameters used to connect to a remote exec-server environment. +#[derive(Clone)] +pub(crate) enum ExecServerTransportParams { + Deferred(Box>), + WebSocketUrl { + websocket_url: String, + connect_timeout: Duration, + initialize_timeout: Duration, + http_headers: HeaderMap, + }, + NoiseRendezvous { + provider: Arc, + identity: NoiseChannelIdentity, + }, + #[allow(dead_code)] + StdioCommand { + command: StdioExecServerCommand, + initialize_timeout: Duration, + }, +} + +impl std::fmt::Debug for ExecServerTransportParams { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Deferred(deferred) => f + .debug_struct("Deferred") + .field("transport", &deferred.transport) + .finish_non_exhaustive(), + Self::WebSocketUrl { + websocket_url, + connect_timeout, + initialize_timeout, + .. + } => f + .debug_struct("WebSocketUrl") + .field("websocket_url", websocket_url) + .field("connect_timeout", connect_timeout) + .field("initialize_timeout", initialize_timeout) + .field("http_headers", &"") + .finish(), + Self::NoiseRendezvous { .. } => { + f.debug_struct("NoiseRendezvous").finish_non_exhaustive() + } + Self::StdioCommand { + command, + initialize_timeout, + } => f + .debug_struct("StdioCommand") + .field("command", command) + .field("initialize_timeout", initialize_timeout) + .finish(), + } + } +} + +impl ExecServerTransportParams { + pub(crate) fn websocket_url(websocket_url: String, connect_timeout: Duration) -> Self { + Self::WebSocketUrl { + websocket_url, + connect_timeout, + initialize_timeout: DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT, + http_headers: HeaderMap::new(), + } + } +} + +/// Sends HTTP requests through a runtime-selected transport. +/// +/// This is the HTTP capability counterpart to [`crate::ExecBackend`]. Callers +/// use it when they need environment-owned network requests but should not +/// depend on the concrete connection type or how that connection is established. +pub trait HttpClient: Send + Sync { + /// Perform an HTTP request and buffer the response body. + fn http_request( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result>; + + /// Perform an HTTP request and return a streamed body handle. + fn http_request_stream( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>>; +} diff --git a/codex-rs/exec-server/src/client_recovery.rs b/codex-rs/exec-server/src/client_recovery.rs new file mode 100644 index 0000000000000000000000000000000000000000..6b4cf61df1dc5ca0795499561440e6de6a6341ac --- /dev/null +++ b/codex-rs/exec-server/src/client_recovery.rs @@ -0,0 +1,900 @@ +use std::collections::hash_map::DefaultHasher; +use std::hash::Hash; +use std::hash::Hasher; +use std::sync::Arc; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_network_proxy::NetworkDecision; +use codex_network_proxy::NetworkPolicyDecision; +use codex_network_proxy::NetworkPolicyRequest; +use codex_network_proxy::NetworkProtocol; +use codex_network_proxy::NetworkRequestCancellation; +use codex_network_proxy::NetworkRequestCancellationReason; +use serde_json::Value; +use tokio::sync::mpsc; +use tokio::time::Instant; +use tokio::time::sleep; +use tokio::time::timeout; +use tokio::time::timeout_at; +use tokio_util::sync::CancellationToken; +use tracing::Instrument; +use tracing::debug; + +use super::ConnectionStatus; +use super::ExecServerClient; +use super::ExecServerError; +use super::Inner; +use super::OrderedSessionEvents; +use super::RecoveryPolicy; +use super::SessionState; +use super::disconnected_message; +use super::fail_all_in_flight_work; +use super::handle_server_notification; +use super::is_transport_closed_error; +use crate::client_transport::ExecServerReconnectStrategy; +use crate::process::ExecProcessEvent; +use crate::protocol::EXEC_READ_METHOD; +use crate::protocol::EXEC_TERMINATE_METHOD; +use crate::protocol::ExecServerNetworkPolicyDecision; +use crate::protocol::ExecServerNetworkProtocol; +use crate::protocol::MAX_NETWORK_POLICY_HOST_BYTES; +use crate::protocol::MAX_NETWORK_POLICY_PROCESS_ID_BYTES; +use crate::protocol::MAX_NETWORK_POLICY_REASON_BYTES; +use crate::protocol::NETWORK_POLICY_REQUEST_METHOD; +use crate::protocol::NetworkPolicyRequestParams; +use crate::protocol::NetworkPolicyRequestResponse; +use crate::protocol::ReadParams; +use crate::protocol::ReadResponse; +use crate::protocol::TerminateParams; +use crate::protocol::TerminateResponse; +use crate::rpc::RpcClient; +use crate::rpc::RpcClientEvent; +use crate::rpc::RpcInboundRequestAdmissionError; +use crate::rpc::SESSION_ALREADY_ATTACHED_ERROR_CODE; +use crate::rpc::invalid_params; +use crate::rpc::method_not_found; + +#[cfg(test)] +const SESSION_RECOVERY_TIMEOUT: Duration = Duration::from_millis(500); +#[cfg(not(test))] +// Leave margin inside the server's 30-second retention windows because the +// client and server start their disconnect clocks independently. +const SESSION_RECOVERY_TIMEOUT: Duration = Duration::from_secs(25); +const SESSION_RECOVERY_RETRY_INTERVAL: Duration = Duration::from_millis(100); +const REGISTRY_RECOVERY_INITIAL_RETRY_INTERVAL: Duration = Duration::from_millis(500); +const REGISTRY_RECOVERY_MAX_RETRY_INTERVAL: Duration = Duration::from_secs(5); +const NETWORK_POLICY_DENIAL_REASON: &str = "not_allowed"; + +struct ClientRequestOutcome { + span: tracing::Span, + result: &'static str, +} + +impl ClientRequestOutcome { + fn complete(&mut self, result: &'static str) { + self.result = result; + } +} + +impl Drop for ClientRequestOutcome { + fn drop(&mut self) { + self.span.record("result", self.result); + } +} + +impl SessionState { + fn last_published_seq(&self) -> u64 { + self.ordered_events + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .last_published_seq + } + + fn recover_events(&self, response: ReadResponse) -> Result { + let ReadResponse { + chunks, + next_seq, + exited, + exit_code, + closed, + failure, + sandbox_denied, + } = response; + if let Some(message) = failure { + return Err(ExecServerError::Protocol(format!( + "process failed while recovering: {message}" + ))); + } + + let target_seq = next_seq.saturating_sub(1); + let published_closed = { + let mut ordered_events = self + .ordered_events + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if ordered_events.failure.is_some() + || ordered_events.closed_published + || target_seq <= ordered_events.last_published_seq + { + return Ok(false); + } + let pending_exit = ordered_events.pending.range_mut(..=target_seq).find_map( + |(_, event)| match event { + ExecProcessEvent::Exited { + sandbox_denied: pending_sandbox_denied, + .. + } => Some(pending_sandbox_denied), + _ => None, + }, + ); + let exit_pending = pending_exit.is_some(); + if let Some(pending_sandbox_denied) = pending_exit { + *pending_sandbox_denied = + Some(pending_sandbox_denied.unwrap_or(false) || sandbox_denied); + } + let mut exit_known = ordered_events.exit_published || exit_pending; + if closed + && (matches!( + ordered_events.pending.get(&target_seq), + Some(event) if !matches!(event, ExecProcessEvent::Closed { .. }) + ) || chunks.iter().any(|chunk| chunk.seq == target_seq)) + { + return Err(ExecServerError::Protocol(format!( + "process close sequence {target_seq} conflicts with recovered output" + ))); + } + let mut published_closed = false; + for chunk in chunks { + if chunk.seq > target_seq { + return Err(ExecServerError::Protocol(format!( + "recovered process output sequence {} exceeds target sequence {target_seq}", + chunk.seq + ))); + } + let next_seq = ordered_events.last_published_seq.saturating_add(1); + if exited && !exit_known && chunk.seq > next_seq { + let exit_code = exit_code.ok_or_else(|| { + ExecServerError::Protocol( + "recovering exited process did not include its exit code".to_string(), + ) + })?; + ordered_events + .insert_pending(ExecProcessEvent::Exited { + seq: next_seq, + exit_code, + sandbox_denied: Some(sandbox_denied), + }) + .map_err(ExecServerError::Protocol)?; + published_closed |= self.publish_ready(&mut ordered_events); + exit_known = true; + } + if chunk.seq > ordered_events.last_published_seq { + ordered_events + .insert_pending(ExecProcessEvent::Output(chunk)) + .map_err(ExecServerError::Protocol)?; + published_closed |= self.publish_ready(&mut ordered_events); + } + } + if closed + && !ordered_events.closed_published + && !matches!( + ordered_events.pending.get(&target_seq), + Some(ExecProcessEvent::Closed { .. }) + ) + { + ordered_events + .insert_pending(ExecProcessEvent::Closed { seq: target_seq }) + .map_err(ExecServerError::Protocol)?; + } + + let event_count = target_seq.saturating_sub(ordered_events.last_published_seq); + let first_unpublished_seq = ordered_events.last_published_seq.saturating_add(1); + let retained_count = if first_unpublished_seq <= target_seq { + ordered_events + .pending + .range(first_unpublished_seq..=target_seq) + .count() as u64 + } else { + 0 + }; + let missing_count = event_count.saturating_sub(retained_count); + if exited && !exit_known { + if missing_count != 1 { + return Err(recovery_gap_error(target_seq)); + } + let seq = first_missing_seq(&ordered_events, target_seq); + let exit_code = exit_code.ok_or_else(|| { + ExecServerError::Protocol( + "recovering exited process did not include its exit code".to_string(), + ) + })?; + ordered_events + .insert_pending(ExecProcessEvent::Exited { + seq, + exit_code, + sandbox_denied: Some(sandbox_denied), + }) + .map_err(ExecServerError::Protocol)?; + } else if missing_count != 0 { + return Err(recovery_gap_error(target_seq)); + } + published_closed |= self.publish_ready(&mut ordered_events); + published_closed + }; + + self.note_change(target_seq); + Ok(published_closed) + } +} + +fn first_missing_seq(events: &OrderedSessionEvents, target_seq: u64) -> u64 { + let mut expected = events.last_published_seq.saturating_add(1); + for seq in events + .pending + .range(expected..=target_seq) + .map(|(seq, _)| *seq) + { + if seq != expected { + break; + } + expected = expected.saturating_add(1); + } + expected +} + +fn recovery_gap_error(target_seq: u64) -> ExecServerError { + ExecServerError::Protocol(format!( + "process events are no longer retained while recovering through sequence {target_seq}" + )) +} + +impl Inner { + pub(super) async fn rpc_client(self: &Arc) -> Result, ExecServerError> { + let mut connection_changed = self.connection_changed.subscribe(); + loop { + if let Some(message) = self.failure_message() { + return Err(ExecServerError::Disconnected(message)); + } + + let rpc_client = { + let connection = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + match &connection.status { + ConnectionStatus::Connected(rpc_client) => Some(Arc::clone(rpc_client)), + ConnectionStatus::Recovering | ConnectionStatus::Failed(_) => None, + } + }; + let Some(rpc_client) = rpc_client else { + let _ = connection_changed.changed().await; + continue; + }; + if !rpc_client.is_disconnected() { + return Ok(rpc_client); + } + + let _ = connection_changed.changed().await; + } + } + + pub(super) fn begin_process_start(&self, expected: &Arc) -> bool { + let mut connection = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let ConnectionStatus::Connected(current) = &connection.status else { + return false; + }; + if !Arc::ptr_eq(current, expected) || expected.is_disconnected() { + return false; + } + connection.active_process_starts += 1; + true + } + + pub(super) fn finish_process_start(&self) { + { + let mut connection = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if connection.active_process_starts == 0 { + tracing::error!("finished an exec-server process start that was not active"); + return; + } + connection.active_process_starts -= 1; + } + self.notify_connection_changed(); + } + + pub(super) fn is_failed(&self) -> bool { + self.failure_message().is_some() + } + + pub(super) fn failure_message(&self) -> Option { + let connection = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + match &connection.status { + ConnectionStatus::Failed(message) => Some(message.clone()), + ConnectionStatus::Connected(_) | ConnectionStatus::Recovering => None, + } + } + + pub(super) fn request_recovery( + self: &Arc, + failed_rpc_client: Arc, + disconnect_message: String, + ) { + let should_recover = { + let mut connection = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + match &connection.status { + ConnectionStatus::Connected(current) + if Arc::ptr_eq(current, &failed_rpc_client) => + { + connection.set_status(ConnectionStatus::Recovering); + true + } + ConnectionStatus::Connected(_) + | ConnectionStatus::Recovering + | ConnectionStatus::Failed(_) => false, + } + }; + if !should_recover { + return; + } + + self.notify_connection_changed(); + let inner = Arc::clone(self); + tokio::spawn(async move { + tokio::select! { + biased; + _ = inner.retired.cancelled() => {}, + _ = inner.recover(disconnect_message) => {}, + } + }); + } + + async fn recover(self: &Arc, disconnect_message: String) { + let deadline = Instant::now() + SESSION_RECOVERY_TIMEOUT; + self.fail_all_http_body_streams(disconnect_message.clone()) + .await; + if timeout_at(deadline, self.wait_for_process_starts()) + .await + .is_err() + { + let message = format!( + "{disconnect_message}; failed to resume exec-server session: recovery timed out after {SESSION_RECOVERY_TIMEOUT:?}" + ); + self.fail(message).await; + return; + } + if self.reconnect_strategy.is_none() { + self.fail(disconnect_message).await; + return; + } + + let Some(session_id) = self.session_id.get().cloned() else { + let message = format!( + "{disconnect_message}; failed to resume exec-server session: missing session id" + ); + self.fail(message).await; + return; + }; + let uses_registry_backoff = matches!( + self.reconnect_strategy.as_ref(), + Some(ExecServerReconnectStrategy::NoiseRendezvous { .. }) + ); + let mut registry_retry_attempt = 0; + let last_error = loop { + match timeout_at(deadline, self.resume_once(&session_id)).await { + Ok(Ok((rpc_client, _attempt))) => { + if !rpc_client.is_disconnected() && self.install_recovered_client(rpc_client) { + return; + } + } + Ok(Err(error)) if !is_retryable_recovery_error(&error) => { + break error.to_string(); + } + Ok(Err(_)) => {} + Err(_) => { + break format!("recovery timed out after {SESSION_RECOVERY_TIMEOUT:?}"); + } + } + + let retry_delay = if uses_registry_backoff { + let delay = registry_recovery_retry_delay(&session_id, registry_retry_attempt); + registry_retry_attempt = registry_retry_attempt.saturating_add(1); + delay + } else { + SESSION_RECOVERY_RETRY_INTERVAL + }; + + let now = Instant::now(); + if now >= deadline { + break format!("recovery timed out after {SESSION_RECOVERY_TIMEOUT:?}"); + } + sleep(retry_delay.min(deadline - now)).await; + }; + + let message = + format!("{disconnect_message}; failed to resume exec-server session: {last_error}"); + self.fail(message).await; + } + + async fn wait_for_process_starts(&self) { + let mut connection_changed = self.connection_changed.subscribe(); + loop { + let starts_are_done = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .active_process_starts + == 0; + if starts_are_done { + return; + } + let _ = connection_changed.changed().await; + } + } + + fn install_recovered_client(&self, rpc_client: Arc) -> bool { + let installed = { + let mut connection = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if !matches!(connection.status, ConnectionStatus::Recovering) + || rpc_client.is_disconnected() + { + false + } else { + connection.set_status(ConnectionStatus::Connected(rpc_client)); + true + } + }; + if installed { + self.notify_connection_changed(); + } + installed + } + + fn notify_connection_changed(&self) { + self.connection_changed.send_replace(()); + } + + async fn resume_once( + self: &Arc, + session_id: &str, + ) -> Result<(Arc, Option), ExecServerError> { + let reconnect_strategy = self + .reconnect_strategy + .as_ref() + .ok_or_else(|| ExecServerError::Protocol("missing reconnect strategy".to_string()))?; + let attempt = reconnect_strategy.resume(session_id).await?; + let (connection, options, attempt_permit, noise_context) = attempt.into_parts(); + let (rpc_client, events_rx) = RpcClient::new(connection); + let rpc_client = Arc::new(rpc_client); + let client = ExecServerClient { + inner: Arc::clone(self), + recovery_policy: RecoveryPolicy::Wait, + }; + // Resuming a session redirects notifications from its running processes + // to this connection during initialize. Drain them immediately so a + // burst cannot fill the bounded event channel and block the initialize + // response behind it. + client.spawn_rpc_reader(&rpc_client, events_rx); + client + .initialize_rpc(&rpc_client, options, noise_context) + .await?; + + self.recover_processes(&rpc_client).await?; + Ok((rpc_client, attempt_permit)) + } + + async fn recover_processes( + self: &Arc, + rpc_client: &RpcClient, + ) -> Result<(), ExecServerError> { + let sessions = self.sessions.load_full(); + for (process_id, session) in sessions.iter() { + if !session.recoverable.load(Ordering::Acquire) { + continue; + } + let response = rpc_client + .call::<_, ReadResponse>( + EXEC_READ_METHOD, + &ReadParams { + process_id: process_id.clone(), + after_seq: Some(session.last_published_seq()), + max_bytes: None, + wait_ms: Some(0), + }, + ) + .await + .map_err(ExecServerError::from); + let recovered = match response { + Ok(response) => session.recover_events(response), + Err(error) if is_transport_closed_error(&error) => return Err(error), + Err(error) => Err(error), + }; + match recovered { + Ok(true) => self.remove_session_if(process_id, session), + Ok(false) => {} + Err(error) => { + session + .network_policy + .cancellation + .record(NetworkRequestCancellationReason::ProcessCancelled); + let terminated: Result = rpc_client + .call_for_cleanup( + EXEC_TERMINATE_METHOD, + &TerminateParams { + process_id: process_id.clone(), + }, + ) + .await + .map_err(ExecServerError::from); + if let Err(terminate_error) = terminated + && is_transport_closed_error(&terminate_error) + { + return Err(terminate_error); + } + self.remove_session_if(process_id, session); + session.set_failure(format!("failed to recover process {process_id}: {error}")); + } + } + } + Ok(()) + } + + async fn fail(self: &Arc, message: String) { + let (message, newly_failed) = { + let mut connection = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + match &connection.status { + ConnectionStatus::Failed(existing) => (existing.clone(), false), + ConnectionStatus::Connected(_) | ConnectionStatus::Recovering => { + connection.set_status(ConnectionStatus::Failed(message.clone())); + (message, true) + } + } + }; + if newly_failed { + self.notify_connection_changed(); + fail_all_in_flight_work(self, message.clone()).await; + } + } +} + +impl ExecServerClient { + pub(super) fn spawn_rpc_reader( + &self, + rpc_client: &Arc, + mut events_rx: mpsc::Receiver, + ) { + let inner = Arc::downgrade(&self.inner); + let rpc_inbound_request_slots = Arc::clone(&self.inner.rpc_inbound_request_slots); + let rpc_client = Arc::downgrade(rpc_client); + let connection_cancelled = CancellationToken::new(); + let connection_cancel_guard = connection_cancelled.clone().drop_guard(); + tokio::spawn(async move { + let _connection_cancel_guard = connection_cancel_guard; + while let Some(event) = events_rx.recv().await { + let (Some(inner), Some(rpc_client)) = (inner.upgrade(), rpc_client.upgrade()) + else { + return; + }; + match event { + RpcClientEvent::Request { + request, + request_span, + } => { + let mut request_outcome = ClientRequestOutcome { + span: request_span, + result: "disconnected", + }; + if request.method != NETWORK_POLICY_REQUEST_METHOD { + let error = method_not_found(format!( + "exec-server client does not implement `{}` yet", + request.method + )); + if rpc_client.respond_error(request.id, error).await.is_err() { + inner.request_recovery( + rpc_client, + disconnected_message(/*reason*/ None), + ); + return; + } + request_outcome.complete("error"); + continue; + } + request_outcome + .span + .record("otel.name", NETWORK_POLICY_REQUEST_METHOD); + + let request_guard = match rpc_client + .admit_inbound_request(&request.id, &rpc_inbound_request_slots) + { + Ok(request_guard) => request_guard, + Err(RpcInboundRequestAdmissionError::InvalidRequestId) => { + rpc_client.close_transport().await; + inner.request_recovery( + rpc_client, + "exec-server sent an invalid request ID".to_string(), + ); + return; + } + Err(RpcInboundRequestAdmissionError::DuplicateRequestId) => { + rpc_client.close_transport().await; + inner.request_recovery( + rpc_client, + "exec-server reused an in-flight request ID".to_string(), + ); + return; + } + Err(RpcInboundRequestAdmissionError::AtCapacity) => { + let response = NetworkPolicyRequestResponse { + decision: ExecServerNetworkPolicyDecision::Deny { + reason: NETWORK_POLICY_DENIAL_REASON.to_string(), + }, + }; + if rpc_client.respond(request.id, &response).await.is_err() { + inner.request_recovery( + rpc_client, + disconnected_message(/*reason*/ None), + ); + return; + } + request_outcome.complete("success"); + continue; + } + }; + let request_id = request.id; + let params: NetworkPolicyRequestParams = + match serde_json::from_value(request.params.unwrap_or(Value::Null)) { + Ok(params) => params, + Err(_) => { + let error = invalid_params( + "invalid network policy request params".to_string(), + ); + if rpc_client.respond_error(request_id, error).await.is_err() { + inner.request_recovery( + rpc_client, + disconnected_message(/*reason*/ None), + ); + return; + } + request_outcome.complete("error"); + continue; + } + }; + let process_id = params.process_id; + let request = params.request; + let process_id_valid = !process_id.is_empty() + && process_id.len() <= MAX_NETWORK_POLICY_PROCESS_ID_BYTES; + let host_valid = !request.host.is_empty() + && request.host.len() <= MAX_NETWORK_POLICY_HOST_BYTES + && !request.host.chars().any(char::is_control) + && !request.host.chars().any(char::is_whitespace); + let session = (process_id_valid && host_valid) + .then(|| inner.get_session(&process_id)) + .flatten(); + let controller = session + .as_ref() + .and_then(|session| session.network_policy.controller.load_full()); + let process_cancelled = session + .as_ref() + .map(|session| session.network_policy.cancelled.clone()); + let process_cancellation = session + .as_ref() + .map(|session| session.network_policy.cancellation.clone()); + let cancellation = NetworkRequestCancellation::default(); + let expected_session = session.as_ref().map(Arc::downgrade); + let policy_request = + (process_id_valid && host_valid).then_some(NetworkPolicyRequest { + protocol: match request.protocol { + ExecServerNetworkProtocol::Http => NetworkProtocol::Http, + ExecServerNetworkProtocol::HttpsConnect => { + NetworkProtocol::HttpsConnect + } + ExecServerNetworkProtocol::Socks5Tcp => { + NetworkProtocol::Socks5Tcp + } + ExecServerNetworkProtocol::Socks5Udp => { + NetworkProtocol::Socks5Udp + } + }, + host: request.host, + port: request.port, + environment_id: None, + client_addr: None, + method: None, + command: None, + exec_policy_hint: None, + execution_id: None, + disconnect: None, + cancellation: Some(cancellation.clone()), + }); + let inner = Arc::downgrade(&inner); + let rpc_client = Arc::downgrade(&rpc_client); + let connection_cancelled = connection_cancelled.clone(); + let task_span = request_outcome.span.clone(); + let task = async move { + let _request_guard = request_guard; + let decision = match (controller, policy_request, process_cancelled) { + (Some(controller), Some(request), Some(process_cancelled)) => { + // Keep the decision future outside select/timeout so its + // guard sees the cancellation cause before it is dropped. + let mut decision = controller.decider.decide(request); + tokio::select! { + biased; + _ = connection_cancelled.cancelled() => { + cancellation.record(NetworkRequestCancellationReason::ConnectionClosed); + return; + }, + _ = process_cancelled.cancelled() => { + cancellation.record(process_cancellation.as_ref() + .and_then(NetworkRequestCancellation::reason) + .unwrap_or(NetworkRequestCancellationReason::ProcessCancelled)); + NetworkDecision::deny(NETWORK_POLICY_DENIAL_REASON) + } + result = timeout( + controller.timeout, + &mut decision, + ) => result.unwrap_or_else(|_| { + cancellation.record(NetworkRequestCancellationReason::TimedOut); + NetworkDecision::deny(NETWORK_POLICY_DENIAL_REASON) + }), + } + } + (None, _, _) | (_, None, _) | (_, _, None) => { + NetworkDecision::deny(NETWORK_POLICY_DENIAL_REASON) + } + }; + if let Some(expected_session) = expected_session { + let (Some(inner), Some(expected_session)) = + (inner.upgrade(), expected_session.upgrade()) + else { + return; + }; + if !inner + .get_session(&process_id) + .is_some_and(|session| Arc::ptr_eq(&session, &expected_session)) + { + return; + } + } + let Some(rpc_client) = rpc_client.upgrade() else { + return; + }; + let decision = match decision { + NetworkDecision::Allow => ExecServerNetworkPolicyDecision::Allow, + NetworkDecision::Deny { + reason, decision, .. + } if reason.len() <= MAX_NETWORK_POLICY_REASON_BYTES + && !reason.chars().any(char::is_control) => + { + match decision { + NetworkPolicyDecision::Deny => { + ExecServerNetworkPolicyDecision::Deny { reason } + } + NetworkPolicyDecision::Ask => { + ExecServerNetworkPolicyDecision::Ask { reason } + } + } + } + NetworkDecision::Deny { .. } => { + ExecServerNetworkPolicyDecision::Deny { + reason: NETWORK_POLICY_DENIAL_REASON.to_string(), + } + } + }; + if let Err(error) = rpc_client + .respond(request_id, &NetworkPolicyRequestResponse { decision }) + .await + { + debug!( + ?error, + "failed to send network policy decision to exec-server" + ); + } else { + request_outcome.complete("success"); + } + }; + tokio::spawn(task.instrument(task_span)); + } + RpcClientEvent::Notification(notification) => { + if let Err(error) = handle_server_notification(&inner, notification).await { + rpc_client.close_transport().await; + inner.request_recovery( + rpc_client, + format!("exec-server notification handling failed: {error}"), + ); + return; + } + } + RpcClientEvent::Disconnected { reason } => { + inner.request_recovery(rpc_client, disconnected_message(reason.as_deref())); + return; + } + } + } + }); + } +} + +pub(crate) fn is_retryable_recovery_error(error: &ExecServerError) -> bool { + if let ExecServerError::ConnectionAttempt(error) = error { + return is_retryable_recovery_error(error.as_ref()); + } + is_transport_closed_error(error) + || matches!( + error, + ExecServerError::ProvisioningFailed(_) + | ExecServerError::WebSocketConnectTimeout { .. } + | ExecServerError::WebSocketConnect { .. } + | ExecServerError::InitializeTimedOut { .. } + ) + || is_retryable_registry_error(error) + || matches!( + error, + ExecServerError::Server { code, .. } + if *code == SESSION_ALREADY_ATTACHED_ERROR_CODE + ) +} + +pub(crate) fn is_retryable_registry_error(error: &ExecServerError) -> bool { + matches!( + error, + ExecServerError::EnvironmentRegistryRequest(error) + if error.is_connect() + || error.is_timeout() + || error.is_body() + || matches!( + error, + codex_http_client::RouteAwareRequestError::Request(error) + if error.is_decode() + ) + ) || matches!( + error, + ExecServerError::EnvironmentRegistryHttp { status, .. } + if status.is_server_error() + || *status == http::StatusCode::REQUEST_TIMEOUT + || *status == http::StatusCode::TOO_MANY_REQUESTS + ) || is_environment_offline_error(error) +} + +pub(crate) fn is_environment_offline_error(error: &ExecServerError) -> bool { + matches!( + error, + ExecServerError::EnvironmentRegistryHttp { status, code, .. } + if *status == http::StatusCode::CONFLICT + && code.as_deref() == Some("environment_offline") + ) +} + +pub(crate) fn registry_recovery_retry_delay(retry_key: &str, attempt: u32) -> Duration { + let multiplier = 1_u32.checked_shl(attempt.min(4)).unwrap_or(u32::MAX); + let base_delay = REGISTRY_RECOVERY_INITIAL_RETRY_INTERVAL + .saturating_mul(multiplier) + .min(REGISTRY_RECOVERY_MAX_RETRY_INTERVAL); + let base_millis = base_delay.as_millis() as u64; + let mut hasher = DefaultHasher::new(); + retry_key.hash(&mut hasher); + attempt.hash(&mut hasher); + + Duration::from_millis(base_millis + hasher.finish() % (base_millis / 2 + 1)) +} + +#[cfg(test)] +#[path = "client_recovery_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/client_recovery_tests.rs b/codex-rs/exec-server/src/client_recovery_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4abc939d323955ae708971fb9bb21bc8a2248c05 --- /dev/null +++ b/codex-rs/exec-server/src/client_recovery_tests.rs @@ -0,0 +1,262 @@ +use std::time::Duration; + +use pretty_assertions::assert_eq; + +use super::*; +use crate::protocol::ExecOutputStream; +use crate::protocol::ProcessOutputChunk; + +fn registry_error(status: http::StatusCode, code: Option<&str>) -> ExecServerError { + ExecServerError::EnvironmentRegistryHttp { + status, + code: code.map(str::to_string), + message: "registry unavailable".to_string(), + } +} + +#[test] +fn registry_recovery_retry_delay_exponentially_backs_off_and_caps() { + let cases = [ + (0, Duration::from_millis(500)), + (1, Duration::from_secs(1)), + (2, Duration::from_secs(2)), + (3, Duration::from_secs(4)), + (4, Duration::from_secs(5)), + (20, Duration::from_secs(5)), + ]; + + for (attempt, base) in cases { + let delay = registry_recovery_retry_delay("session-1", attempt); + assert!(delay >= base, "delay {delay:?} for attempt {attempt}"); + assert!( + delay <= base + base / 2, + "delay {delay:?} for attempt {attempt}" + ); + } +} + +#[test] +fn recovery_retries_transient_registry_errors() { + for status in [ + http::StatusCode::REQUEST_TIMEOUT, + http::StatusCode::TOO_MANY_REQUESTS, + http::StatusCode::INTERNAL_SERVER_ERROR, + http::StatusCode::BAD_GATEWAY, + http::StatusCode::SERVICE_UNAVAILABLE, + ] { + let error = registry_error(status, /*code*/ None); + + assert!(is_retryable_registry_error(&error)); + assert!(is_retryable_recovery_error(&error)); + assert!(is_retryable_recovery_error( + &ExecServerError::ConnectionAttempt(Arc::new(error)) + )); + } +} + +#[test] +fn recovery_retries_registry_request_timeouts() { + let error = ExecServerError::EnvironmentRegistryRequest( + codex_http_client::RouteAwareRequestError::Timeout, + ); + + assert!(is_retryable_registry_error(&error)); + assert!(is_retryable_recovery_error(&error)); +} + +#[test] +fn recovery_retries_environment_offline_conflicts() { + let error = registry_error(http::StatusCode::CONFLICT, Some("environment_offline")); + + assert!(is_retryable_registry_error(&error)); + assert!(is_retryable_recovery_error(&error)); +} + +#[test] +fn recovery_does_not_retry_other_registry_conflicts() { + let error = registry_error(http::StatusCode::CONFLICT, Some("registration_conflict")); + + assert!(!is_retryable_registry_error(&error)); + assert!(!is_retryable_recovery_error(&error)); + assert!(!is_retryable_recovery_error( + &ExecServerError::ConnectionAttempt(Arc::new(error)) + )); +} + +#[test] +fn process_event_reorder_rejects_oversized_output() { + let state = SessionState::new(/*recoverable*/ true); + + let error = state + .publish_ordered_event(ExecProcessEvent::Output(ProcessOutputChunk { + seq: 1, + stream: ExecOutputStream::Stdout, + chunk: vec![0; super::super::MAX_PENDING_PROCESS_EVENT_BYTES + 1].into(), + })) + .expect_err("oversized pending process output should be rejected"); + + assert!(error.contains("bytes")); +} + +#[test] +fn process_event_reorder_accepts_gap_closing_event_at_limits() { + let state = SessionState::new(/*recoverable*/ true); + let chunk_size = + super::super::MAX_PENDING_PROCESS_EVENT_BYTES / super::super::MAX_PENDING_PROCESS_EVENTS; + let last_seq = super::super::MAX_PENDING_PROCESS_EVENTS as u64 + 1; + + for seq in 2..=last_seq { + assert!( + !state + .publish_ordered_event(ExecProcessEvent::Output(ProcessOutputChunk { + seq, + stream: ExecOutputStream::Stdout, + chunk: vec![0; chunk_size].into(), + })) + .expect("future output should fit within reorder limits") + ); + } + assert!( + !state + .publish_ordered_event(ExecProcessEvent::Output(ProcessOutputChunk { + seq: 1, + stream: ExecOutputStream::Stdout, + chunk: b"x".to_vec().into(), + })) + .expect("gap-closing output should drain the reorder buffer") + ); + + let ordered_events = state + .ordered_events + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + assert_eq!( + ( + ordered_events.last_published_seq, + ordered_events.pending.len(), + ordered_events.pending_bytes, + ), + (last_seq, 0, 0) + ); +} + +#[test] +fn recovery_handles_dense_tail_output_and_newer_notification() { + let state = SessionState::new(/*recoverable*/ true); + let last_seq = super::super::MAX_PENDING_PROCESS_EVENTS as u64 + 2; + let live_seq = last_seq + 1; + assert!( + !state + .publish_ordered_event(ExecProcessEvent::Output(ProcessOutputChunk { + seq: live_seq, + stream: ExecOutputStream::Stdout, + chunk: b"live".to_vec().into(), + })) + .expect("live output should remain bounded while recovery fills the gap") + ); + let chunks = (2..=last_seq) + .map(|seq| ProcessOutputChunk { + seq, + stream: ExecOutputStream::Stdout, + chunk: b"x".to_vec().into(), + }) + .collect(); + + assert!( + !state + .recover_events(ReadResponse { + chunks, + next_seq: last_seq + 1, + exited: true, + exit_code: Some(17), + closed: false, + failure: None, + sandbox_denied: false, + }) + .expect("dense retained output should recover") + ); + + let ordered_events = state + .ordered_events + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + assert_eq!( + ( + ordered_events.last_published_seq, + ordered_events.pending.len(), + ordered_events.pending_bytes, + ), + (live_seq, 0, 0) + ); +} + +#[test] +fn recovery_rejects_output_at_closed_sequence() { + let state = SessionState::new(/*recoverable*/ true); + + let error = state + .recover_events(ReadResponse { + chunks: vec![ProcessOutputChunk { + seq: 1, + stream: ExecOutputStream::Stdout, + chunk: b"output".to_vec().into(), + }], + next_seq: 2, + exited: false, + exit_code: None, + closed: true, + failure: None, + sandbox_denied: false, + }) + .expect_err("output should not occupy the closed sequence"); + + assert!( + error + .to_string() + .contains("conflicts with recovered output") + ); +} + +#[tokio::test] +async fn recovery_adds_sandbox_denial_to_pending_exit_event() { + let state = SessionState::new(/*recoverable*/ true); + assert!( + !state + .publish_ordered_event(ExecProcessEvent::Exited { + seq: 2, + exit_code: 1, + sandbox_denied: None, + }) + .expect("pending exit should fit within reorder limits") + ); + + state + .recover_events(ReadResponse { + chunks: vec![ProcessOutputChunk { + seq: 1, + stream: ExecOutputStream::Stderr, + chunk: b"sandbox denied".to_vec().into(), + }], + next_seq: 3, + exited: true, + exit_code: Some(1), + closed: false, + failure: None, + sandbox_denied: true, + }) + .expect("recovery should publish the pending exit"); + + let mut events = state.subscribe_events(); + assert!(matches!( + events.recv().await, + Ok(ExecProcessEvent::Output(_)) + )); + assert_eq!( + events.recv().await, + Ok(ExecProcessEvent::Exited { + seq: 2, + exit_code: 1, + sandbox_denied: Some(true), + }) + ); +} diff --git a/codex-rs/exec-server/src/client_refresh.rs b/codex-rs/exec-server/src/client_refresh.rs new file mode 100644 index 0000000000000000000000000000000000000000..9918227b7147b828b6ba5486c203ea29839bcd31 --- /dev/null +++ b/codex-rs/exec-server/src/client_refresh.rs @@ -0,0 +1,269 @@ +//! Explicit connection refresh after a planned executor replacement. +//! +//! Ordinary recovery tries to resume the same executor session after a transient +//! disconnect. Replacement needs a fresh session, without waiting for old recovery +//! to give up. The caller supplies the ordering: register the replacement first, +//! then refresh. Executor identity stays inside this connection layer. +//! +//! Flow: fresh registry lookup -> reuse or retire session -> connect if needed -> +//! live status probe. The lazy client and its `Environment` remain the same objects; +//! only the underlying `ExecServerClient` may change. The public caller contract is on +//! `Environment::refresh_connection`. +//! +//! Two races determine the synchronization here. A client installed during the +//! lookup makes that lookup stale, so refresh checks again. A connection attempt +//! cancelled by refresh must never install later. Cancellation and installation +//! synchronize on the `current_client` lock; acquire `reconnect` first when both are needed. +//! `refresh_lock` serializes only explicit refreshes; ordinary connection and recovery +//! work can continue concurrently. Retired sessions cannot publish environment state. + +use std::sync::Arc; +use std::sync::Mutex as StdMutex; + +use futures::future::BoxFuture; +use tokio::sync::OnceCell; +use tokio::sync::watch; +use tokio::time::timeout; +use tokio_util::sync::CancellationToken; + +use super::ConnectionResult; +use super::ConnectionStatus; +use super::ExecServerClient; +use super::ExecServerError; +use super::Inner; +use super::LazyRemoteExecServerClient; +use super::fail_all_in_flight_work; +use crate::EnvironmentConnectionState; +use crate::NoiseChannelPublicKey; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::client_api::ExecServerTransportParams; +use crate::client_api::NoiseRendezvousConnectBundle; +use crate::client_api::NoiseRendezvousConnectProvider; +use crate::client_transport::ExecServerReconnectStrategy; + +/// Shared startup/reconnect result plus cancellation for work superseded by refresh. +/// The optional transport carries a refresh lookup's bundle into the normal connector. +#[derive(Default)] +pub(super) struct ConnectionAttempt { + pub(super) result: OnceCell, + pub(super) cancelled: CancellationToken, + pub(super) transport: Option, +} + +// Use the compared bundle intact for the first connection: address, key and authorization +// belong together. Later lookups, including authorization refresh, use the real provider. +struct PrefetchedConnectProvider { + bundle: StdMutex>, + provider: Arc, +} + +impl NoiseRendezvousConnectProvider for PrefetchedConnectProvider { + fn connect_bundle( + &self, + harness_public_key: NoiseChannelPublicKey, + ) -> BoxFuture<'_, Result> { + Box::pin(async move { + let bundle = self + .bundle + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .take(); + match bundle { + Some(bundle) => Ok(bundle), + None => self.provider.connect_bundle(harness_public_key).await, + } + }) + } +} + +impl LazyRemoteExecServerClient { + #[expect( + clippy::await_holding_invalid_type, + reason = "serialize explicit refreshes, not ordinary connection or recovery attempts" + )] + pub(crate) async fn refresh_connection(&self) -> Result<(), ExecServerError> { + let _refresh = self.refresh_lock.lock().await; + let (previous, attempt) = loop { + let observed = self.cached_client(); + let mut transport = self.transport_params.clone().ok_or_else(|| { + ExecServerError::Protocol( + "connection refresh requires a Noise registry".to_string(), + ) + })?; + let target = match &mut transport { + ExecServerTransportParams::Deferred(deferred) => &mut deferred.transport, + transport => transport, + }; + let ExecServerTransportParams::NoiseRendezvous { provider, identity } = target else { + return Err(ExecServerError::Protocol( + "connection refresh requires a Noise registry".to_string(), + )); + }; + // This lookup is independent of the old session and its recovery deadline. + let bundle = timeout( + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + provider.connect_bundle(identity.public_key()), + ) + .await + .map_err(|_| { + ExecServerError::EnvironmentRegistryRequest( + codex_http_client::RouteAwareRequestError::Timeout, + ) + })??; + let executor_public_key = bundle.executor_public_key.clone(); + *provider = Arc::new(PrefetchedConnectProvider { + bundle: StdMutex::new(Some(bundle)), + provider: Arc::clone(provider), + }); + + let mut reconnect = self + .reconnect + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let current = self + .current_client + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + // An ordinary connection may have finished during the lookup. Re-read the + // registry rather than retire a newer client using a superseded response. + if !match (&observed, &*current) { + (Some(observed), Some(current)) => Arc::ptr_eq(&observed.inner, ¤t.inner), + (None, None) => true, + _ => false, + } { + continue; + } + // is_disconnected means terminally failed, not temporarily recovering. + // Preserve same-executor recovery; the final probe fails fast if still recovering. + if current.as_ref().is_some_and(|client| { + !client.is_disconnected() + && matches!( + client.inner.reconnect_strategy.as_ref(), + Some(ExecServerReconnectStrategy::NoiseRendezvous { + executor_public_key: key, .. + }) if key == &executor_public_key + ) + }) { + break (current.clone(), None); + } + // Cancellation and connection installation use the same lock. A late + // handshake cannot install a client after its attempt has been superseded. + self.startup.cancelled.cancel(); + if let Some(attempt) = reconnect.as_ref() { + attempt.cancelled.cancel(); + } + self.environment_connection_state_tx + .send_replace(EnvironmentConnectionState::Disconnected); + let attempt = Arc::new(ConnectionAttempt { + transport: Some(transport), + ..Default::default() + }); + *reconnect = Some(Arc::clone(&attempt)); + break (current.clone(), Some(attempt)); + }; + let client = match attempt { + Some(attempt) => { + if let Some(previous) = previous { + previous.inner.retire().await; + } + let result = attempt + .result + .get_or_init(|| self.connect_once(&attempt)) + .await + .clone(); + let mut reconnect = self + .reconnect + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if reconnect + .as_ref() + .is_some_and(|current| Arc::ptr_eq(current, &attempt)) + { + *reconnect = None; + } + result.map_err(ExecServerError::ConnectionAttempt)? + } + None => previous.ok_or_else(|| { + ExecServerError::Protocol("current executor session is missing".to_string()) + })?, + }; + // Metadata may be cached; readiness requires a live, non-recovering probe. + client.environment_status().await.map(drop) + } + + #[tracing::instrument(name = "codex.exec_server.remote.connect", skip_all)] + pub(super) fn connect_once<'a>( + &'a self, + attempt: &'a ConnectionAttempt, + ) -> BoxFuture<'a, ConnectionResult> { + // Keep the transport future out of every caller's async layout, including + // the CLI entry point, which otherwise exceeds rustc's query-depth limit. + Box::pin(async move { + let transport = attempt + .transport + .as_ref() + .or(self.transport_params.as_ref()) + .ok_or_else(|| { + Arc::new(ExecServerError::Protocol( + "missing transport params for lazy exec-server connection".to_string(), + )) + })?; + let client = tokio::select! { + biased; + _ = attempt.cancelled.cancelled() => return Err(Arc::new(ExecServerError::Disconnected("connection attempt was superseded".to_string()))), + result = ExecServerClient::connect_for_transport(transport.clone(), self.http_client_factory.clone()) => result.map_err(Arc::new)?, + }; + // Cancellation can race with a completed handshake. Recheck before attaching + // state or installing the client, under the same lock used by refresh. + { + let mut current = self + .current_client + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if !attempt.cancelled.is_cancelled() { + client.attach_environment_connection_state( + self.environment_connection_state_tx.clone(), + ); + *current = Some(client.clone()); + return Ok(client); + } + } + client.inner.retire().await; + Err(Arc::new(ExecServerError::Disconnected( + "connection attempt was superseded".to_string(), + ))) + }) + } +} + +impl Inner { + async fn retire(self: &Arc) { + let message = "exec-server executor was replaced".to_string(); + let rpc_client = { + let mut connection = self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + // Detach before a later transport completion can publish stale state. + connection.environment_connection_state_tx = + watch::channel(EnvironmentConnectionState::Disconnected).0; + let rpc_client = match &connection.status { + ConnectionStatus::Connected(client) => Some(Arc::clone(client)), + ConnectionStatus::Recovering | ConnectionStatus::Failed(_) => None, + }; + self.retired.cancel(); + connection.set_status(ConnectionStatus::Failed(message.clone())); + rpc_client + }; + self.connection_changed.send_replace(()); + // Drain pending RPCs before stream cleanup, which may wait for other work. + if let Some(rpc_client) = rpc_client { + rpc_client.close_transport().await; + } + fail_all_in_flight_work(self, message).await; + } +} + +#[cfg(test)] +#[path = "client_refresh_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/client_refresh_tests.rs b/codex-rs/exec-server/src/client_refresh_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4e45a85ec50d40c5082fe661b9c4105e10b15be6 --- /dev/null +++ b/codex-rs/exec-server/src/client_refresh_tests.rs @@ -0,0 +1,621 @@ +use std::sync::Arc; +use std::sync::Mutex; +use std::time::Duration; + +use anyhow::Result; +use futures::future::BoxFuture; +use pretty_assertions::assert_eq; +use tokio::net::TcpListener; +use tokio::sync::Notify; +use tokio::sync::oneshot; +use tokio::task::JoinSet; +use tokio_util::task::AbortOnDropHandle; + +use super::*; +use crate::ExecServerRuntimePaths; +use crate::NoiseChannelIdentity; +use crate::ProcessId; +use crate::relay::HarnessKeyValidator; +use crate::relay::run_multiplexed_environment; +use crate::server::ConnectionProcessor; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; + +// These tests exercise real sockets, not keepalive deadlines. The shared unit-test +// Pong timeout is only 100 ms and can expire during Noise handshakes under load. +// A blocking task prevents paused time from auto-advancing while socket I/O is +// pending; dropping the returned sender releases it without polling or sleeping. +fn freeze_clock() -> std::sync::mpsc::Sender<()> { + tokio::time::pause(); + let (guard, dropped) = std::sync::mpsc::channel(); + tokio::task::spawn_blocking(move || { + let _ = dropped.recv(); + }); + guard +} + +#[derive(Clone)] +struct Target { + url: String, + identity: NoiseChannelIdentity, + registration: String, +} + +struct Registry { + target: Mutex, + next_lookup: Mutex>>, + lookup_started: Notify, +} + +impl Registry { + fn block_next_lookup(&self) -> oneshot::Sender<()> { + let (tx, rx) = oneshot::channel(); + *self.next_lookup.lock().unwrap() = Some(rx); + tx + } +} + +impl NoiseRendezvousConnectProvider for Registry { + fn connect_bundle( + &self, + _: NoiseChannelPublicKey, + ) -> BoxFuture<'_, Result> { + Box::pin(async move { + let target = self.target.lock().unwrap().clone(); + let block = self.next_lookup.lock().unwrap().take(); + if let Some(block) = block { + self.lookup_started.notify_one(); + block + .await + .map_err(|_| ExecServerError::Protocol("test lookup failed".to_owned()))?; + } + Ok(NoiseRendezvousConnectBundle { + websocket_url: target.url, + environment_id: "environment".to_owned(), + executor_registration_id: target.registration, + executor_public_key: target.identity.public_key(), + harness_key_authorization: "authorization".to_owned(), + }) + }) + } +} + +#[derive(Clone, Default)] +struct Validator { + handshake: Option>, + started: Arc, +} + +impl HarnessKeyValidator for Validator { + async fn validate_harness_key( + &self, + _: &NoiseChannelPublicKey, + _: &str, + ) -> Result<(), ExecServerError> { + if let Some(handshake) = &self.handshake { + self.started.notify_one(); + handshake.notified().await; + } + Ok(()) + } +} + +struct Executor { + target: Target, + _server: AbortOnDropHandle<()>, +} + +impl Executor { + async fn start(validator: Validator) -> Result { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let target = Target { + url: format!("ws://{}", listener.local_addr()?), + identity: NoiseChannelIdentity::generate()?, + registration: uuid::Uuid::new_v4().to_string(), + }; + let executor = target.clone(); + let processor = ConnectionProcessor::new(ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?); + let server = tokio::spawn(async move { + let mut connections = JoinSet::new(); + loop { + tokio::select! { + socket = listener.accept() => { + let (socket, _) = socket.unwrap(); + let executor = executor.clone(); + let processor = processor.clone(); + let validator = validator.clone(); + connections.spawn(async move { + let socket = tokio_tungstenite::accept_async(socket).await.unwrap(); + run_multiplexed_environment(socket, processor, "environment".to_owned(), executor.registration, executor.identity, validator).await; + }); + } + _ = connections.join_next(), if !connections.is_empty() => {} + } + } + }); + Ok(Self { + target, + _server: AbortOnDropHandle::new(server), + }) + } + + fn client(&self) -> Result<(LazyRemoteExecServerClient, Arc)> { + let registry = Arc::new(Registry { + target: Mutex::new(self.target.clone()), + next_lookup: Mutex::new(None), + lookup_started: Notify::new(), + }); + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::NoiseRendezvous { + provider: registry.clone(), + identity: NoiseChannelIdentity::generate()?, + }, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ); + Ok((client, registry)) + } +} + +async fn disconnect(client: &ExecServerClient) { + let rpc_client = { + let connection = client.inner.connection.lock().unwrap(); + let ConnectionStatus::Connected(rpc_client) = &connection.status else { + panic!("expected connected session") + }; + Arc::clone(rpc_client) + }; + rpc_client.close_transport().await; +} + +#[tokio::test] +async fn refresh_cancels_old_recovery_and_connects_without_resuming_old_session() -> Result<()> { + let _clock = freeze_clock(); + let old = Executor::start(Validator::default()).await?; + let new = Executor::start(Validator::default()).await?; + let (client, registry) = old.client()?; + let original = client.get().await?; + let process = original + .register_session(&ProcessId::from("old-process")) + .await?; + let blocked_recovery = registry.block_next_lookup(); + disconnect(&original).await; + registry.lookup_started.notified().await; + *registry.target.lock().unwrap() = new.target.clone(); + let refresh = client.refresh_connection(); + tokio::pin!(refresh); + let started = tokio::time::Instant::now(); + assert!(futures::poll!(refresh.as_mut()).is_pending()); + assert!(original.inner.retired.is_cancelled()); + assert!(original.is_disconnected()); + assert_eq!(started.elapsed(), Duration::ZERO); + refresh.await?; + let replacement = client.get().await?; + assert_ne!(original.session_id(), replacement.session_id()); + assert_eq!( + *client.environment_connection_state_tx.borrow(), + EnvironmentConnectionState::Connected + ); + assert!(matches!( + process.write(b"never replay".to_vec()).await, + Err(ExecServerError::Disconnected(_)) + )); + // Releasing an old registry response cannot reinstall or disconnect the replacement. + let _ = blocked_recovery.send(()); + tokio::task::yield_now().await; + assert!(Arc::ptr_eq(&client.get().await?.inner, &replacement.inner)); + replacement.environment_status().await?; + Ok(()) +} + +#[tokio::test] +async fn refresh_retires_a_still_connected_old_executor() -> Result<()> { + let _clock = freeze_clock(); + let old = Executor::start(Validator::default()).await?; + let new = Executor::start(Validator::default()).await?; + let (client, registry) = old.client()?; + let original = client.get().await?; + *registry.target.lock().unwrap() = new.target.clone(); + client.refresh_connection().await?; + assert!(original.is_disconnected()); + assert_ne!(original.session_id(), client.get().await?.session_id()); + Ok(()) +} + +#[tokio::test] +async fn refresh_preserves_a_current_session_across_registration_renewal() -> Result<()> { + let _clock = freeze_clock(); + let executor = Executor::start(Validator::default()).await?; + let (client, registry) = executor.client()?; + let original = client.get().await?; + registry.target.lock().unwrap().registration = "renewed-registration".to_owned(); + let concurrent = client.clone(); + let concurrent = tokio::spawn(async move { concurrent.refresh_connection().await }); + client.refresh_connection().await?; + concurrent.await??; + assert!(Arc::ptr_eq(&original.inner, &client.get().await?.inner)); + original.environment_status().await?; + Ok(()) +} + +#[tokio::test] +async fn refresh_cancels_a_stalled_initial_lookup() -> Result<()> { + let _clock = freeze_clock(); + let old = Executor::start(Validator::default()).await?; + let new = Executor::start(Validator::default()).await?; + let (client, registry) = old.client()?; + let release = registry.block_next_lookup(); + let initial = client.get(); + tokio::pin!(initial); + assert!(futures::poll!(initial.as_mut()).is_pending()); + *registry.target.lock().unwrap() = new.target.clone(); + client.refresh_connection().await?; + assert!(initial.await.is_err()); + let _ = release.send(()); + client.get().await?.environment_status().await?; + Ok(()) +} + +#[tokio::test] +async fn refresh_cancels_a_stalled_noise_handshake() -> Result<()> { + let _clock = freeze_clock(); + let validator = Validator { + handshake: Some(Arc::new(Notify::new())), + ..Default::default() + }; + let old = Executor::start(validator.clone()).await?; + let new = Executor::start(Validator::default()).await?; + let (client, registry) = old.client()?; + let connecting = client.clone(); + let initial = tokio::spawn(async move { connecting.get().await }); + validator.started.notified().await; + *registry.target.lock().unwrap() = new.target.clone(); + client.refresh_connection().await?; + assert!(initial.await?.is_err()); + validator.handshake.unwrap().notify_one(); + client.get().await?.environment_status().await?; + assert_eq!( + *client.environment_connection_state_tx.borrow(), + EnvironmentConnectionState::Connected + ); + Ok(()) +} + +#[tokio::test] +async fn failed_refresh_lookup_leaves_the_existing_session_usable() -> Result<()> { + let _clock = freeze_clock(); + let executor = Executor::start(Validator::default()).await?; + let (client, registry) = executor.client()?; + let original = client.get().await?; + drop(registry.block_next_lookup()); + assert!(client.refresh_connection().await.is_err()); + assert!(Arc::ptr_eq(&original.inner, &client.get().await?.inner)); + original.environment_status().await?; + Ok(()) +} + +#[tokio::test] +async fn failed_replacement_connection_keeps_old_handles_retired_and_get_retries() -> Result<()> { + let _clock = freeze_clock(); + let old = Executor::start(Validator::default()).await?; + let new = Executor::start(Validator::default()).await?; + let (client, registry) = old.client()?; + let original = client.get().await?; + let process = original + .register_session(&ProcessId::from("old-process")) + .await?; + + // The registry knows the replacement, but its endpoint drops the new connection. + let unavailable = TcpListener::bind("127.0.0.1:0").await?; + let target = Target { + url: format!("ws://{}", unavailable.local_addr()?), + ..new.target.clone() + }; + let _rejected_connection = AbortOnDropHandle::new(tokio::spawn(async move { + drop(unavailable.accept().await.unwrap()); + })); + *registry.target.lock().unwrap() = target; + assert!(matches!( + client.refresh_connection().await, + Err(ExecServerError::ConnectionAttempt(_)) + )); + assert!(original.inner.retired.is_cancelled()); + assert!(matches!( + original.environment_status().await, + Err(ExecServerError::Disconnected(_)) + )); + assert!(matches!( + process.write(b"never replay".to_vec()).await, + Err(ExecServerError::Disconnected(_)) + )); + assert_eq!( + *client.environment_connection_state_tx.borrow(), + EnvironmentConnectionState::Disconnected + ); + + // A later caller retries against the registry instead of reviving the retired client. + *registry.target.lock().unwrap() = new.target.clone(); + let replacement = client.get().await?; + assert_ne!(original.session_id(), replacement.session_id()); + replacement.environment_status().await?; + assert_eq!( + *client.environment_connection_state_tx.borrow(), + EnvironmentConnectionState::Connected + ); + assert!(matches!( + process.write(b"still retired".to_vec()).await, + Err(ExecServerError::Disconnected(_)) + )); + Ok(()) +} + +#[tokio::test] +async fn superseded_refresh_lookup_does_not_retire_a_newer_session() -> Result<()> { + let _clock = freeze_clock(); + let old = Executor::start(Validator::default()).await?; + let new = Executor::start(Validator::default()).await?; + let (client, registry) = old.client()?; + let original = client.get().await?; + let release = registry.block_next_lookup(); + let refreshing = client.refresh_connection(); + tokio::pin!(refreshing); + assert!(futures::poll!(refreshing.as_mut()).is_pending()); + *registry.target.lock().unwrap() = new.target.clone(); + original.inner.retire().await; + let replacement = client.get().await?; + release.send(()).unwrap(); + refreshing.await?; + assert!(Arc::ptr_eq(&replacement.inner, &client.get().await?.inner)); + assert!(!replacement.inner.retired.is_cancelled()); + replacement.environment_status().await?; + Ok(()) +} + +#[tokio::test] +async fn environment_refresh_preserves_environment_and_filesystem_handles() -> Result<()> { + let _clock = freeze_clock(); + let old = Executor::start(Validator::default()).await?; + let new = Executor::start(Validator::default()).await?; + let (_, registry) = old.client()?; + let manager = crate::EnvironmentManager::from_snapshot( + crate::environment_provider::EnvironmentProviderSnapshot { + environments: Vec::new(), + default: crate::environment_provider::EnvironmentDefault::Disabled, + include_local: false, + }, + /*local_runtime_paths*/ None, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )?; + let environment = manager + .materialize_pending_noise_environment("environment".to_owned(), registry.clone())?; + manager.report_environment_provisioning_status( + "environment".to_owned(), + Ok(crate::EnvironmentReadyInfo { + selected_capability_roots: Vec::new(), + }), + registry.clone(), + )?; + environment.info().await?; + let filesystem = environment.get_filesystem(); + *registry.target.lock().unwrap() = new.target.clone(); + environment.refresh_connection().await?; + assert!(Arc::ptr_eq( + &environment, + &manager.get_environment("environment").unwrap() + )); + assert!(Arc::ptr_eq(&filesystem, &environment.get_filesystem())); + environment.info().await?; + Ok(()) +} + +#[tokio::test] +async fn refresh_does_not_accept_cached_metadata_during_recovery() -> Result<()> { + let _clock = freeze_clock(); + let executor = Executor::start(Validator::default()).await?; + let (client, registry) = executor.client()?; + let original = client.get().await?; + original.environment_info().await?; + let blocked_recovery = registry.block_next_lookup(); + disconnect(&original).await; + registry.lookup_started.notified().await; + let started = tokio::time::Instant::now(); + assert!(matches!( + client.refresh_connection().await, + Err(ExecServerError::Disconnected(_)) + )); + assert_eq!(started.elapsed(), Duration::ZERO); + assert!(!original.inner.retired.is_cancelled()); + drop(blocked_recovery); + Ok(()) +} + +#[tokio::test] +async fn refresh_before_startup_marks_startup_finished() -> Result<()> { + let _clock = freeze_clock(); + let executor = Executor::start(Validator::default()).await?; + let (client, _) = executor.client()?; + assert!(!client.startup_finished()); + client.refresh_connection().await?; + assert!(client.startup_finished()); + assert!(matches!(client.readiness_result(), Some(Ok(())))); + Ok(()) +} + +struct ControlledRpc { + client: ExecServerClient, + requests: tokio::sync::mpsc::Receiver, + responses: tokio::sync::mpsc::Sender, +} + +async fn controlled_rpc() -> Result { + use crate::connection::JsonRpcConnection; + use crate::connection::JsonRpcConnectionEvent; + use crate::connection::JsonRpcTransport; + use codex_exec_server_protocol::JSONRPCMessage; + use codex_exec_server_protocol::JSONRPCResponse; + let (outgoing_tx, mut requests) = tokio::sync::mpsc::channel(/*buffer*/ 8); + let (responses, incoming_rx) = tokio::sync::mpsc::channel(/*buffer*/ 8); + let connection = JsonRpcConnection { + outgoing_tx, + incoming_rx, + disconnected_rx: tokio::sync::watch::channel(/*init*/ false).1, + task_handles: Vec::new(), + transport: JsonRpcTransport::Plain, + }; + let connecting = ExecServerClient::connect(connection, /*options*/ Default::default()); + tokio::pin!(connecting); + assert!(futures::poll!(connecting.as_mut()).is_pending()); + let Some(JSONRPCMessage::Request(initialize)) = requests.recv().await else { + anyhow::bail!("expected initialize request"); + }; + responses + .send(JsonRpcConnectionEvent::message(JSONRPCMessage::Response( + JSONRPCResponse { + id: initialize.id, + result: serde_json::json!({"sessionId": "controlled-session"}), + }, + ))) + .await?; + let client = connecting.await?; + assert!(matches!( + requests.recv().await, + Some(JSONRPCMessage::Notification(_)) + )); + Ok(ControlledRpc { + client, + requests, + responses, + }) +} + +#[tokio::test] +#[expect( + clippy::await_holding_invalid_type, + reason = "hold stream cleanup pending to exercise retirement ordering" +)] +async fn retirement_rejects_pending_mutation_before_stream_cleanup() -> Result<()> { + use crate::connection::JsonRpcConnectionEvent; + use codex_exec_server_protocol::JSONRPCMessage; + use codex_exec_server_protocol::JSONRPCResponse; + for response_queued in [false, true] { + let mut rpc = controlled_rpc().await?; + let call = rpc.client.fs_remove(crate::protocol::FsRemoveParams { + path: "file:///retired-file".parse()?, + recursive: None, + force: None, + follow_symlinks: None, + sandbox: None, + }); + tokio::pin!(call); + assert!(futures::poll!(call.as_mut()).is_pending()); + let Some(JSONRPCMessage::Request(request)) = rpc.requests.recv().await else { + anyhow::bail!("expected filesystem request"); + }; + let response = JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::json!({}), + }); + if response_queued { + rpc.responses + .send(JsonRpcConnectionEvent::message(response.clone())) + .await?; + let transport = rpc.client.rpc_client_without_recovery()?; + tokio::time::timeout(Duration::from_secs(5), async { + while transport.pending_request_count().await != 0 { + tokio::task::yield_now().await; + } + }) + .await?; + } + // Hold stream cleanup. Cover both a late response and one already queued + // for the caller when retirement begins; closing the socket alone misses the latter. + let streams = rpc.client.inner.http_body_streams_write_lock.lock().await; + let retirement = rpc.client.inner.retire(); + tokio::pin!(retirement); + assert!(futures::poll!(retirement.as_mut()).is_pending()); + if !response_queued { + rpc.responses + .send(JsonRpcConnectionEvent::message(response)) + .await?; + } + assert!(matches!(call.await, Err(ExecServerError::Disconnected(_)))); + drop(streams); + retirement.await; + } + Ok(()) +} + +#[tokio::test] +#[expect( + clippy::await_holding_invalid_type, + reason = "hold stream cleanup pending to exercise retirement ordering" +)] +async fn retirement_rejects_pending_process_start_before_stream_cleanup() -> Result<()> { + use crate::connection::JsonRpcConnectionEvent; + use codex_exec_server_protocol::JSONRPCMessage; + use codex_exec_server_protocol::JSONRPCResponse; + for response_queued in [false, true] { + let mut rpc = controlled_rpc().await?; + let process_id = ProcessId::from("retired-process"); + let call = rpc.client.start_process( + crate::protocol::ExecParams { + metadata: Default::default(), + process_id: process_id.clone(), + argv: vec!["unused".to_owned()], + cwd: "file:///".parse()?, + shell_snapshot: None, + env_policy: None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }, + /*network_policy_decider*/ None, + ); + tokio::pin!(call); + assert!(futures::poll!(call.as_mut()).is_pending()); + let Some(JSONRPCMessage::Request(request)) = rpc.requests.recv().await else { + anyhow::bail!("expected process start request"); + }; + let response = JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::json!({"processId": "retired-process"}), + }); + if response_queued { + rpc.responses + .send(JsonRpcConnectionEvent::message(response.clone())) + .await?; + let state = rpc + .client + .inner + .get_session(&process_id) + .expect("pending process"); + // The start task has a second response channel to the original caller. + tokio::time::timeout(Duration::from_secs(5), async { + while !state.recoverable.load(std::sync::atomic::Ordering::Acquire) { + tokio::task::yield_now().await; + } + }) + .await?; + } + let streams = rpc.client.inner.http_body_streams_write_lock.lock().await; + let retirement = rpc.client.inner.retire(); + tokio::pin!(retirement); + assert!(futures::poll!(retirement.as_mut()).is_pending()); + if !response_queued { + rpc.responses + .send(JsonRpcConnectionEvent::message(response)) + .await?; + } + assert!(matches!(call.await, Err(ExecServerError::Disconnected(_)))); + drop(streams); + retirement.await; + } + Ok(()) +} diff --git a/codex-rs/exec-server/src/client_telemetry.rs b/codex-rs/exec-server/src/client_telemetry.rs new file mode 100644 index 0000000000000000000000000000000000000000..f35b17af3808805e271cc2222a625422653c5b76 --- /dev/null +++ b/codex-rs/exec-server/src/client_telemetry.rs @@ -0,0 +1,23 @@ +//! Caller-side RPC attempt counts using the client's protocol method names, independent of tracing. + +use codex_otel::EXEC_SERVER_CLIENT_REQUEST_COUNT_METRIC; +use codex_otel::MetricsClient; + +pub(crate) fn record_client_request(metrics: Option<&MetricsClient>, method: &str) { + let Some(metrics) = metrics else { + return; + }; + // Record before local admission so failures and cancelled calls still count + // as attempts. Notifications and responses never enter these call paths. + if metrics + .counter_with_description( + EXEC_SERVER_CLIENT_REQUEST_COUNT_METRIC, + "Total number of client-side exec-server RPC attempts, including local failures.", + /*inc*/ 1, + &[("method", method)], + ) + .is_err() + { + tracing::warn!("failed to emit exec-server client request counter"); + } +} diff --git a/codex-rs/exec-server/src/client_transport.rs b/codex-rs/exec-server/src/client_transport.rs new file mode 100644 index 0000000000000000000000000000000000000000..bbdf769038fa3622e9b0811431d498d9eafbf9ee --- /dev/null +++ b/codex-rs/exec-server/src/client_transport.rs @@ -0,0 +1,814 @@ +use std::process::Stdio; +use std::sync::Arc; +use std::time::Duration; + +use tokio::io::AsyncBufReadExt; +use tokio::io::BufReader; +use tokio::process::Command; +use tokio::sync::OwnedSemaphorePermit; +use tokio::time::Instant; +use tokio::time::sleep; +use tokio::time::timeout; +use tokio::time::timeout_at; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tracing::debug; +use tracing::warn; + +use codex_api::AuthError; +use codex_api::AuthProvider; +use codex_http_client::HttpClientFactory; +use codex_http_client::Request; +use codex_http_client::RequestCompression; +use codex_protocol::shell_environment::scrub_non_inheritable_env_vars; +use codex_utils_rustls_provider::ensure_rustls_crypto_provider; +use codex_websocket_client::WebSocketConnection; +use codex_websocket_client::WebSocketConnector; +use codex_websocket_client::WebSocketTlsMode; +use http::HeaderMap; + +use crate::ExecServerClient; +use crate::ExecServerError; +use crate::client::NoiseInitializeContext; +use crate::client::accepted::AcceptedConnectionSource; +use crate::client::is_retryable_registry_error; +use crate::client::registry_recovery_retry_delay; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT; +use crate::client_api::ExecServerClientConnectOptions; +use crate::client_api::ExecServerTransportParams; +use crate::client_api::NoiseRendezvousConnectArgs; +use crate::client_api::NoiseRendezvousConnectBundle; +use crate::client_api::NoiseRendezvousConnectProvider; +use crate::client_api::RemoteExecServerConnectArgs; +use crate::client_api::StdioExecServerCommand; +use crate::client_api::StdioExecServerConnectArgs; +use crate::connection::JsonRpcConnection; +use crate::noise_channel::NoiseChannelIdentity; +use crate::noise_relay::NoiseHarnessConnectionArgs; +use crate::noise_relay::noise_harness_connection_from_websocket_with_readiness; +use crate::noise_relay::noise_relay_websocket_config; +use crate::relay::harness_connection_from_websocket; +use crate::trace_context::current_rendezvous_headers; + +const ENVIRONMENT_CLIENT_NAME: &str = "codex-environment"; +const INITIAL_REGISTRY_MAX_RETRIES: u32 = 4; +const INITIAL_REGISTRY_REQUEST_TIMEOUT: Duration = Duration::from_secs(6); +const INITIAL_REGISTRY_OPERATION_TIMEOUT: Duration = Duration::from_secs(14); + +pub(crate) async fn connect_websocket_request( + request: http::Request<()>, + diagnostic_url: String, + connector: WebSocketConnector, + connect_timeout: Duration, + use_loopback_direct: bool, +) -> Result { + let websocket_config = tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default(); + timeout(connect_timeout, async { + if use_loopback_direct { + connector + .connect_loopback_direct(request, websocket_config) + .await + } else { + connector.connect(request, websocket_config).await + } + }) + .await + .map_err(|_| ExecServerError::WebSocketConnectTimeout { + url: diagnostic_url.clone(), + timeout: connect_timeout, + })? + .map(|(websocket, _)| websocket) + .map_err(|source| ExecServerError::WebSocketConnect { + url: diagnostic_url, + source, + }) +} + +pub(crate) async fn authenticate_websocket_request( + request: &mut http::Request<()>, + auth_provider: &dyn AuthProvider, +) -> Result<(), AuthError> { + let url = request.uri().to_string(); + let signing_url = if let Some(rest) = url.strip_prefix("wss://") { + format!("https://{rest}") + } else if let Some(rest) = url.strip_prefix("ws://") { + format!("http://{rest}") + } else { + url + }; + let mut auth_request = Request::new(request.method().clone(), signing_url); + // Intermediaries may rewrite WebSocket and hop-by-hop headers after signing. + if let Some(host) = request.headers().get(http::header::HOST) { + auth_request + .headers + .insert(http::header::HOST, host.clone()); + } + let authenticated = auth_provider.apply_auth(auth_request).await?; + if authenticated.method != *request.method() { + return Err(AuthError::Build( + "authentication changed the WebSocket request method".to_string(), + )); + } + if authenticated.body.is_some() || authenticated.compression != RequestCompression::None { + return Err(AuthError::Build( + "authentication added a body or compression to the WebSocket request".to_string(), + )); + } + + let authenticated_websocket_url = websocket_url_from_authenticated_url(&authenticated.url)?; + let authenticated_uri = authenticated_websocket_url.parse().map_err(|error| { + AuthError::Build(format!("invalid authenticated WebSocket URL: {error}")) + })?; + let original_host = request.headers().get(http::header::HOST).cloned(); + for (name, value) in &authenticated.headers { + if is_websocket_handshake_header(name) { + if name == http::header::HOST && original_host.as_ref() == Some(value) { + continue; + } + return Err(AuthError::Build(format!( + "authentication changed WebSocket handshake header {name}" + ))); + } + request.headers_mut().insert(name, value.clone()); + } + *request.uri_mut() = authenticated_uri; + Ok(()) +} + +fn websocket_url_from_authenticated_url(url: &str) -> Result { + let mut url = url::Url::parse(url) + .map_err(|error| AuthError::Build(format!("invalid authenticated request URL: {error}")))?; + let websocket_scheme = match url.scheme() { + "https" => "wss", + "http" => "ws", + scheme => { + return Err(AuthError::Build(format!( + "authentication returned unsupported WebSocket URL scheme: {scheme}" + ))); + } + }; + url.set_scheme(websocket_scheme).map_err(|_| { + AuthError::Build("failed to convert authenticated URL to WebSocket scheme".to_string()) + })?; + Ok(url.into()) +} + +fn is_websocket_handshake_header(name: &http::header::HeaderName) -> bool { + name == http::header::HOST + || name == http::header::CONNECTION + || name == http::header::UPGRADE + || name == http::header::CONTENT_LENGTH + || name == http::header::TRANSFER_ENCODING + || name.as_str().starts_with("sec-websocket-") +} + +/// Everything the recovery loop needs for one connection attempt. +/// +/// An attempt may also carry a permit whose lifetime must extend until the +/// attempt finishes. +pub(crate) struct ReconnectAttempt { + connection: JsonRpcConnection, + options: ExecServerClientConnectOptions, + attempt_permit: Option, + noise_context: Option, +} + +struct OpenNoiseRendezvousConnection { + connection: JsonRpcConnection, + options: ExecServerClientConnectOptions, + handshake_ready: tokio::sync::oneshot::Receiver<()>, +} + +struct ReadyNoiseRendezvousConnection { + connection: JsonRpcConnection, + options: ExecServerClientConnectOptions, + noise_context: NoiseInitializeContext, +} + +impl ReconnectAttempt { + pub(crate) fn new( + connection: JsonRpcConnection, + options: ExecServerClientConnectOptions, + ) -> Self { + Self { + connection, + options, + attempt_permit: None, + noise_context: None, + } + } + + fn with_noise_context( + connection: JsonRpcConnection, + options: ExecServerClientConnectOptions, + noise_context: NoiseInitializeContext, + ) -> Self { + Self { + connection, + options, + attempt_permit: None, + noise_context: Some(noise_context), + } + } + + pub(crate) fn with_attempt_permit( + connection: JsonRpcConnection, + options: ExecServerClientConnectOptions, + attempt_permit: OwnedSemaphorePermit, + ) -> Self { + Self { + connection, + options, + attempt_permit: Some(attempt_permit), + noise_context: None, + } + } + + pub(crate) fn into_parts( + self, + ) -> ( + JsonRpcConnection, + ExecServerClientConnectOptions, + Option, + Option, + ) { + ( + self.connection, + self.options, + self.attempt_permit, + self.noise_context, + ) + } +} + +/// Reopens the transport for one logical exec-server client session. +/// +/// URL connections reuse their configured endpoint. Noise connections retain +/// the harness identity but fetch a fresh single-use authorization bundle for +/// every physical connection attempt. +#[derive(Clone)] +pub(crate) enum ExecServerReconnectStrategy { + Accepted(AcceptedConnectionSource), + WebSocket { + args: RemoteExecServerConnectArgs, + http_headers: HeaderMap, + }, + NoiseRendezvous { + // The executor that created the session, not the latest recovery lookup. + executor_public_key: crate::NoiseChannelPublicKey, + provider: Arc, + identity: NoiseChannelIdentity, + client_name: String, + connect_timeout: Duration, + initialize_timeout: Duration, + http_client_factory: HttpClientFactory, + }, +} + +impl ExecServerReconnectStrategy { + pub(crate) async fn resume( + &self, + session_id: &str, + ) -> Result { + match self { + Self::Accepted(source) => source.next_connection(session_id).await, + Self::WebSocket { args, http_headers } => { + let mut args = args.clone(); + args.resume_session_id = Some(session_id.to_string()); + let connection = + ExecServerClient::open_websocket_connection(&args, http_headers).await?; + Ok(ReconnectAttempt::new(connection, args.into())) + } + Self::NoiseRendezvous { + executor_public_key: _, + provider, + identity, + client_name, + connect_timeout, + initialize_timeout, + http_client_factory, + } => { + let bundle = provider.connect_bundle(identity.public_key()).await?; + let opened = ExecServerClient::open_noise_rendezvous_connection( + NoiseRendezvousConnectArgs { + bundle, + harness_identity: identity.clone(), + client_name: client_name.clone(), + connect_timeout: *connect_timeout, + initialize_timeout: *initialize_timeout, + resume_session_id: Some(session_id.to_string()), + http_client_factory: http_client_factory.clone(), + }, + ) + .await?; + let ready = ExecServerClient::finish_noise_rendezvous_connection(opened).await?; + Ok(ReconnectAttempt::with_noise_context( + ready.connection, + ready.options, + ready.noise_context, + )) + } + } + } +} + +impl ExecServerClient { + /// Open the selected transport and run the common JSON-RPC initialization. + /// Noise connection details are fetched here so reconnects get a fresh URL + /// and authorization without replacing the harness identity. + pub(crate) async fn connect_for_transport( + transport_params: ExecServerTransportParams, + http_client_factory: HttpClientFactory, + ) -> Result { + let (transport_params, deferred_readiness) = match transport_params { + ExecServerTransportParams::Deferred(deferred) => { + (deferred.transport, Some(deferred.readiness)) + } + transport_params => (transport_params, None), + }; + + if let Some(mut readiness) = deferred_readiness { + let provisioning_result = readiness + .wait_for(Option::is_some) + .await + .map_err(|_| { + ExecServerError::Disconnected( + "environment unavailable: environment provisioning ended before completion" + .to_string(), + ) + })? + .clone() + .ok_or_else(|| { + ExecServerError::Disconnected( + "environment unavailable: provisioning remained pending after completion" + .to_string(), + ) + })?; + provisioning_result.map_err(ExecServerError::ProvisioningFailed)?; + } + + let websocket = match transport_params { + ExecServerTransportParams::Deferred(_) => { + return Err(ExecServerError::Protocol( + "nested deferred exec-server transports are unsupported".to_string(), + )); + } + ExecServerTransportParams::WebSocketUrl { + websocket_url, + connect_timeout, + initialize_timeout, + http_headers, + } => ( + websocket_url, + connect_timeout, + initialize_timeout, + http_headers, + ), + ExecServerTransportParams::NoiseRendezvous { provider, identity } => { + let (ready, executor_public_key) = Self::open_initial_noise_rendezvous_connection( + &provider, + &identity, + http_client_factory.clone(), + ) + .await?; + let reconnect_strategy = ExecServerReconnectStrategy::NoiseRendezvous { + executor_public_key, + provider, + identity, + client_name: ENVIRONMENT_CLIENT_NAME.to_string(), + connect_timeout: DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + initialize_timeout: DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT, + http_client_factory, + }; + return Self::connect_with_recovery_and_noise_context( + ready.connection, + ready.options, + Some(reconnect_strategy), + ready.noise_context, + ) + .await; + } + ExecServerTransportParams::StdioCommand { + command, + initialize_timeout, + } => { + return Self::connect_stdio_command(StdioExecServerConnectArgs { + command, + client_name: ENVIRONMENT_CLIENT_NAME.to_string(), + initialize_timeout, + resume_session_id: None, + }) + .await; + } + }; + let (websocket_url, connect_timeout, initialize_timeout, http_headers) = websocket; + Self::connect_websocket_with_headers( + RemoteExecServerConnectArgs { + websocket_url, + client_name: ENVIRONMENT_CLIENT_NAME.to_string(), + connect_timeout, + initialize_timeout, + resume_session_id: None, + http_client_factory, + }, + http_headers, + ) + .await + } + + #[tracing::instrument(name = "codex.exec_server.remote.noise.connect", skip_all)] + async fn open_initial_noise_rendezvous_connection( + provider: &Arc, + identity: &NoiseChannelIdentity, + http_client_factory: HttpClientFactory, + ) -> Result<(ReadyNoiseRendezvousConnection, crate::NoiseChannelPublicKey), ExecServerError> + { + let open_connection = |bundle: NoiseRendezvousConnectBundle| { + Self::open_noise_rendezvous_connection(NoiseRendezvousConnectArgs { + bundle, + harness_identity: identity.clone(), + client_name: ENVIRONMENT_CLIENT_NAME.to_string(), + connect_timeout: DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + initialize_timeout: DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT, + resume_session_id: None, + http_client_factory: http_client_factory.clone(), + }) + }; + let mut deadline = Instant::now() + INITIAL_REGISTRY_OPERATION_TIMEOUT; + let retry_key = uuid::Uuid::new_v4().to_string(); + let mut retries = 0; + let mut refreshed_unauthorized_bundle = false; + let connect_bundle = || async { + timeout( + INITIAL_REGISTRY_REQUEST_TIMEOUT, + provider.connect_bundle(identity.public_key()), + ) + .await + .unwrap_or_else(|_| { + Err(ExecServerError::EnvironmentRegistryRequest( + codex_http_client::RouteAwareRequestError::Timeout, + )) + }) + }; + let mut result = connect_bundle().await; + loop { + let bundle = match result { + Ok(bundle) => bundle, + Err(error) + if is_retryable_registry_error(&error) + && retries < INITIAL_REGISTRY_MAX_RETRIES => + { + // Session resumption owns its separate recovery deadline. + let delay = registry_recovery_retry_delay(&retry_key, retries); + retries += 1; + result = match timeout_at(deadline, async { + sleep(delay).await; + connect_bundle().await + }) + .await + { + Ok(result) => result, + Err(_) => return Err(error), + }; + continue; + } + Err(error) => return Err(error), + }; + let executor_public_key = bundle.executor_public_key.clone(); + match open_connection(bundle).await { + Err(error) + if !refreshed_unauthorized_bundle + && matches!( + &error, + ExecServerError::WebSocketConnect { source, .. } + if matches!( + source, + tokio_tungstenite::tungstenite::Error::Http(response) + if response.status().as_u16() == 401 + ) + ) => + { + refreshed_unauthorized_bundle = true; + deadline = Instant::now() + INITIAL_REGISTRY_OPERATION_TIMEOUT; + retries = 0; + result = connect_bundle().await; + } + result => { + let opened = result?; + let ready = Self::finish_noise_rendezvous_connection(opened).await?; + return Ok((ready, executor_public_key)); + } + } + } + } + + pub async fn connect_websocket( + args: RemoteExecServerConnectArgs, + ) -> Result { + Self::connect_websocket_with_headers(args, HeaderMap::new()).await + } + + async fn connect_websocket_with_headers( + args: RemoteExecServerConnectArgs, + http_headers: HeaderMap, + ) -> Result { + let connection = Self::open_websocket_connection(&args, &http_headers).await?; + let options = args.clone().into(); + Self::connect_with_recovery( + connection, + options, + Some(ExecServerReconnectStrategy::WebSocket { args, http_headers }), + ) + .await + } + + pub(crate) async fn open_websocket_connection( + args: &RemoteExecServerConnectArgs, + http_headers: &HeaderMap, + ) -> Result { + ensure_rustls_crypto_provider(); + let websocket_url = args.websocket_url.clone(); + let connect_timeout = args.connect_timeout; + let mut request = websocket_url + .as_str() + .into_client_request() + .map_err(|source| ExecServerError::WebSocketConnect { + url: websocket_url.clone(), + source, + })?; + request.headers_mut().extend(http_headers.clone()); + let connector = WebSocketConnector::new_with_tls_mode( + &args.http_client_factory, + WebSocketTlsMode::TungsteniteDefault, + ) + .map_err(|error| ExecServerError::WebSocketConfiguration(error.to_string()))?; + let stream = connect_websocket_request( + request, + websocket_url.clone(), + connector, + connect_timeout, + !http_headers.is_empty() && websocket_url.starts_with("ws://"), + ) + .await?; + + let connection_label = format!("exec-server websocket {websocket_url}"); + let connection = if is_rendezvous_harness_url(&websocket_url) { + harness_connection_from_websocket(stream, connection_label) + } else { + JsonRpcConnection::from_websocket(stream, connection_label) + }; + Ok(connection) + } + + /// Connect to one exec-server through an authenticated rendezvous stream + /// using a caller-supplied single-use authorization bundle. + /// + /// The executor key is pinned before JSON-RPC starts; the websocket carries + /// only ciphertext after that. Environment-managed connections use a + /// retained [`NoiseRendezvousConnectProvider`] so recovery can fetch a fresh + /// bundle for each reconnect. + #[tracing::instrument( + name = "codex.exec_server.remote.harness.connect", + skip_all, + fields( + otel.kind = "client", + otel.name = "codex.exec_server.remote.harness.connect", + ) + )] + pub async fn connect_noise_rendezvous( + args: NoiseRendezvousConnectArgs, + ) -> Result { + let opened = Self::open_noise_rendezvous_connection(args).await?; + let ready = Self::finish_noise_rendezvous_connection(opened).await?; + Self::connect_with_recovery_and_noise_context( + ready.connection, + ready.options, + /*reconnect_strategy*/ None, + ready.noise_context, + ) + .await + } + + #[tracing::instrument( + name = "codex.exec_server.remote.noise.websocket_connect", + skip_all, + fields( + otel.kind = "client", + otel.name = "codex.exec_server.remote.noise.websocket_connect", + environment_id = %args.bundle.environment_id, + executor_registration_id = %args.bundle.executor_registration_id, + ) + )] + async fn open_noise_rendezvous_connection( + args: NoiseRendezvousConnectArgs, + ) -> Result { + ensure_rustls_crypto_provider(); + // Keep the registry-issued URL, key, and authorization together for this + // connection attempt. + let NoiseRendezvousConnectArgs { + bundle, + harness_identity, + client_name, + connect_timeout, + initialize_timeout, + resume_session_id, + http_client_factory, + } = args; + let NoiseRendezvousConnectBundle { + websocket_url, + environment_id, + executor_registration_id, + executor_public_key, + harness_key_authorization, + } = bundle; + let diagnostic_url = websocket_url + .split(['?', '#']) + .next() + .unwrap_or(websocket_url.as_str()) + .to_string(); + let mut request = websocket_url + .as_str() + .into_client_request() + .map_err(|source| ExecServerError::WebSocketConnect { + url: diagnostic_url.clone(), + source, + })?; + request.headers_mut().extend(current_rendezvous_headers()); + let (stream, _) = timeout( + connect_timeout, + WebSocketConnector::new_with_tls_mode( + &http_client_factory, + WebSocketTlsMode::TungsteniteDefault, + ) + .map_err(|error| ExecServerError::WebSocketConfiguration(error.to_string()))? + .with_tcp_nodelay() + .connect(request, noise_relay_websocket_config()), + ) + .await + .map_err(|_| ExecServerError::WebSocketConnectTimeout { + url: diagnostic_url.clone(), + timeout: connect_timeout, + })? + .map_err(|source| ExecServerError::WebSocketConnect { + url: diagnostic_url.clone(), + source, + })?; + + let connection_label = format!("Noise exec-server rendezvous websocket {diagnostic_url}"); + let connection = noise_harness_connection_from_websocket_with_readiness( + stream, + NoiseHarnessConnectionArgs { + connection_label, + environment_id, + executor_registration_id, + identity: harness_identity, + responder_public_key: executor_public_key, + harness_key_authorization, + }, + ); + Ok(OpenNoiseRendezvousConnection { + connection: connection.connection, + options: ExecServerClientConnectOptions { + client_name, + initialize_timeout, + resume_session_id, + }, + handshake_ready: connection.handshake_ready, + }) + } + + #[tracing::instrument( + name = "codex.exec_server.remote.noise.handshake", + skip_all, + parent = initialize_span, + fields( + otel.kind = "client", + otel.name = "codex.exec_server.remote.noise.handshake", + ) + )] + async fn wait_for_noise_handshake( + handshake_ready: &mut tokio::sync::oneshot::Receiver<()>, + deadline: Instant, + initialize_timeout: Duration, + initialize_span: &tracing::Span, + ) -> Result<(), ExecServerError> { + match timeout_at(deadline, handshake_ready).await { + Ok(Ok(())) => Ok(()), + Ok(Err(_)) => Err(ExecServerError::Disconnected( + "Noise harness handshake failed before connection became ready".to_string(), + )), + Err(_) => Err(ExecServerError::InitializeTimedOut { + timeout: initialize_timeout, + }), + } + } + + async fn finish_noise_rendezvous_connection( + mut connection: OpenNoiseRendezvousConnection, + ) -> Result { + // Preserve the legacy initialize request span as the post-WebSocket + // startup parent while making its two child operations visible. + let initialize_timeout = connection.options.initialize_timeout; + let noise_context = NoiseInitializeContext { + span: tracing::info_span!( + "codex.exec_server.request", + otel.kind = "client", + otel.name = "initialize", + method = "initialize", + ), + timeout_for_error: initialize_timeout, + }; + let deadline = Instant::now() + initialize_timeout; + let readiness = Self::wait_for_noise_handshake( + &mut connection.handshake_ready, + deadline, + initialize_timeout, + &noise_context.span, + ) + .await; + if let Err(error) = readiness { + // Unlike the normal connect path, the connection has not reached + // RpcClient yet, so its Drop implementation cannot abort the + // transport task for us. + connection.connection.transport.terminate(); + for task in &connection.connection.task_handles { + task.abort(); + } + return Err(error); + } + let mut options = connection.options; + options.initialize_timeout = deadline.saturating_duration_since(Instant::now()); + Ok(ReadyNoiseRendezvousConnection { + connection: connection.connection, + options, + noise_context, + }) + } + + pub(crate) async fn connect_stdio_command( + args: StdioExecServerConnectArgs, + ) -> Result { + let mut child = stdio_command_process(&args.command) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .map_err(ExecServerError::Spawn)?; + + let stdin = child.stdin.take().ok_or_else(|| { + ExecServerError::Protocol("spawned exec-server command has no stdin".to_string()) + })?; + let stdout = child.stdout.take().ok_or_else(|| { + ExecServerError::Protocol("spawned exec-server command has no stdout".to_string()) + })?; + if let Some(stderr) = child.stderr.take() { + tokio::spawn(async move { + let mut lines = BufReader::new(stderr).lines(); + loop { + match lines.next_line().await { + Ok(Some(line)) => debug!("exec-server stdio stderr: {line}"), + Ok(None) => break, + Err(err) => { + warn!("failed to read exec-server stdio stderr: {err}"); + break; + } + } + } + }); + } + + Self::connect( + JsonRpcConnection::from_stdio(stdout, stdin, "exec-server stdio command".to_string()) + .with_child_process(child), + args.into(), + ) + .await + } +} + +fn is_rendezvous_harness_url(websocket_url: &str) -> bool { + let Some((_path, query)) = websocket_url.split_once('?') else { + return false; + }; + query + .split('&') + .filter_map(|pair| pair.split_once('=')) + .any(|(key, value)| key == "role" && value == "harness") +} + +fn stdio_command_process(stdio_command: &StdioExecServerCommand) -> Command { + let mut command = Command::new(&stdio_command.program); + command.args(&stdio_command.args); + command.envs(&stdio_command.env); + scrub_non_inheritable_env_vars(command.as_std_mut()); + if let Some(cwd) = &stdio_command.cwd { + command.current_dir(cwd); + } + #[cfg(unix)] + command.process_group(0); + command +} + +#[cfg(test)] +#[path = "client_transport_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/client_transport_tests.rs b/codex-rs/exec-server/src/client_transport_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..63345c227abd010e7fe4ed6bf9a643124b5d1485 --- /dev/null +++ b/codex-rs/exec-server/src/client_transport_tests.rs @@ -0,0 +1,581 @@ +use std::collections::VecDeque; +use std::future::Future; +use std::sync::Arc; +use std::sync::Mutex; + +use anyhow::Result; +use codex_exec_server_protocol::JSONRPCMessage; +use futures::FutureExt; +use futures::SinkExt; +use futures::StreamExt; +use futures::future::BoxFuture; +use pretty_assertions::assert_eq; +use tokio::io::AsyncBufReadExt; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::io::BufReader; +use tokio::io::duplex; +use tokio::net::TcpListener; +use tokio_tungstenite::accept_async; +use tokio_tungstenite::tungstenite::Message; + +use super::ExecServerClient; +use super::ExecServerReconnectStrategy; +use super::INITIAL_REGISTRY_MAX_RETRIES; +use super::INITIAL_REGISTRY_OPERATION_TIMEOUT; +use super::INITIAL_REGISTRY_REQUEST_TIMEOUT; +use crate::ExecServerError; +use crate::NoiseChannelIdentity; +use crate::NoiseChannelPublicKey; +use crate::NoiseRendezvousConnectArgs; +use crate::NoiseRendezvousConnectBundle; +use crate::NoiseRendezvousConnectProvider; +use crate::client::NoiseInitializeContext; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT; +use crate::client_api::ExecServerClientConnectOptions; +use crate::connection::JsonRpcConnection; +use crate::noise_channel::PendingResponderHandshake; +use crate::noise_channel::noise_channel_prologue; +use crate::protocol::INITIALIZE_METHOD; +use crate::relay::RelayFrameBodyKind; +use crate::relay::decode_relay_message_frame; +use crate::relay::encode_relay_message_frame; +use crate::relay_proto::RelayMessageFrame; + +#[derive(Default)] +struct SequenceNoiseConnectProvider { + bundles: + Mutex>>>, + returned_urls: Mutex>, + requested_keys: Mutex>, +} + +impl SequenceNoiseConnectProvider { + fn push_response( + &self, + response: impl Future> + + Send + + 'static, + ) { + self.bundles.lock().unwrap().push_back(response.boxed()); + } + + fn push_error(&self, error: ExecServerError) { + self.push_response(futures::future::ready(Err(error))); + } + + fn push_pending(&self) { + self.push_response(futures::future::pending()); + } + + fn requested_keys(&self) -> Vec { + self.requested_keys.lock().unwrap().clone() + } + + fn assert_requested_identity(&self, identity: &NoiseChannelIdentity, requests: usize) { + assert_eq!(self.requested_keys(), vec![identity.public_key(); requests]); + } + + fn returned_urls(&self) -> Vec { + self.returned_urls + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + } + + async fn connect( + self: &Arc, + identity: &NoiseChannelIdentity, + ) -> Result< + ( + super::JsonRpcConnection, + super::ExecServerClientConnectOptions, + ), + ExecServerError, + > { + let provider: Arc = self.clone(); + ExecServerClient::open_initial_noise_rendezvous_connection( + &provider, + identity, + codex_http_client::HttpClientFactory::new( + codex_http_client::OutboundProxyPolicy::ReqwestDefault, + ), + ) + .await + .map(|(ready, _)| (ready.connection, ready.options)) + } +} + +impl NoiseRendezvousConnectProvider for SequenceNoiseConnectProvider { + fn connect_bundle( + &self, + harness_public_key: NoiseChannelPublicKey, + ) -> BoxFuture<'_, Result> { + self.requested_keys.lock().unwrap().push(harness_public_key); + let response = self + .bundles + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .pop_front() + .expect("test Noise provider exhausted"); + Box::pin(async move { + let result = response.await; + if let Ok(bundle) = &result { + self.returned_urls + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .push(bundle.websocket_url.clone()); + } + result + }) + } +} + +fn test_bundle(websocket_url: String) -> Result { + Ok(NoiseRendezvousConnectBundle { + websocket_url, + environment_id: "environment".to_string(), + executor_registration_id: "registration".to_string(), + executor_public_key: NoiseChannelIdentity::generate()?.public_key(), + harness_key_authorization: "authorization".to_string(), + }) +} + +fn registry_error(status: http::StatusCode, code: &str) -> ExecServerError { + ExecServerError::EnvironmentRegistryHttp { + status, + code: Some(code.to_string()), + message: "registry unavailable".to_string(), + } +} + +#[tokio::test] +async fn noise_handshake_uses_initialize_timeout() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let server = tokio::spawn(async move { + let (socket, _) = listener.accept().await?; + let mut websocket = accept_async(socket).await?; + // Drain the frames sent before the harness waits for the responder, + // then verify that a timed-out readiness wait closes the socket. + assert!(websocket.next().await.is_some()); + assert!(websocket.next().await.is_some()); + let closed = + tokio::time::timeout(std::time::Duration::from_secs(1), websocket.next()).await?; + assert!( + matches!(closed, None | Some(Ok(Message::Close(_))) | Some(Err(_))), + "timed-out Noise handshake must close its websocket" + ); + anyhow::Ok(()) + }); + let initialize_timeout = std::time::Duration::from_millis(1); + let opened = ExecServerClient::open_noise_rendezvous_connection(NoiseRendezvousConnectArgs { + bundle: NoiseRendezvousConnectBundle { + websocket_url, + environment_id: "environment".to_string(), + executor_registration_id: "registration".to_string(), + executor_public_key: NoiseChannelIdentity::generate()?.public_key(), + harness_key_authorization: "authorization".to_string(), + }, + harness_identity: NoiseChannelIdentity::generate()?, + client_name: "test".to_string(), + connect_timeout: DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + initialize_timeout, + resume_session_id: None, + http_client_factory: codex_http_client::HttpClientFactory::new( + codex_http_client::OutboundProxyPolicy::ReqwestDefault, + ), + }) + .await?; + + let error = ExecServerClient::finish_noise_rendezvous_connection(opened) + .await + .err() + .expect("stalled Noise handshake must time out"); + assert!(matches!( + error, + ExecServerError::InitializeTimedOut { timeout } if timeout == initialize_timeout + )); + + server.await??; + Ok(()) +} + +#[tokio::test] +async fn deferred_initialize_timeout_reports_configured_budget() { + let (client_stdin, server_reader) = duplex(1 << 20); + let (server_writer, client_stdout) = duplex(1 << 20); + let server = tokio::spawn(async move { + let _server_writer = server_writer; + let mut lines = BufReader::new(server_reader).lines(); + let line = lines + .next_line() + .await + .expect("initialize read should succeed") + .expect("initialize request should arrive"); + let request: JSONRPCMessage = + serde_json::from_str(&line).expect("initialize request should parse"); + assert!( + matches!( + request, + JSONRPCMessage::Request(ref request) if request.method == INITIALIZE_METHOD + ), + "expected initialize request, got {request:?}" + ); + futures::future::pending::<()>().await; + }); + let configured_timeout = std::time::Duration::from_secs(10); + let error = ExecServerClient::connect_with_recovery_and_noise_context( + JsonRpcConnection::from_stdio( + client_stdout, + client_stdin, + "timeout-test-client".to_string(), + ), + ExecServerClientConnectOptions { + client_name: "timeout-test-client".to_string(), + initialize_timeout: std::time::Duration::from_millis(1), + resume_session_id: None, + }, + /*reconnect_strategy*/ None, + NoiseInitializeContext { + span: tracing::info_span!("codex.exec_server.request"), + timeout_for_error: configured_timeout, + }, + ) + .await + .err() + .expect("initialize RPC must time out"); + assert!(matches!( + error, + ExecServerError::InitializeTimedOut { timeout } if timeout == configured_timeout + )); + server.abort(); + let _ = server.await; +} + +#[tokio::test(start_paused = true)] +async fn initial_noise_connection_bounds_offline_retries() -> Result<()> { + let sequence = Arc::new(SequenceNoiseConnectProvider::default()); + for _ in 0..=INITIAL_REGISTRY_MAX_RETRIES { + sequence.push_error(registry_error( + http::StatusCode::CONFLICT, + "environment_offline", + )); + } + let identity = NoiseChannelIdentity::generate()?; + let started = tokio::time::Instant::now(); + let error = sequence + .connect(&identity) + .await + .err() + .expect("offline retries must end"); + + assert!(crate::client::is_environment_offline_error(&error)); + let requests = sequence.requested_keys().len(); + assert!((4..=INITIAL_REGISTRY_MAX_RETRIES as usize + 1).contains(&requests)); + sequence.assert_requested_identity(&identity, requests); + assert!(started.elapsed() <= INITIAL_REGISTRY_OPERATION_TIMEOUT); + Ok(()) +} + +#[tokio::test(start_paused = true)] +async fn initial_noise_connection_bounds_a_stalled_retry_request() -> Result<()> { + let sequence = Arc::new(SequenceNoiseConnectProvider::default()); + sequence.push_error(registry_error( + http::StatusCode::CONFLICT, + "environment_offline", + )); + for _ in 0..INITIAL_REGISTRY_MAX_RETRIES { + sequence.push_pending(); + } + let identity = NoiseChannelIdentity::generate()?; + let started = tokio::time::Instant::now(); + let error = sequence + .connect(&identity) + .await + .err() + .expect("stalled retry must time out"); + + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryRequest(error) if error.is_timeout() + )); + assert_eq!(started.elapsed(), INITIAL_REGISTRY_OPERATION_TIMEOUT); + let requests = sequence.requested_keys().len(); + assert!((2..=3).contains(&requests)); + sequence.assert_requested_identity(&identity, requests); + Ok(()) +} + +#[tokio::test(start_paused = true)] +async fn initial_noise_connection_bounds_a_stalled_initial_request() -> Result<()> { + let sequence = Arc::new(SequenceNoiseConnectProvider::default()); + for _ in 0..=INITIAL_REGISTRY_MAX_RETRIES { + sequence.push_pending(); + } + let identity = NoiseChannelIdentity::generate()?; + let started = tokio::time::Instant::now(); + + let error = sequence + .connect(&identity) + .await + .err() + .expect("stalled initial request must time out"); + + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryRequest(error) if error.is_timeout() + )); + assert_eq!(started.elapsed(), INITIAL_REGISTRY_OPERATION_TIMEOUT); + let requests = sequence.requested_keys().len(); + assert!((2..=3).contains(&requests)); + sequence.assert_requested_identity(&identity, requests); + Ok(()) +} + +#[tokio::test(start_paused = true)] +async fn initial_noise_connection_retries_a_stalled_initial_request() -> Result<()> { + let sequence = Arc::new(SequenceNoiseConnectProvider::default()); + sequence.push_pending(); + sequence.push_error(registry_error(http::StatusCode::FORBIDDEN, "forbidden")); + let identity = NoiseChannelIdentity::generate()?; + let started = tokio::time::Instant::now(); + + let error = sequence + .connect(&identity) + .await + .err() + .expect("terminal response must stop the retry sequence"); + + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryHttp { + status: http::StatusCode::FORBIDDEN, + .. + } + )); + assert!(started.elapsed() >= INITIAL_REGISTRY_REQUEST_TIMEOUT); + assert!(started.elapsed() < INITIAL_REGISTRY_OPERATION_TIMEOUT); + sequence.assert_requested_identity(&identity, /*requests*/ 2); + Ok(()) +} + +#[tokio::test(start_paused = true)] +async fn initial_noise_connection_retries_transient_registry_statuses() -> Result<()> { + for status in [ + http::StatusCode::REQUEST_TIMEOUT, + http::StatusCode::TOO_MANY_REQUESTS, + http::StatusCode::INTERNAL_SERVER_ERROR, + http::StatusCode::BAD_GATEWAY, + http::StatusCode::SERVICE_UNAVAILABLE, + ] { + let sequence = Arc::new(SequenceNoiseConnectProvider::default()); + sequence.push_error(registry_error(status, "temporarily_unavailable")); + sequence.push_error(registry_error(http::StatusCode::FORBIDDEN, "forbidden")); + let identity = NoiseChannelIdentity::generate()?; + + let error = sequence + .connect(&identity) + .await + .err() + .expect("terminal response must stop the retry sequence"); + + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryHttp { + status: http::StatusCode::FORBIDDEN, + .. + } + )); + sequence.assert_requested_identity(&identity, /*requests*/ 2); + } + Ok(()) +} + +#[tokio::test(start_paused = true)] +async fn initial_noise_connection_retries_registry_request_timeouts() -> Result<()> { + let sequence = Arc::new(SequenceNoiseConnectProvider::default()); + sequence.push_error(ExecServerError::EnvironmentRegistryRequest( + codex_http_client::RouteAwareRequestError::Timeout, + )); + sequence.push_error(registry_error(http::StatusCode::FORBIDDEN, "forbidden")); + let identity = NoiseChannelIdentity::generate()?; + + let error = sequence + .connect(&identity) + .await + .err() + .expect("terminal response must stop the retry sequence"); + + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryHttp { + status: http::StatusCode::FORBIDDEN, + .. + } + )); + sequence.assert_requested_identity(&identity, /*requests*/ 2); + Ok(()) +} + +#[tokio::test(start_paused = true)] +async fn initial_noise_connection_does_not_retry_permanent_registry_errors() -> Result<()> { + for (status, code) in [ + (http::StatusCode::UNAUTHORIZED, "unauthorized"), + (http::StatusCode::FORBIDDEN, "forbidden"), + (http::StatusCode::BAD_REQUEST, "bad_request"), + (http::StatusCode::NOT_FOUND, "environment_not_found"), + (http::StatusCode::CONFLICT, "registration_conflict"), + (http::StatusCode::CONFLICT, "route_unavailable"), + ] { + // A terminal error must also stop a retry sequence already in progress. + for initial_offline in [false, true] { + let sequence = Arc::new(SequenceNoiseConnectProvider::default()); + if initial_offline { + sequence.push_error(registry_error( + http::StatusCode::CONFLICT, + "environment_offline", + )); + } + sequence.push_error(registry_error(status, code)); + let identity = NoiseChannelIdentity::generate()?; + let error = sequence + .connect(&identity) + .await + .err() + .expect("other errors must propagate"); + assert!( + matches!(error, ExecServerError::EnvironmentRegistryHttp { status: actual_status, code: Some(actual_code), .. } if actual_status == status && actual_code == code) + ); + sequence.assert_requested_identity(&identity, 1 + usize::from(initial_offline)); + } + } + Ok(()) +} + +#[tokio::test(start_paused = true)] +async fn noise_session_resume_leaves_offline_retries_to_recovery() -> Result<()> { + let sequence = Arc::new(SequenceNoiseConnectProvider::default()); + sequence.push_error(registry_error( + http::StatusCode::CONFLICT, + "environment_offline", + )); + let identity = NoiseChannelIdentity::generate()?; + let strategy = ExecServerReconnectStrategy::NoiseRendezvous { + executor_public_key: NoiseChannelIdentity::generate()?.public_key(), + provider: sequence.clone(), + identity: identity.clone(), + client_name: "test".to_string(), + connect_timeout: DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + initialize_timeout: DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT, + http_client_factory: codex_http_client::HttpClientFactory::new( + codex_http_client::OutboundProxyPolicy::ReqwestDefault, + ), + }; + let started = tokio::time::Instant::now(); + let error = strategy + .resume("session") + .await + .err() + .expect("resume must return the offline error"); + assert!(crate::client::is_environment_offline_error(&error)); + assert_eq!(started.elapsed(), std::time::Duration::ZERO); + sequence.assert_requested_identity(&identity, /*requests*/ 1); + Ok(()) +} + +#[tokio::test] +async fn initial_noise_connection_refreshes_bundle_after_exhausting_initial_retries() -> Result<()> +{ + let unauthorized_listener = TcpListener::bind("127.0.0.1:0").await?; + let unauthorized_url = format!("ws://{}", unauthorized_listener.local_addr()?); + let unauthorized_server = tokio::spawn(async move { + let (mut socket, _) = unauthorized_listener.accept().await?; + let mut request = [0_u8; 4096]; + let _ = socket.read(&mut request).await?; + socket + .write_all( + b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\nConnection: close\r\n\r\n", + ) + .await?; + socket.shutdown().await?; + anyhow::Ok(()) + }); + let accepted_listener = TcpListener::bind("127.0.0.1:0").await?; + let accepted_url = format!("ws://{}", accepted_listener.local_addr()?); + let executor_identity = NoiseChannelIdentity::generate()?; + let executor_public_key = executor_identity.public_key(); + let accepted_server = tokio::spawn(async move { + let (socket, _) = accepted_listener.accept().await?; + let mut websocket = accept_async(socket).await?; + let Message::Binary(resume_payload) = websocket.next().await.unwrap()? else { + anyhow::bail!("expected Noise relay resume frame"); + }; + let resume = decode_relay_message_frame(resume_payload.as_ref())?; + assert_eq!(resume.validate()?, RelayFrameBodyKind::Resume); + let Message::Binary(handshake_payload) = websocket.next().await.unwrap()? else { + anyhow::bail!("expected Noise relay handshake frame"); + }; + let handshake = decode_relay_message_frame(handshake_payload.as_ref())?; + let stream_id = handshake.stream_id.clone(); + let prologue = noise_channel_prologue("environment", "registration", &stream_id); + let pending = PendingResponderHandshake::read_request( + &executor_identity, + &prologue, + &handshake.into_handshake_payload()?, + )?; + let (_transport, response) = pending.complete()?; + websocket + .send(Message::Binary( + encode_relay_message_frame(&RelayMessageFrame::handshake(stream_id, response)) + .into(), + )) + .await?; + anyhow::Ok(()) + }); + let sequence = Arc::new(SequenceNoiseConnectProvider::default()); + let unauthorized_bundle = test_bundle(unauthorized_url.clone())?; + let mut accepted_bundle = test_bundle(accepted_url.clone())?; + accepted_bundle.executor_public_key = executor_public_key; + sequence.push_response(async { + tokio::time::pause(); + Err(registry_error( + http::StatusCode::CONFLICT, + "environment_offline", + )) + }); + for _ in 1..INITIAL_REGISTRY_MAX_RETRIES { + sequence.push_error(registry_error( + http::StatusCode::CONFLICT, + "environment_offline", + )); + } + sequence.push_response(async move { + tokio::time::resume(); + Ok(unauthorized_bundle) + }); + sequence.push_response(async { + tokio::time::pause(); + Err(registry_error( + http::StatusCode::CONFLICT, + "environment_offline", + )) + }); + sequence.push_response(async move { + tokio::time::resume(); + Ok(accepted_bundle) + }); + let identity = NoiseChannelIdentity::generate()?; + + let _connection = sequence.connect(&identity).await?; + + assert_eq!( + sequence.returned_urls(), + vec![unauthorized_url, accepted_url] + ); + sequence.assert_requested_identity(&identity, INITIAL_REGISTRY_MAX_RETRIES as usize + 3); + unauthorized_server.await??; + accepted_server.await??; + Ok(()) +} diff --git a/codex-rs/exec-server/src/connection.rs b/codex-rs/exec-server/src/connection.rs new file mode 100644 index 0000000000000000000000000000000000000000..249b63330e59cd8bb553892976809d48866936c6 --- /dev/null +++ b/codex-rs/exec-server/src/connection.rs @@ -0,0 +1,1042 @@ +#[cfg(windows)] +use std::process::Stdio; +use std::sync::Arc; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; +use std::time::Duration; +use std::time::Instant; + +use axum::extract::ws::Message as AxumWebSocketMessage; +use axum::extract::ws::WebSocket as AxumWebSocket; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCRequest; +use futures::Sink; +use futures::SinkExt; +use futures::Stream; +use futures::StreamExt; +use tokio::io::AsyncRead; +use tokio::io::AsyncWrite; +use tokio::process::Child; +use tokio::sync::mpsc; +use tokio::sync::watch; +use tokio::time::timeout; +#[cfg(test)] +use tokio_tungstenite::WebSocketStream; +use tokio_tungstenite::tungstenite::Message; +use tracing::debug; +use tracing::warn; + +use tokio::io::AsyncBufReadExt; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::io::BufReader; +use tokio::io::BufWriter; + +pub(crate) const CHANNEL_CAPACITY: usize = 128; +// Match the existing serialized JSON-RPC message ceiling used by Noise and +// WebSocket transports so stdio has the same per-message bound. +const MAX_STDIO_JSONRPC_MESSAGE_LEN: usize = 64 * 1024 * 1024; +const STDIO_TERMINATION_GRACE_PERIOD: Duration = Duration::from_secs(2); +#[cfg(test)] +pub(crate) const WEBSOCKET_KEEPALIVE_INTERVAL: Duration = Duration::from_millis(25); +#[cfg(not(test))] +pub(crate) const WEBSOCKET_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(30); + +#[derive(Debug)] +pub(crate) enum JsonRpcConnectionEvent { + Message(JSONRPCMessage), + QueuedRequest { + request: JSONRPCRequest, + request_span: tracing::Span, + queued_at: Instant, + }, + MalformedMessage { + reason: String, + }, + Disconnected { + reason: Option, + }, +} + +impl JsonRpcConnectionEvent { + pub(crate) fn message(message: JSONRPCMessage) -> Self { + let JSONRPCMessage::Request(request) = message else { + return Self::Message(message); + }; + + let queued_at = Instant::now(); + let request_span = tracing::info_span!( + "codex.exec_server.request", + otel.kind = "server", + otel.name = "unknown", + method = request.method.as_str(), + result = tracing::field::Empty, + ); + if let Some(trace) = &request.trace + && !codex_otel::set_parent_from_w3c_trace_context(&request_span, trace) + { + warn!( + method = request.method.as_str(), + "ignoring invalid inbound exec-server trace carrier" + ); + } + + Self::QueuedRequest { + request, + request_span, + queued_at, + } + } +} + +#[derive(Clone)] +pub(crate) enum JsonRpcTransport { + // Plain means no child process; transport bytes may still be encrypted. + Plain, + Stdio { transport: StdioTransport }, +} + +impl JsonRpcTransport { + fn from_child_process(child_process: Child) -> Self { + Self::Stdio { + transport: StdioTransport::spawn(child_process), + } + } + + pub(crate) fn terminate(&self) { + match self { + Self::Plain => {} + Self::Stdio { transport } => transport.terminate(), + } + } +} + +#[derive(Clone)] +pub(crate) struct StdioTransport { + handle: Arc, +} + +struct StdioTransportHandle { + terminate_tx: watch::Sender, + terminate_requested: AtomicBool, +} + +impl StdioTransport { + fn spawn(child_process: Child) -> Self { + let (terminate_tx, terminate_rx) = watch::channel(false); + let handle = Arc::new(StdioTransportHandle { + terminate_tx, + terminate_requested: AtomicBool::new(false), + }); + spawn_stdio_child_supervisor(child_process, terminate_rx); + Self { handle } + } + + fn terminate(&self) { + self.handle.terminate(); + } +} + +impl StdioTransportHandle { + fn terminate(&self) { + if !self.terminate_requested.swap(true, Ordering::AcqRel) { + let _ = self.terminate_tx.send(true); + } + } +} + +impl Drop for StdioTransportHandle { + fn drop(&mut self) { + self.terminate(); + } +} + +fn spawn_stdio_child_supervisor(mut child_process: Child, mut terminate_rx: watch::Receiver) { + let process_group_id = child_process.id(); + tokio::spawn(async move { + tokio::select! { + result = child_process.wait() => { + log_stdio_child_wait_result(result); + kill_process_tree(&mut child_process, process_group_id); + } + () = wait_for_stdio_termination(&mut terminate_rx) => { + terminate_stdio_child(&mut child_process, process_group_id).await; + } + } + }); +} + +async fn wait_for_stdio_termination(terminate_rx: &mut watch::Receiver) { + loop { + if *terminate_rx.borrow() { + return; + } + if terminate_rx.changed().await.is_err() { + return; + } + } +} + +async fn terminate_stdio_child(child_process: &mut Child, process_group_id: Option) { + terminate_process_tree(child_process, process_group_id); + match timeout(STDIO_TERMINATION_GRACE_PERIOD, child_process.wait()).await { + Ok(result) => { + log_stdio_child_wait_result(result); + } + Err(_) => { + kill_process_tree(child_process, process_group_id); + log_stdio_child_wait_result(child_process.wait().await); + } + } +} + +fn terminate_process_tree(child_process: &mut Child, process_group_id: Option) { + let Some(process_group_id) = process_group_id else { + kill_direct_child(child_process, "terminate"); + return; + }; + + #[cfg(unix)] + if let Err(err) = codex_utils_pty::process_group::terminate_process_group(process_group_id) { + warn!("failed to terminate exec-server stdio process group {process_group_id}: {err}"); + kill_direct_child(child_process, "terminate"); + } + + #[cfg(windows)] + if !kill_windows_process_tree(process_group_id) { + kill_direct_child(child_process, "terminate"); + } + + #[cfg(not(any(unix, windows)))] + { + let _ = process_group_id; + kill_direct_child(child_process, "terminate"); + } +} + +fn kill_process_tree(child_process: &mut Child, process_group_id: Option) { + let Some(process_group_id) = process_group_id else { + kill_direct_child(child_process, "kill"); + return; + }; + + #[cfg(unix)] + if let Err(err) = codex_utils_pty::process_group::kill_process_group(process_group_id) { + warn!("failed to kill exec-server stdio process group {process_group_id}: {err}"); + } + + #[cfg(windows)] + if !kill_windows_process_tree(process_group_id) { + kill_direct_child(child_process, "kill"); + } + + #[cfg(not(any(unix, windows)))] + { + let _ = process_group_id; + kill_direct_child(child_process, "kill"); + } +} + +fn kill_direct_child(child_process: &mut Child, action: &str) { + if let Err(err) = child_process.start_kill() { + debug!("failed to {action} exec-server stdio child: {err}"); + } +} + +#[cfg(windows)] +fn kill_windows_process_tree(pid: u32) -> bool { + let pid = pid.to_string(); + match std::process::Command::new("taskkill") + .args(["/PID", pid.as_str(), "/T", "/F"]) + .stdin(Stdio::null()) + .stdout(Stdio::null()) + .stderr(Stdio::null()) + .status() + { + Ok(status) => status.success(), + Err(err) => { + warn!("failed to run taskkill for exec-server stdio process tree {pid}: {err}"); + false + } + } +} + +fn log_stdio_child_wait_result(result: std::io::Result) { + if let Err(err) = result { + debug!("failed to wait for exec-server stdio child: {err}"); + } +} + +pub(crate) struct JsonRpcConnection { + pub(crate) outgoing_tx: mpsc::Sender, + pub(crate) incoming_rx: mpsc::Receiver, + pub(crate) disconnected_rx: watch::Receiver, + pub(crate) task_handles: Vec>, + pub(crate) transport: JsonRpcTransport, +} + +impl JsonRpcConnection { + pub(crate) fn from_stdio(reader: R, writer: W, connection_label: String) -> Self + where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, + { + Self::from_stdio_with_max_message_len( + reader, + writer, + connection_label, + MAX_STDIO_JSONRPC_MESSAGE_LEN, + ) + } + + fn from_stdio_with_max_message_len( + reader: R, + writer: W, + connection_label: String, + max_message_len: usize, + ) -> Self + where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, + { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (incoming_tx, incoming_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (disconnected_tx, disconnected_rx) = watch::channel(false); + + let reader_label = connection_label.clone(); + let incoming_tx_for_reader = incoming_tx.clone(); + let disconnected_tx_for_reader = disconnected_tx.clone(); + // Read one byte past the payload limit so an unterminated oversized + // message fails promptly. A trailing CR gets one more byte of lookahead + // because it may be the first half of a valid CRLF terminator. + let read_limit = u64::try_from(max_message_len.saturating_add(1)).unwrap_or(u64::MAX); + let reader_task = tokio::spawn(async move { + let mut reader = BufReader::new(reader); + let mut line = String::new(); + loop { + line.clear(); + let read_result = (&mut reader).take(read_limit).read_line(&mut line).await; + match read_result { + Ok(0) => { + send_disconnected( + &incoming_tx_for_reader, + &disconnected_tx_for_reader, + /*reason*/ None, + ) + .await; + break; + } + Ok(_) => { + if line.ends_with('\n') { + line.pop(); + if line.ends_with('\r') { + line.pop(); + } + } else if line.len() > max_message_len && line.ends_with('\r') { + match reader.read_u8().await { + Ok(b'\n') => { + line.pop(); + } + Ok(_) => {} + Err(err) if err.kind() == std::io::ErrorKind::UnexpectedEof => {} + Err(err) => { + send_disconnected( + &incoming_tx_for_reader, + &disconnected_tx_for_reader, + Some(format!( + "failed to read JSON-RPC message from {reader_label}: {err}" + )), + ) + .await; + break; + } + } + } + if line.len() > max_message_len { + send_disconnected( + &incoming_tx_for_reader, + &disconnected_tx_for_reader, + Some(format!( + "JSON-RPC message from {reader_label} exceeds maximum length of {max_message_len} bytes" + )), + ) + .await; + break; + } + if line.trim().is_empty() { + continue; + } + match serde_json::from_str::(&line) { + Ok(message) => { + if incoming_tx_for_reader + .send(JsonRpcConnectionEvent::message(message)) + .await + .is_err() + { + break; + } + } + Err(err) => { + send_malformed_message( + &incoming_tx_for_reader, + Some(format!( + "failed to parse JSON-RPC message from {reader_label}: {err}" + )), + ) + .await; + } + } + } + Err(err) => { + send_disconnected( + &incoming_tx_for_reader, + &disconnected_tx_for_reader, + Some(format!( + "failed to read JSON-RPC message from {reader_label}: {err}" + )), + ) + .await; + break; + } + } + } + }); + + let writer_task = tokio::spawn(async move { + let mut writer = BufWriter::new(writer); + while let Some(message) = outgoing_rx.recv().await { + if let Err(err) = write_jsonrpc_line_message(&mut writer, &message).await { + send_disconnected( + &incoming_tx, + &disconnected_tx, + Some(format!( + "failed to write JSON-RPC message to {connection_label}: {err}" + )), + ) + .await; + break; + } + } + }); + + Self { + outgoing_tx, + incoming_rx, + disconnected_rx, + task_handles: vec![reader_task, writer_task], + transport: JsonRpcTransport::Plain, + } + } + + pub(crate) fn from_websocket(stream: T, connection_label: String) -> Self + where + T: Sink + Stream> + Unpin + Send + 'static, + E: std::fmt::Display + Send + 'static, + { + Self::from_websocket_stream(stream, connection_label, /*ping_interval*/ None) + } + + pub(crate) fn from_axum_websocket(stream: AxumWebSocket, connection_label: String) -> Self { + Self::from_websocket_stream(stream, connection_label, Some(WEBSOCKET_KEEPALIVE_INTERVAL)) + } + + fn from_websocket_stream( + mut websocket: T, + connection_label: String, + ping_interval: Option, + ) -> Self + where + T: Sink + Stream> + Unpin + Send + 'static, + M: JsonRpcWebSocketMessage, + E: std::fmt::Display + Send + 'static, + { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (incoming_tx, incoming_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (disconnected_tx, disconnected_rx) = watch::channel(false); + + let websocket_task = tokio::spawn(async move { + let mut ping_interval = ping_interval.map(|ping_interval| { + let mut interval = tokio::time::interval_at( + tokio::time::Instant::now() + ping_interval, + ping_interval, + ); + interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + interval + }); + + loop { + tokio::select! { + maybe_message = outgoing_rx.recv() => { + let Some(message) = maybe_message else { + break; + }; + if let Err(reason) = send_websocket_jsonrpc_message( + &mut websocket, + &connection_label, + &message, + ) + .await + { + send_disconnected(&incoming_tx, &disconnected_tx, Some(reason)).await; + break; + } + } + _ = async { + match ping_interval.as_mut() { + Some(interval) => interval.tick().await, + None => std::future::pending().await, + } + } => { + if let Err(err) = websocket.send(M::ping()).await { + send_disconnected( + &incoming_tx, + &disconnected_tx, + Some(format!( + "failed to write websocket ping to {connection_label}: {err}" + )), + ) + .await; + break; + } + } + incoming_message = websocket.next() => { + match incoming_message { + Some(Ok(message)) => match message.parse_jsonrpc_frame() { + Ok(JsonRpcWebSocketFrame::Message(message)) => { + if incoming_tx + .send(JsonRpcConnectionEvent::message(message)) + .await + .is_err() + { + break; + } + } + Ok(JsonRpcWebSocketFrame::Close) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + /*reason*/ None, + ) + .await; + break; + } + Ok(JsonRpcWebSocketFrame::Ignore) => {} + Err(err) => { + send_malformed_message( + &incoming_tx, + Some(format!( + "failed to parse websocket JSON-RPC message from {connection_label}: {err}" + )), + ) + .await; + } + }, + Some(Err(err)) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + Some(format!( + "failed to read websocket JSON-RPC message from {connection_label}: {err}" + )), + ) + .await; + break; + } + None => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + /*reason*/ None, + ) + .await; + break; + } + } + } + } + } + }); + + Self { + outgoing_tx, + incoming_rx, + disconnected_rx, + task_handles: vec![websocket_task], + transport: JsonRpcTransport::Plain, + } + } + + pub(crate) fn with_child_process(mut self, child_process: Child) -> Self { + self.transport = JsonRpcTransport::from_child_process(child_process); + self + } +} + +enum JsonRpcWebSocketFrame { + Message(JSONRPCMessage), + Close, + Ignore, +} + +trait JsonRpcWebSocketMessage: Send + 'static { + fn parse_jsonrpc_frame(self) -> Result; + fn from_text(text: String) -> Self; + fn ping() -> Self; +} + +impl JsonRpcWebSocketMessage for Message { + fn parse_jsonrpc_frame(self) -> Result { + match self { + Message::Text(text) => { + serde_json::from_str(text.as_ref()).map(JsonRpcWebSocketFrame::Message) + } + Message::Binary(bytes) => { + serde_json::from_slice(bytes.as_ref()).map(JsonRpcWebSocketFrame::Message) + } + Message::Close(_) => Ok(JsonRpcWebSocketFrame::Close), + Message::Ping(_) | Message::Pong(_) | Message::Frame(_) => { + Ok(JsonRpcWebSocketFrame::Ignore) + } + } + } + + fn from_text(text: String) -> Self { + Self::Text(text.into()) + } + + fn ping() -> Self { + Self::Ping(Vec::new().into()) + } +} + +impl JsonRpcWebSocketMessage for AxumWebSocketMessage { + fn parse_jsonrpc_frame(self) -> Result { + match self { + AxumWebSocketMessage::Text(text) => { + serde_json::from_str(text.as_ref()).map(JsonRpcWebSocketFrame::Message) + } + AxumWebSocketMessage::Binary(bytes) => { + serde_json::from_slice(bytes.as_ref()).map(JsonRpcWebSocketFrame::Message) + } + AxumWebSocketMessage::Close(_) => Ok(JsonRpcWebSocketFrame::Close), + AxumWebSocketMessage::Ping(_) | AxumWebSocketMessage::Pong(_) => { + Ok(JsonRpcWebSocketFrame::Ignore) + } + } + } + + fn from_text(text: String) -> Self { + Self::Text(text.into()) + } + + fn ping() -> Self { + Self::Ping(Vec::new().into()) + } +} + +async fn send_disconnected( + incoming_tx: &mpsc::Sender, + disconnected_tx: &watch::Sender, + reason: Option, +) { + let _ = disconnected_tx.send(true); + let _ = incoming_tx + .send(JsonRpcConnectionEvent::Disconnected { reason }) + .await; +} + +async fn send_malformed_message( + incoming_tx: &mpsc::Sender, + reason: Option, +) { + let _ = incoming_tx + .send(JsonRpcConnectionEvent::MalformedMessage { + reason: reason.unwrap_or_else(|| "malformed JSON-RPC message".to_string()), + }) + .await; +} + +async fn write_jsonrpc_line_message( + writer: &mut BufWriter, + message: &JSONRPCMessage, +) -> std::io::Result<()> +where + W: AsyncWrite + Unpin, +{ + let encoded = + serialize_jsonrpc_message(message).map_err(|err| std::io::Error::other(err.to_string()))?; + writer.write_all(encoded.as_bytes()).await?; + writer.write_all(b"\n").await?; + writer.flush().await +} + +async fn send_websocket_jsonrpc_message( + websocket_writer: &mut W, + connection_label: &str, + message: &JSONRPCMessage, +) -> Result<(), String> +where + W: Sink + Unpin, + M: JsonRpcWebSocketMessage, + E: std::fmt::Display, +{ + match serialize_jsonrpc_message(message) { + Ok(encoded) => websocket_writer + .send(M::from_text(encoded)) + .await + .map_err(|err| { + format!("failed to write websocket JSON-RPC message to {connection_label}: {err}") + }), + Err(err) => Err(format!( + "failed to serialize JSON-RPC message for {connection_label}: {err}" + )), + } +} + +fn serialize_jsonrpc_message(message: &JSONRPCMessage) -> Result { + serde_json::to_string(message) +} + +#[cfg(test)] +mod tests { + use std::pin::Pin; + use std::sync::Arc; + use std::sync::atomic::AtomicBool; + use std::sync::atomic::Ordering; + use std::task::Context; + use std::task::Poll; + + use codex_exec_server_protocol::JSONRPCRequest; + use codex_exec_server_protocol::RequestId; + use futures::channel::mpsc as futures_mpsc; + use futures::task::AtomicWaker; + use pretty_assertions::assert_eq; + use tokio::net::TcpListener; + use tokio::time::timeout; + use tokio_tungstenite::accept_async; + use tokio_tungstenite::connect_async; + + use super::*; + + #[tokio::test] + async fn stdio_connection_accepts_message_at_size_limit() -> anyhow::Result<()> { + let message = test_jsonrpc_message(); + let encoded = serde_json::to_string(&message)?; + let max_message_len = encoded.len(); + + for line_ending in [b"\n".as_slice(), b"\r\n".as_slice()] { + let (reader, mut peer) = + tokio::io::duplex(max_message_len.saturating_add(line_ending.len())); + let mut connection = JsonRpcConnection::from_stdio_with_max_message_len( + reader, + tokio::io::sink(), + "test stdio peer".to_string(), + max_message_len, + ); + + peer.write_all(encoded.as_bytes()).await?; + peer.write_all(line_ending).await?; + let event = timeout(Duration::from_secs(1), connection.incoming_rx.recv()) + .await? + .expect("stdio connection should report the message"); + match event { + JsonRpcConnectionEvent::QueuedRequest { request, .. } => { + assert_eq!(JSONRPCMessage::Request(request), message) + } + event => anyhow::bail!("expected JSON-RPC message, got {event:?}"), + } + + drop(peer); + drop(connection); + } + + Ok(()) + } + + #[tokio::test] + async fn stdio_connection_rejects_overlong_unterminated_message() -> anyhow::Result<()> { + let max_message_len: usize = 32; + let (reader, mut peer) = tokio::io::duplex(max_message_len.saturating_add(1)); + let mut connection = JsonRpcConnection::from_stdio_with_max_message_len( + reader, + tokio::io::sink(), + "hostile stdio peer".to_string(), + max_message_len, + ); + let overlong_message = vec![b'x'; max_message_len + 1]; + + peer.write_all(&overlong_message).await?; + let event = timeout(Duration::from_secs(1), connection.incoming_rx.recv()) + .await? + .expect("stdio connection should report the framing violation"); + match event { + JsonRpcConnectionEvent::Disconnected { reason } => assert_eq!( + reason, + Some( + "JSON-RPC message from hostile stdio peer exceeds maximum length of 32 bytes" + .to_string() + ) + ), + event => anyhow::bail!("expected stdio disconnect, got {event:?}"), + } + + drop(peer); + drop(connection); + Ok(()) + } + + #[tokio::test] + async fn websocket_connection_sends_configured_ping() -> anyhow::Result<()> { + let (client_websocket, mut server_websocket) = websocket_pair().await?; + let connection = JsonRpcConnection::from_websocket_stream( + client_websocket, + "test".into(), + Some(WEBSOCKET_KEEPALIVE_INTERVAL), + ); + + let message = timeout(Duration::from_secs(1), server_websocket.next()) + .await? + .expect("websocket should stay open")?; + assert!(matches!(message, Message::Ping(_))); + + drop(connection); + Ok(()) + } + + #[tokio::test] + async fn websocket_connection_ignores_server_pong() -> anyhow::Result<()> { + let (client_websocket, mut server_websocket) = websocket_pair().await?; + let mut connection = JsonRpcConnection::from_websocket(client_websocket, "test".into()); + + server_websocket + .send(Message::Pong(b"check".to_vec().into())) + .await?; + assert!( + timeout(Duration::from_millis(50), connection.incoming_rx.recv()) + .await + .is_err() + ); + + drop(connection); + Ok(()) + } + + #[tokio::test] + async fn websocket_connection_reports_server_close() -> anyhow::Result<()> { + let (client_websocket, mut server_websocket) = websocket_pair().await?; + let mut connection = JsonRpcConnection::from_websocket(client_websocket, "test".into()); + + server_websocket.close(None).await?; + assert!(matches!( + timeout(Duration::from_secs(1), connection.incoming_rx.recv()).await?, + Some(JsonRpcConnectionEvent::Disconnected { reason: None }) + )); + + drop(connection); + Ok(()) + } + + #[tokio::test] + async fn websocket_connection_accepts_binary_jsonrpc_message() -> anyhow::Result<()> { + let (client_websocket, mut server_websocket) = websocket_pair().await?; + let mut connection = JsonRpcConnection::from_websocket(client_websocket, "test".into()); + let message = JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(1), + method: "test".to_string(), + params: None, + trace: None, + }); + + server_websocket + .send(Message::Binary(serde_json::to_vec(&message)?.into())) + .await?; + let Some(JsonRpcConnectionEvent::QueuedRequest { request, .. }) = + timeout(Duration::from_secs(1), connection.incoming_rx.recv()).await? + else { + anyhow::bail!("expected a queued JSON-RPC request"); + }; + assert_eq!(JSONRPCMessage::Request(request), message); + + drop(connection); + Ok(()) + } + + #[tokio::test] + async fn websocket_connection_keeps_outbound_message_while_send_is_backpressured() + -> anyhow::Result<()> { + let (websocket, control, mut outbound_rx) = + ControlledWebSocket::new(/*write_ready*/ false); + let mut connection = JsonRpcConnection::from_websocket_stream( + websocket, + "test".into(), + /*ping_interval*/ None, + ); + let message = test_jsonrpc_message(); + + connection.outgoing_tx.send(message.clone()).await?; + control.wait_for_blocked_write().await?; + control.send_inbound(Message::Pong(b"check".to_vec().into()))?; + assert!( + timeout(Duration::from_millis(50), connection.incoming_rx.recv()) + .await + .is_err() + ); + + control.set_write_ready(); + assert!(matches!( + timeout(Duration::from_secs(1), outbound_rx.next()).await?, + Some(Message::Text(text)) if serde_json::from_str::(&text)? == message + )); + drop(connection); + Ok(()) + } + + async fn websocket_pair() -> anyhow::Result<( + WebSocketStream>, + WebSocketStream, + )> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let server_task = tokio::spawn(async move { + let (stream, _) = listener.accept().await?; + accept_async(stream).await.map_err(anyhow::Error::from) + }); + let (client_websocket, _) = connect_async(websocket_url).await?; + let server_websocket = server_task.await??; + Ok((client_websocket, server_websocket)) + } + + fn test_jsonrpc_message() -> JSONRPCMessage { + JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(1), + method: "test".to_string(), + params: None, + trace: None, + }) + } + + struct ControlledWebSocket { + inbound_rx: futures_mpsc::UnboundedReceiver>, + outbound_tx: futures_mpsc::UnboundedSender, + write_ready: Arc, + write_blocked: Arc, + write_blocked_waker: Arc, + write_waker: Arc, + } + + struct ControlledWebSocketHandle { + inbound_tx: futures_mpsc::UnboundedSender>, + write_ready: Arc, + write_blocked: Arc, + write_blocked_waker: Arc, + write_waker: Arc, + } + + impl ControlledWebSocket { + fn new( + write_ready: bool, + ) -> ( + Self, + ControlledWebSocketHandle, + futures_mpsc::UnboundedReceiver, + ) { + let (inbound_tx, inbound_rx) = futures_mpsc::unbounded(); + let (outbound_tx, outbound_rx) = futures_mpsc::unbounded(); + let write_ready = Arc::new(AtomicBool::new(write_ready)); + let write_blocked = Arc::new(AtomicBool::new(false)); + let write_blocked_waker = Arc::new(AtomicWaker::new()); + let write_waker = Arc::new(AtomicWaker::new()); + ( + Self { + inbound_rx, + outbound_tx, + write_ready: Arc::clone(&write_ready), + write_blocked: Arc::clone(&write_blocked), + write_blocked_waker: Arc::clone(&write_blocked_waker), + write_waker: Arc::clone(&write_waker), + }, + ControlledWebSocketHandle { + inbound_tx, + write_ready, + write_blocked, + write_blocked_waker, + write_waker, + }, + outbound_rx, + ) + } + } + + impl ControlledWebSocketHandle { + fn send_inbound(&self, message: Message) -> anyhow::Result<()> { + self.inbound_tx + .unbounded_send(Ok(message)) + .map_err(anyhow::Error::from) + } + + fn set_write_ready(&self) { + self.write_ready.store(true, Ordering::Release); + self.write_waker.wake(); + } + + async fn wait_for_blocked_write(&self) -> anyhow::Result<()> { + timeout( + Duration::from_secs(1), + futures::future::poll_fn(|cx| { + if self.write_blocked.load(Ordering::Acquire) { + Poll::Ready(()) + } else { + self.write_blocked_waker.register(cx.waker()); + Poll::Pending + } + }), + ) + .await?; + Ok(()) + } + } + + impl Sink for ControlledWebSocket { + type Error = std::convert::Infallible; + + fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + if self.write_ready.load(Ordering::Acquire) { + Poll::Ready(Ok(())) + } else { + self.write_blocked.store(true, Ordering::Release); + self.write_blocked_waker.wake(); + self.write_waker.register(cx.waker()); + Poll::Pending + } + } + + fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> { + self.outbound_tx + .unbounded_send(item) + .expect("test outbound receiver should stay open"); + Ok(()) + } + + fn poll_flush( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } + } + + impl Stream for ControlledWebSocket { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.inbound_rx).poll_next(cx) + } + } +} diff --git a/codex-rs/exec-server/src/environment.rs b/codex-rs/exec-server/src/environment.rs new file mode 100644 index 0000000000000000000000000000000000000000..26b8825d727f5aa73b63fe50c5866b52df7748f2 --- /dev/null +++ b/codex-rs/exec-server/src/environment.rs @@ -0,0 +1,1906 @@ +mod connect_options; + +use std::collections::HashMap; +use std::collections::HashSet; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::RwLock; + +use arc_swap::ArcSwapOption; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_protocol::capabilities::CapabilityRootLocation; +use codex_protocol::capabilities::SelectedCapabilityRoot; +use codex_protocol::shell_environment::CODEX_EXEC_SERVER_NOISE_AUTH_TOKEN_ENV_VAR; + +use crate::CapabilityRootsDiscoverParams; +use crate::CapabilityRootsDiscoverResponse; +use crate::EnvironmentConfigReadParams; +use crate::EnvironmentConfigReadResponse; +use crate::ExecServerError; +use crate::ExecServerRuntimePaths; +use crate::ExecutorFileSystem; +use crate::HttpClient; +use crate::NoiseChannelIdentity; +use crate::NoiseRendezvousConnectProvider; +use crate::client::LazyRemoteExecServerClient; +use crate::client::http_client::RouteAwareHttpClient; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::client_api::ExecServerTransportParams; +use crate::environment_bootstrap::PreparedEnvironmentManager; +use crate::environment_bootstrap::PreparedEnvironmentSource; +use crate::environment_config::read_environment_config; +use crate::environment_provider::DefaultEnvironmentProvider; +use crate::environment_provider::EnvironmentDefault; +use crate::environment_provider::EnvironmentProvider; +use crate::environment_provider::EnvironmentProviderSnapshot; +use crate::environment_provider::normalize_exec_server_url; +use crate::environment_toml::environment_provider_from_codex_home; +use crate::local_file_system::LocalFileSystem; +use crate::local_process::LocalProcess; +use crate::process::ExecBackend; +use crate::protocol::EnvironmentInfo; +use crate::remote::NoiseRendezvousEnvironmentConfig; +use crate::remote_file_system::RemoteFileSystem; +use crate::remote_process::RemoteProcess; +use tokio::sync::watch; +use tokio_util::task::AbortOnDropHandle; +use tracing::Instrument; +use tracing::instrument::WithSubscriber; + +#[path = "environment/accepted.rs"] +mod accepted; + +pub use connect_options::RemoteEnvironmentOptions; + +pub const CODEX_EXEC_SERVER_URL_ENV_VAR: &str = "CODEX_EXEC_SERVER_URL"; +pub const CODEX_EXEC_SERVER_NOISE_REGISTRY_URL_ENV_VAR: &str = + "CODEX_EXEC_SERVER_NOISE_REGISTRY_URL"; +pub const CODEX_EXEC_SERVER_NOISE_ENVIRONMENT_ID_ENV_VAR: &str = + "CODEX_EXEC_SERVER_NOISE_ENVIRONMENT_ID"; +pub const CODEX_EXEC_SERVER_NOISE_CHATGPT_ACCOUNT_ID_ENV_VAR: &str = + "CODEX_EXEC_SERVER_NOISE_CHATGPT_ACCOUNT_ID"; + +/// The current connection state for one concrete environment. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum EnvironmentConnectionState { + /// An initialized exec-server connection is available. + Connected, + /// No initialized exec-server connection is currently available. + Disconnected, +} + +/// Owns the execution/filesystem environments available to the Codex runtime. +/// +/// `EnvironmentManager` is a shared registry for concrete environments. Its +/// default constructor preserves the legacy `CODEX_EXEC_SERVER_URL` behavior +/// while configured construction accepts a provider-supplied snapshot. +/// +/// Setting `CODEX_EXEC_SERVER_URL=none` disables environment access by leaving +/// the default environment unset and omitting the local environment. Callers +/// use `default_environment().is_some()` as the signal for model-facing +/// shell/filesystem tool availability. +/// +/// Ordinary remote environments begin connecting when added to the manager. +/// Provisioned remote environments connect only after they are selected for use; +/// their deferred transport waits for provisioning to complete first. Filesystem +/// and execution backends share the resulting startup and reconnect as needed. +#[derive(Debug)] +pub struct EnvironmentManager { + default_environment: Option, + pub(super) environments: RwLock>>, + local_environment: Option>, + local_runtime_paths: Option, + http_client_factory: HttpClientFactory, +} + +/// Information supplied by the environment owner when an environment is ready. +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct EnvironmentReadyInfo { + /// Ordered capability roots selected for this environment. + pub selected_capability_roots: Vec, +} + +/// Maximum capability roots accepted from environment ready information. +pub const MAX_SELECTED_CAPABILITY_ROOTS: usize = 256; + +pub const LOCAL_ENVIRONMENT_ID: &str = "local"; +pub const REMOTE_ENVIRONMENT_ID: &str = "remote"; + +/// Non-mutating connection status observed by an environment owner. +/// +/// Computing this status never starts, waits for, or reconnects an exec-server +/// transport. Already-ready remote environments may receive a fail-fast probe +/// over their existing connection. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum EnvironmentObservedStatus { + /// A local environment, or a remote environment whose existing connection answered a probe. + Ready, + /// The configured environment has no ready connection and no observed connection failure. + /// + /// This includes lazy transports that have never been started and initial startup that has + /// not finished. Computing status does not start the environment or wait for startup. + Pending, + /// A connection attempt, prior connection, or fail-fast status probe observed a failure. + /// + /// This does not promise that the failure is terminal: later normal environment use may + /// recover the connection. Computing status itself does not trigger recovery. + Disconnected { + /// Human-readable reason recorded by the failed connection attempt or probe. + error: String, + }, +} + +impl EnvironmentManager { + /// Builds a test-only manager without configured sandbox helper paths. + pub fn default_for_tests() -> Self { + Self { + default_environment: Some(LOCAL_ENVIRONMENT_ID.to_string()), + environments: RwLock::new(HashMap::from([( + LOCAL_ENVIRONMENT_ID.to_string(), + Arc::new(Environment::default_for_tests()), + )])), + local_environment: Some(Arc::new(Environment::default_for_tests())), + local_runtime_paths: None, + // Test-only construction has no application config from which to resolve proxy policy. + http_client_factory: HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + } + } + + /// Builds a manager with no configured execution environments. + pub fn without_environments(http_client_factory: HttpClientFactory) -> Self { + Self { + default_environment: None, + environments: RwLock::new(HashMap::new()), + local_environment: None, + local_runtime_paths: None, + http_client_factory, + } + } + + /// Builds a test-only manager from a raw exec-server URL value. + pub async fn create_for_tests( + exec_server_url: Option, + local_runtime_paths: Option, + ) -> Self { + let provider = DefaultEnvironmentProvider::new(exec_server_url); + match Self::from_snapshot( + provider.snapshot_inner(), + local_runtime_paths, + // Test-only construction has no application config from which to resolve proxy policy. + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) { + Ok(manager) => manager, + Err(err) => panic!("default provider should create valid environments: {err}"), + } + } + + /// Discovers configured environments without starting remote connections. + /// + /// If `CODEX_HOME/environments.toml` is present, it defines the configured + /// environments. Otherwise this preserves the legacy + /// `CODEX_EXEC_SERVER_URL` behavior. + pub async fn prepare_from_codex_home( + codex_home: impl AsRef, + ) -> Result { + let source = if let Some(config) = noise_environment_config_from_env()? { + PreparedEnvironmentSource::Noise(config) + } else { + let provider = environment_provider_from_codex_home(codex_home.as_ref())?; + PreparedEnvironmentSource::Snapshot(provider.snapshot().await?) + }; + Ok(PreparedEnvironmentManager { source }) + } + + /// Builds a manager from `CODEX_HOME` with an explicit outbound HTTP policy. + pub async fn from_codex_home( + codex_home: impl AsRef, + local_runtime_paths: Option, + http_client_factory: HttpClientFactory, + ) -> Result { + Self::prepare_from_codex_home(codex_home) + .await? + .build(local_runtime_paths, http_client_factory) + } + + /// Discovers environment-variable environments without starting connections. + pub async fn prepare_from_env() -> Result { + let source = if let Some(config) = noise_environment_config_from_env()? { + PreparedEnvironmentSource::Noise(config) + } else { + let provider = DefaultEnvironmentProvider::from_env(); + PreparedEnvironmentSource::Snapshot(provider.snapshot().await?) + }; + Ok(PreparedEnvironmentManager { source }) + } + + /// Builds a manager from environment variables with an explicit outbound HTTP policy. + pub async fn from_env( + local_runtime_paths: Option, + http_client_factory: HttpClientFactory, + ) -> Result { + Self::prepare_from_env() + .await? + .build(local_runtime_paths, http_client_factory) + } + + pub(crate) fn from_noise_environment_config( + config: NoiseRendezvousEnvironmentConfig, + local_runtime_paths: Option, + http_client_factory: HttpClientFactory, + ) -> Result { + let connect_provider = config.into_connect_provider(http_client_factory.clone())?; + let manager = Self { + default_environment: Some(REMOTE_ENVIRONMENT_ID.to_string()), + environments: RwLock::new(HashMap::new()), + local_environment: None, + local_runtime_paths, + http_client_factory, + }; + let identity = noise_channel_identity()?; + let environment = Arc::new(Environment::remote_with_transport( + ExecServerTransportParams::NoiseRendezvous { + provider: connect_provider, + identity, + }, + manager.local_runtime_paths.clone(), + manager.http_client_factory.clone(), + )); + manager.insert_environment(REMOTE_ENVIRONMENT_ID.to_string(), environment)?; + Ok(manager) + } + + /// Builds a test-only manager that keeps the provider default while also + /// allowing tests to select the local environment explicitly. + pub async fn create_for_tests_with_local( + exec_server_url: Option, + local_runtime_paths: ExecServerRuntimePaths, + ) -> Self { + let mut snapshot = DefaultEnvironmentProvider::new(exec_server_url).snapshot_inner(); + snapshot.include_local = true; + match Self::from_snapshot( + snapshot, + Some(local_runtime_paths), + // Test-only construction has no application config from which to resolve proxy policy. + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) { + Ok(manager) => manager, + Err(err) => panic!("test provider with local should create valid environments: {err}"), + } + } + + pub(crate) fn from_snapshot( + snapshot: EnvironmentProviderSnapshot, + local_runtime_paths: Option, + http_client_factory: HttpClientFactory, + ) -> Result { + let EnvironmentProviderSnapshot { + environments, + default, + include_local, + } = snapshot; + let mut environment_map = + HashMap::with_capacity(environments.len() + usize::from(include_local)); + let local_environment = if include_local { + let local_runtime_paths = local_runtime_paths.clone().ok_or_else(|| { + ExecServerError::Protocol( + "local environment requires configured runtime paths".to_string(), + ) + })?; + let local_environment = Arc::new(Environment::local( + local_runtime_paths, + http_client_factory.clone(), + )); + environment_map.insert( + LOCAL_ENVIRONMENT_ID.to_string(), + Arc::clone(&local_environment), + ); + Some(local_environment) + } else { + None + }; + for (id, transport) in environments { + if id.is_empty() { + return Err(ExecServerError::Protocol( + "environment id cannot be empty".to_string(), + )); + } + if id == LOCAL_ENVIRONMENT_ID { + return Err(ExecServerError::Protocol(format!( + "environment id `{LOCAL_ENVIRONMENT_ID}` is reserved for EnvironmentManager" + ))); + } + let environment = Environment::remote_with_transport( + transport, + /*local_runtime_paths*/ None, + http_client_factory.clone(), + ); + if environment_map + .insert(id.clone(), Arc::new(environment)) + .is_some() + { + return Err(ExecServerError::Protocol(format!( + "environment id `{id}` is duplicated" + ))); + } + } + let default_environment = match default { + EnvironmentDefault::Disabled => None, + EnvironmentDefault::EnvironmentId(environment_id) => { + if !environment_map.contains_key(&environment_id) { + return Err(ExecServerError::Protocol(format!( + "default environment `{environment_id}` is not configured" + ))); + } + Some(environment_id) + } + }; + // The snapshot is valid; start connecting its remote environments in the background. + for environment in environment_map.values() { + environment.start_connecting(); + } + Ok(Self { + default_environment, + environments: RwLock::new(environment_map), + local_environment, + local_runtime_paths, + http_client_factory, + }) + } + + /// Returns the default environment instance. + pub fn default_environment(&self) -> Option> { + self.default_environment + .as_deref() + .and_then(|environment_id| self.get_environment(environment_id)) + } + + /// Returns the id of the default environment. + pub fn default_environment_id(&self) -> Option<&str> { + self.default_environment.as_deref() + } + + /// Returns the ordered environment ids used for new thread startup. + pub fn default_environment_ids(&self) -> Vec { + let Some(default_environment_id) = self.default_environment.as_ref() else { + return Vec::new(); + }; + let environments = self + .environments + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let mut environment_ids = Vec::with_capacity(environments.len()); + environment_ids.push(default_environment_id.clone()); + environment_ids.extend( + environments + .keys() + .filter(|environment_id| *environment_id != default_environment_id) + .cloned(), + ); + environment_ids + } + + /// Returns the local environment instance when one is configured. + pub fn try_local_environment(&self) -> Option> { + self.local_environment.as_ref().map(Arc::clone) + } + + /// Returns a named environment instance. + pub fn get_environment(&self, environment_id: &str) -> Option> { + self.environments + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .get(environment_id) + .cloned() + } + + /// Records a Ready or Failed provisioning result for an environment. + /// + /// Ordinary environments are ignored. A provisioned environment keeps the same `Arc` from + /// Pending through Ready or Failed, and is created if the report arrives first. + /// + /// Ready updates capability roots and can recover a failed provisioning attempt. Failed keeps + /// the first error until a Ready report arrives; a late failure cannot replace Ready. Invalid + /// Ready information fails an existing Pending environment but does not create a missing one. + /// + /// This only updates provisioning. The connection starts when the environment is selected. + pub fn report_environment_provisioning_status( + &self, + environment_id: String, + readiness: Result, + provider_if_missing: Arc, + ) -> Result>, ExecServerError> { + validate_environment_id(&environment_id)?; + let mut environments = self + .environments + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(environment) = environments.get(&environment_id).cloned() { + if environment.provisioning_status_tx.is_none() { + return Ok(None); + } + match readiness { + Ok(ready_info) => { + environment.apply_ready_report(&environment_id, ready_info)?; + } + Err(error) => { + environment.apply_error_report(&environment_id, error)?; + } + } + return Ok(Some(environment)); + } + + let environment = match readiness { + Ok(ready_info) => { + validate_environment_ready_info(&environment_id, &ready_info)?; + let environment = Arc::new( + self.provisioning_noise_environment(provider_if_missing, Some(Ok(())))?, + ); + environment.ready_info.store(Some(Arc::new(ready_info))); + environment + } + Err(error) => Arc::new( + self.provisioning_noise_environment(provider_if_missing, Some(Err(error)))?, + ), + }; + environments.insert(environment_id, Arc::clone(&environment)); + Ok(Some(environment)) + } + + /// Returns the outbound HTTP policy carried by this manager. + pub fn http_client_factory(&self) -> &HttpClientFactory { + &self.http_client_factory + } + + /// Returns the current status of one named environment when it is configured. + pub async fn get_environment_status( + &self, + environment_id: &str, + ) -> Option { + let environment = self.get_environment(environment_id)?; + Some(environment.status().await) + } + + /// Adds or replaces a named remote environment without changing the + /// manager's default environment selection. Uses the default WebSocket + /// connection timeout when none is provided. + pub fn upsert_environment( + &self, + environment_id: String, + exec_server_url: String, + connect_timeout: Option, + ) -> Result<(), ExecServerError> { + self.upsert_environment_with_options( + environment_id, + RemoteEnvironmentOptions { + exec_server_url, + connect_timeout, + http_headers: HashMap::new(), + }, + ) + } + + /// Adds or replaces a direct environment with trusted host-owned connection options. + /// + /// Invalid headers and WebSocket-controlled handshake headers are rejected + /// before the environment is registered. Valid headers are retained for + /// automatic reconnects without changing existing URL-only environment APIs. + pub fn upsert_environment_with_options( + &self, + environment_id: String, + options: RemoteEnvironmentOptions, + ) -> Result<(), ExecServerError> { + validate_environment_id(&environment_id)?; + let transport = options.into_transport_params()?; + let environment = Arc::new(Environment::remote_with_transport( + transport, + self.local_runtime_paths.clone(), + self.http_client_factory.clone(), + )); + self.insert_environment(environment_id, environment) + } + + /// Returns the stable environment for an ID, creating it as pending when absent. + pub fn materialize_pending_noise_environment( + &self, + environment_id: String, + provider: Arc, + ) -> Result, ExecServerError> { + validate_environment_id(&environment_id)?; + let mut environments = self + .environments + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(environment) = environments.get(&environment_id) { + if environment.provisioning_status_tx.is_none() { + return Err(ExecServerError::ProvisioningModeConflict { environment_id }); + } + return Ok(Arc::clone(environment)); + } + + let environment = + Arc::new(self.provisioning_noise_environment(provider, /*initial_result*/ None)?); + environments.insert(environment_id, Arc::clone(&environment)); + Ok(environment) + } + + fn provisioning_noise_environment( + &self, + provider: Arc, + initial_result: Option>, + ) -> Result { + let identity = noise_channel_identity()?; + let (provisioning_status_tx, provisioning_status_rx) = watch::channel(initial_result); + let mut environment = Environment::remote_with_transport( + ExecServerTransportParams::Deferred(Box::new(crate::client_api::Deferred { + readiness: provisioning_status_rx, + transport: ExecServerTransportParams::NoiseRendezvous { provider, identity }, + })), + self.local_runtime_paths.clone(), + self.http_client_factory.clone(), + ); + environment.provisioning_status_tx = Some(provisioning_status_tx); + Ok(environment) + } + + fn insert_environment( + &self, + environment_id: String, + environment: Arc, + ) -> Result<(), ExecServerError> { + let replaced = { + let mut environments = self + .environments + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + environments.insert(environment_id, Arc::clone(&environment)) + }; + drop(replaced); + environment.start_connecting(); + Ok(()) + } +} + +fn validate_environment_ready_info( + environment_id: &str, + ready_info: &EnvironmentReadyInfo, +) -> Result<(), ExecServerError> { + if ready_info.selected_capability_roots.len() > MAX_SELECTED_CAPABILITY_ROOTS { + return Err(ExecServerError::Protocol(format!( + "environment ready info contains more than {MAX_SELECTED_CAPABILITY_ROOTS} selected capability roots" + ))); + } + + let mut root_ids = HashSet::with_capacity(ready_info.selected_capability_roots.len()); + for root in &ready_info.selected_capability_roots { + let CapabilityRootLocation::Environment { + environment_id: root_environment_id, + .. + } = &root.location; + if root.id.trim().is_empty() + || root_environment_id != environment_id + || !root_ids.insert(root.id.as_str()) + { + return Err(ExecServerError::Protocol(format!( + "selected capability roots must have unique non-empty IDs and belong to environment `{environment_id}`" + ))); + } + } + + Ok(()) +} + +fn noise_channel_identity() -> Result { + NoiseChannelIdentity::generate().map_err(|error| { + ExecServerError::Protocol(format!( + "failed to generate Noise harness identity: {error}" + )) + }) +} + +fn validate_environment_id(environment_id: &str) -> Result<(), ExecServerError> { + if environment_id.is_empty() { + return Err(ExecServerError::Protocol( + "environment id cannot be empty".to_string(), + )); + } + if environment_id == LOCAL_ENVIRONMENT_ID { + return Err(ExecServerError::Protocol(format!( + "environment id `{LOCAL_ENVIRONMENT_ID}` is reserved for EnvironmentManager" + ))); + } + Ok(()) +} + +fn noise_environment_config_from_env() +-> Result, ExecServerError> { + noise_environment_config_from_values( + optional_environment_value(CODEX_EXEC_SERVER_NOISE_REGISTRY_URL_ENV_VAR), + optional_environment_value(CODEX_EXEC_SERVER_NOISE_ENVIRONMENT_ID_ENV_VAR), + optional_environment_value(CODEX_EXEC_SERVER_NOISE_AUTH_TOKEN_ENV_VAR), + optional_environment_value(CODEX_EXEC_SERVER_NOISE_CHATGPT_ACCOUNT_ID_ENV_VAR), + ) +} + +fn noise_environment_config_from_values( + registry_url: Option, + environment_id: Option, + auth_token: Option, + chatgpt_account_id: Option, +) -> Result, ExecServerError> { + let (registry_url, environment_id, auth_token) = + match (registry_url, environment_id, auth_token) { + (None, None, None) => return Ok(None), + (Some(registry_url), Some(environment_id), Some(auth_token)) => { + (registry_url, environment_id, auth_token) + } + _ => { + return Err(ExecServerError::EnvironmentRegistryConfig(format!( + "Noise environment requires {CODEX_EXEC_SERVER_NOISE_REGISTRY_URL_ENV_VAR}, \ +{CODEX_EXEC_SERVER_NOISE_ENVIRONMENT_ID_ENV_VAR}, and \ +{CODEX_EXEC_SERVER_NOISE_AUTH_TOKEN_ENV_VAR}" + ))); + } + }; + + NoiseRendezvousEnvironmentConfig::new( + registry_url, + environment_id, + auth_token, + chatgpt_account_id, + ) + .map(Some) +} + +fn optional_environment_value(name: &str) -> Option { + std::env::var(name) + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) +} + +/// Concrete execution/filesystem environment selected for a session. +/// +/// This bundles the selected backend metadata together with the local runtime +/// paths used by filesystem helpers. +#[derive(Clone)] +pub struct Environment { + remote_client: Option, + ready_info: Arc>, + // No sender means an ordinary environment. A provisioned environment retains a sender whose + // value is None while Pending, Some(Ok(())) when Ready, or Some(Err(error)) when Failed. + provisioning_status_tx: Option>>>, + // Dropping the environment stops unfinished background startup work. + startup_task: Arc>>>, + exec_backend: Arc, + filesystem: Arc, + http_client: Arc, + local_runtime_paths: Option, +} + +impl Environment { + /// Builds a test-only local environment without configured sandbox helper paths. + pub fn default_for_tests() -> Self { + Self { + remote_client: None, + ready_info: Arc::new(ArcSwapOption::empty()), + provisioning_status_tx: None, + startup_task: Arc::new(Mutex::new(None)), + exec_backend: Arc::new(LocalProcess::default()), + filesystem: Arc::new(LocalFileSystem::unsandboxed()), + // Test-only construction has no application config from which to resolve proxy policy. + http_client: Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + local_runtime_paths: None, + } + } +} + +impl std::fmt::Debug for Environment { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Environment").finish_non_exhaustive() + } +} + +impl Environment { + /// Builds an environment using the caller's effective outbound HTTP policy. + pub fn create( + exec_server_url: Option, + local_runtime_paths: ExecServerRuntimePaths, + http_client_factory: HttpClientFactory, + ) -> Result { + Self::create_inner( + exec_server_url, + Some(local_runtime_paths), + http_client_factory, + ) + } + + /// Builds a test-only environment without configured sandbox helper paths. + pub fn create_for_tests(exec_server_url: Option) -> Result { + Self::create_inner( + exec_server_url, + /*local_runtime_paths*/ None, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + } + + /// Builds an environment from the raw `CODEX_EXEC_SERVER_URL` value and + /// local runtime paths used when creating local filesystem helpers. + fn create_inner( + exec_server_url: Option, + local_runtime_paths: Option, + http_client_factory: HttpClientFactory, + ) -> Result { + let (exec_server_url, disabled) = normalize_exec_server_url(exec_server_url); + if disabled { + return Err(ExecServerError::Protocol( + "disabled mode does not create an Environment".to_string(), + )); + } + + Ok(match exec_server_url { + Some(exec_server_url) => Self::remote_with_transport( + ExecServerTransportParams::websocket_url( + exec_server_url, + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + ), + local_runtime_paths, + http_client_factory, + ), + None => match local_runtime_paths { + Some(local_runtime_paths) => Self::local(local_runtime_paths, http_client_factory), + None => Self::default_for_tests(), + }, + }) + } + + pub(crate) fn local( + local_runtime_paths: ExecServerRuntimePaths, + http_client_factory: HttpClientFactory, + ) -> Self { + Self { + remote_client: None, + ready_info: Arc::new(ArcSwapOption::empty()), + provisioning_status_tx: None, + startup_task: Arc::new(Mutex::new(None)), + exec_backend: Arc::new(LocalProcess::with_local_runtime_paths( + local_runtime_paths.clone(), + )), + filesystem: Arc::new(LocalFileSystem::with_runtime_paths( + local_runtime_paths.clone(), + )), + http_client: Arc::new(RouteAwareHttpClient::new(http_client_factory)), + local_runtime_paths: Some(local_runtime_paths), + } + } + + pub(crate) fn remote_with_transport( + remote_transport: ExecServerTransportParams, + local_runtime_paths: Option, + http_client_factory: HttpClientFactory, + ) -> Self { + let client = LazyRemoteExecServerClient::new(remote_transport, http_client_factory); + Self::remote_with_client(client, local_runtime_paths) + } + + pub(crate) fn remote_with_client( + client: LazyRemoteExecServerClient, + local_runtime_paths: Option, + ) -> Self { + let exec_backend: Arc = Arc::new(RemoteProcess::new(client.clone())); + let filesystem: Arc = + Arc::new(RemoteFileSystem::new(client.clone())); + + Self { + remote_client: Some(client.clone()), + ready_info: Arc::new(ArcSwapOption::empty()), + provisioning_status_tx: None, + startup_task: Arc::new(Mutex::new(None)), + exec_backend, + filesystem, + http_client: Arc::new(client), + local_runtime_paths, + } + } + + pub fn is_remote(&self) -> bool { + self.remote_client.is_some() + } + + fn apply_error_report( + &self, + environment_id: &str, + error: String, + ) -> Result<(), ExecServerError> { + let Some(provisioning_status_tx) = &self.provisioning_status_tx else { + return Ok(()); + }; + let mut transition_error = None; + provisioning_status_tx.send_if_modified(|current| match current.as_ref() { + None => { + *current = Some(Err(error.clone())); + true + } + Some(Ok(())) => { + transition_error = Some(ExecServerError::Protocol(format!( + "environment `{environment_id}` is already ready, but a later provisioning report failed: {error}" + ))); + false + } + Some(Err(_)) => false, + }); + + transition_error.map_or(Ok(()), Err) + } + + fn apply_ready_report( + &self, + environment_id: &str, + ready_info: EnvironmentReadyInfo, + ) -> Result<(), ExecServerError> { + let Some(provisioning_status_tx) = &self.provisioning_status_tx else { + return Ok(()); + }; + let mut transition_error = None; + provisioning_status_tx.send_if_modified(|current| { + if let Err(error) = validate_environment_ready_info(environment_id, &ready_info) { + let pending = current.is_none(); + if pending { + *current = Some(Err(error.to_string())); + } + transition_error = Some(error); + return pending; + } + self.ready_info.store(Some(Arc::new(ready_info.clone()))); + let was_ready = matches!(current, Some(Ok(()))); + *current = Some(Ok(())); + !was_ready + }); + + transition_error.map_or(Ok(()), Err) + } + + /// Returns a snapshot of the last accepted Ready report. + /// + /// `None` means no Ready report has been accepted, including for ordinary environments. + /// A report with no capability roots is distinct from `None`. The snapshot does not change + /// when later reports arrive and does not indicate whether the connection is healthy. + pub fn last_ready_info(&self) -> Option> { + self.ready_info.load_full() + } + + /// Returns the capability roots most recently reported for this environment. + pub fn selected_capability_roots(&self) -> Vec { + self.ready_info + .load() + .as_ref() + .map_or_else(Vec::new, |ready_info| { + ready_info.selected_capability_roots.clone() + }) + } + + /// Subscribes to the current connection state for this remote environment. + pub fn subscribe_connection_state( + &self, + ) -> Option> { + self.remote_client + .as_ref() + .map(LazyRemoteExecServerClient::subscribe_connection_state) + } + + pub fn local_runtime_paths(&self) -> Option<&ExecServerRuntimePaths> { + self.local_runtime_paths.as_ref() + } + + /// Returns environment information from the selected execution/filesystem environment. + /// Remote metadata is cached for the current client's lifetime. + #[tracing::instrument( + name = "exec_server.environment.info", + skip_all, + fields(remote = self.is_remote()) + )] + pub async fn info(&self) -> Result { + match &self.remote_client { + Some(client) => client.environment_info().await, + None => Ok(EnvironmentInfo::local()), + } + } + + /// Refresh the connection to the executor currently registered for this environment. + /// + /// # Caller contract + /// + /// Call after a planned replacement has registered and become available under the + /// same environment ID. This method does not provision or destroy executors, or + /// wait for the registry to identify a particular replacement. It requires a remote + /// Noise registry-backed environment; other environment types return an error. + /// + /// # Session behavior + /// + /// A fresh registry lookup determines whether the current session can be reused. + /// A changed executor key, or a failed or missing session, causes a fresh connection + /// without resuming the old session. Retirement cancels old recovery, fails its + /// outstanding work and process handles, and never replays commands. The environment + /// object and filesystem handle remain usable through the new connection. + /// A matching executor key preserves a session that has not failed, including one + /// that is recovering; the live readiness check rejects a recovering connection. + /// + /// # Completion and errors + /// + /// Success means the selected connection answered a live status RPC, not merely that + /// metadata was cached. Refresh bypasses the old session's recovery deadline, but + /// registry lookup, connection, and status RPC timeouts still apply. If the initial + /// registry lookup fails, refresh leaves the old session untouched; errors after + /// retirement do not restore it. Ordinary disconnect recovery is unchanged unless + /// refresh retires the session. + #[tracing::instrument( + name = "exec_server.environment.refresh_connection", + skip_all, + fields(remote = self.is_remote()) + )] + pub async fn refresh_connection(&self) -> Result<(), ExecServerError> { + let client = self.remote_client.as_ref().ok_or_else(|| { + ExecServerError::Protocol( + "connection refresh requires a remote environment".to_string(), + ) + })?; + client.refresh_connection().await + } + + /// Fetches uncached metadata, connecting or waiting for recovery as needed. + // TODO: Remove after app-server migrates off of force_environment_info. + #[tracing::instrument( + name = "exec_server.environment.force_info", + skip_all, + fields(remote = self.is_remote()) + )] + pub async fn force_info(&self) -> Result { + match &self.remote_client { + Some(client) => client.get().await?.force_environment_info().await, + None => Ok(EnvironmentInfo::local()), + } + } + + /// Reads selected executor-local configuration fields for this environment. + pub async fn read_environment_config( + &self, + params: EnvironmentConfigReadParams, + ) -> Result { + match &self.remote_client { + Some(client) => client.get().await?.read_environment_config(params).await, + None => read_environment_config(self.filesystem.as_ref(), params) + .await + .map_err(|error| ExecServerError::Protocol(error.to_string())), + } + } + + /// Discovers plugin and skill manifests through the environment's high-level discovery API. + pub async fn discover_capability_roots( + &self, + params: CapabilityRootsDiscoverParams, + ) -> Result { + match &self.remote_client { + Some(client) => { + let mut connection_state = client.subscribe_connection_state(); + let client = client.get().await?; + let discover = || async { + if params.roots.iter().any(|root| { + root.sandbox + .as_ref() + .is_some_and(crate::FileSystemSandboxContext::should_run_in_sandbox) + }) && !client + .environment_info() + .await? + .capabilities + .capability_discovery_sandbox + { + return Err(ExecServerError::Protocol( + "exec-server does not support sandboxed capability discovery" + .to_string(), + )); + } + client.discover_capability_roots(params.clone()).await + }; + match discover().await { + Err(error) if crate::client::is_retryable_recovery_error(&error) => { + tracing::warn!(%error, "replaying capability discovery after executor recovery"); + let recovered = + tokio::time::timeout(std::time::Duration::from_secs(8), async { + while self.readiness_result().is_none_or(|result| result.is_err()) { + if connection_state.changed().await.is_err() { + return false; + } + } + true + }) + .await + .unwrap_or(false); + if recovered { + discover().await + } else { + Err(error) + } + } + response => response, + } + } + None => crate::discover_capability_roots(self.filesystem.as_ref(), params) + .await + .map_err(|error| ExecServerError::Protocol(error.to_string())), + } + } + + /// Starts connecting a remote environment without waiting for it. + /// Requires an active Tokio runtime when background startup is supported. + pub fn start_connecting(&self) { + let Some(client) = &self.remote_client else { + return; + }; + let mut startup_task = self + .startup_task + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if startup_task.is_none() { + *startup_task = client.start_connecting(); + } + } + + /// Starts the initial connection after an environment is actually selected for use. + pub(crate) fn start_connecting_for_use(environment: &Arc) { + let Some(client) = &environment.remote_client else { + return; + }; + let mut startup_task = environment + .startup_task + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if startup_task.is_none() { + let client = client.clone(); + *startup_task = Some(AbortOnDropHandle::new(tokio::spawn( + async move { + if let Err(error) = client.wait_until_ready().await { + tracing::debug!(%error, "exec-server environment startup failed"); + } + } + .in_current_span() + .with_current_subscriber(), + ))); + } + } + + /// Returns whether startup has completed, including a first connection made by refresh. + pub fn startup_finished(&self) -> bool { + self.remote_client + .as_ref() + .is_none_or(LazyRemoteExecServerClient::startup_finished) + } + + /// Waits for initial startup, retrying a previous transient failure when possible. + #[tracing::instrument( + name = "exec_server.environment.wait_until_ready", + skip_all, + fields(remote = self.is_remote()) + )] + pub async fn wait_until_ready(&self) -> Result<(), ExecServerError> { + match &self.remote_client { + Some(client) => client.wait_until_ready().await, + None => Ok(()), + } + } + + /// Returns whether the environment can serve a request without waiting or reconnecting. + pub(crate) fn readiness_result(&self) -> Option> { + match &self.remote_client { + Some(client) => client.readiness_result(), + None => Some(Ok(())), + } + } + + /// Returns the environment's status without starting or recovering it. + /// + /// Local environments are always ready. Remote environments with an + /// already-ready cached connection receive a fail-fast `environment/status` + /// probe; other remote states are returned from cached connection state + /// without waiting for startup or recovery. + pub async fn status(&self) -> EnvironmentObservedStatus { + match &self.remote_client { + Some(client) => client.status().await, + None => EnvironmentObservedStatus::Ready, + } + } + + pub fn get_exec_backend(&self) -> Arc { + Arc::clone(&self.exec_backend) + } + + pub fn get_http_client(&self) -> Arc { + Arc::clone(&self.http_client) + } + + pub fn get_filesystem(&self) -> Arc { + Arc::clone(&self.filesystem) + } + + /// Returns a filesystem view that fails instead of starting or waiting for a connection. + pub fn get_filesystem_without_reconnect(&self) -> Arc { + match &self.remote_client { + Some(client) => Arc::new(RemoteFileSystem::new(client.fail_fast())), + None => Arc::clone(&self.filesystem), + } + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + use std::sync::Arc; + use std::time::Duration; + + use super::Environment; + use super::EnvironmentManager; + use super::EnvironmentObservedStatus; + use super::LOCAL_ENVIRONMENT_ID; + use super::REMOTE_ENVIRONMENT_ID; + use super::noise_environment_config_from_values; + use crate::ExecServerRuntimePaths; + use crate::ProcessId; + use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; + use crate::client_api::ExecServerTransportParams; + use crate::client_api::StdioExecServerCommand; + use crate::environment_provider::EnvironmentDefault; + use crate::environment_provider::EnvironmentProviderSnapshot; + use codex_http_client::HttpClientFactory; + use codex_http_client::OutboundProxyPolicy; + use codex_utils_path_uri::PathUri; + use pretty_assertions::assert_eq; + use tokio::net::TcpListener; + use tokio::time::timeout; + + fn legacy_http_client_factory() -> HttpClientFactory { + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault) + } + + fn prepared_websocket_environment() -> ExecServerTransportParams { + ExecServerTransportParams::websocket_url( + "ws://127.0.0.1:8765".to_string(), + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + ) + } + + fn test_runtime_paths() -> ExecServerRuntimePaths { + ExecServerRuntimePaths::new( + std::env::current_exe().expect("current exe"), + /*codex_linux_sandbox_exe*/ None, + ) + .expect("runtime paths") + } + + fn assert_local_environment_unavailable(manager: &EnvironmentManager) { + assert!(manager.try_local_environment().is_none()); + } + + #[test] + fn local_environment_info_includes_current_directory() { + let info = super::EnvironmentInfo::local(); + + assert_eq!( + info.cwd, + Some( + PathUri::from_host_native_path(std::env::current_dir().expect("current directory")) + .expect("cwd URI") + ) + ); + assert_eq!( + info.temp_dir, + PathUri::from_host_native_path(std::env::temp_dir()).ok() + ); + } + + #[tokio::test] + async fn noise_environment_config_selects_remote_as_default() { + let config = noise_environment_config_from_values( + Some("http://registry.example/api".to_string()), + Some("environment-requested".to_string()), + Some("registry-token".to_string()), + Some("workspace-123".to_string()), + ) + .expect("parse noise environment configuration") + .expect("noise environment configuration"); + + let manager = EnvironmentManager::from_noise_environment_config( + config, + /*local_runtime_paths*/ None, + HttpClientFactory::new(OutboundProxyPolicy::RespectSystemProxy), + ) + .expect("build environment manager"); + + assert_eq!( + manager.http_client_factory().outbound_proxy_policy(), + OutboundProxyPolicy::RespectSystemProxy + ); + assert_eq!( + manager.default_environment_id(), + Some(REMOTE_ENVIRONMENT_ID) + ); + assert!( + manager + .default_environment() + .expect("remote environment") + .is_remote() + ); + assert_local_environment_unavailable(&manager); + } + + #[tokio::test] + async fn create_local_environment_does_not_connect() { + let environment = Environment::create( + /*exec_server_url*/ None, + test_runtime_paths(), + legacy_http_client_factory(), + ) + .expect("create environment"); + + assert!(!environment.is_remote()); + assert!(environment.info().await.is_ok()); + } + + #[tokio::test] + async fn environment_manager_normalizes_empty_url() { + let manager = + EnvironmentManager::create_for_tests(Some(String::new()), Some(test_runtime_paths())) + .await; + + let environment = manager.default_environment().expect("default environment"); + assert_eq!(manager.default_environment_id(), Some(LOCAL_ENVIRONMENT_ID)); + assert!(Arc::ptr_eq( + &environment, + &manager + .get_environment(LOCAL_ENVIRONMENT_ID) + .expect("local environment") + )); + assert!(Arc::ptr_eq( + &environment, + &manager.try_local_environment().expect("local environment") + )); + assert!(manager.try_local_environment().is_some()); + assert!(manager.get_environment(REMOTE_ENVIRONMENT_ID).is_none()); + assert!(!environment.is_remote()); + } + + #[tokio::test] + async fn disabled_environment_manager_has_no_default_or_local_environment() { + let manager = EnvironmentManager::without_environments(HttpClientFactory::new( + OutboundProxyPolicy::RespectSystemProxy, + )); + + assert!(manager.default_environment().is_none()); + assert_eq!(manager.default_environment_id(), None); + assert_local_environment_unavailable(&manager); + assert!(manager.get_environment(LOCAL_ENVIRONMENT_ID).is_none()); + assert!(manager.get_environment(REMOTE_ENVIRONMENT_ID).is_none()); + assert_eq!( + manager.http_client_factory().outbound_proxy_policy(), + OutboundProxyPolicy::RespectSystemProxy + ); + } + + #[tokio::test] + async fn environment_manager_creates_remote_environment_for_url() { + let manager = EnvironmentManager::create_for_tests( + Some("ws://127.0.0.1:8765".to_string()), + Some(test_runtime_paths()), + ) + .await; + + let environment = manager.default_environment().expect("default environment"); + assert_eq!( + manager.default_environment_id(), + Some(REMOTE_ENVIRONMENT_ID) + ); + assert!(environment.is_remote()); + assert!(Arc::ptr_eq( + &environment, + &manager + .get_environment(REMOTE_ENVIRONMENT_ID) + .expect("remote environment") + )); + assert!(manager.get_environment(LOCAL_ENVIRONMENT_ID).is_none()); + assert_local_environment_unavailable(&manager); + } + + #[tokio::test] + async fn environment_manager_default_environment_caches_environment() { + let manager = EnvironmentManager::default_for_tests(); + + let first = manager.default_environment().expect("default environment"); + let second = manager.default_environment().expect("default environment"); + + assert!(Arc::ptr_eq(&first, &second)); + assert!(Arc::ptr_eq( + &first.get_filesystem(), + &second.get_filesystem() + )); + } + + #[tokio::test] + async fn environment_manager_builds_from_snapshot() { + let snapshot = EnvironmentProviderSnapshot { + environments: vec![( + REMOTE_ENVIRONMENT_ID.to_string(), + prepared_websocket_environment(), + )], + default: EnvironmentDefault::EnvironmentId(REMOTE_ENVIRONMENT_ID.to_string()), + include_local: false, + }; + let manager = EnvironmentManager::from_snapshot( + snapshot, + Some(test_runtime_paths()), + legacy_http_client_factory(), + ) + .expect("environment manager"); + + assert_eq!( + manager.default_environment_id(), + Some(REMOTE_ENVIRONMENT_ID) + ); + assert!( + manager + .get_environment(REMOTE_ENVIRONMENT_ID) + .expect("remote environment") + .is_remote() + ); + assert!(manager.get_environment(LOCAL_ENVIRONMENT_ID).is_none()); + assert_local_environment_unavailable(&manager); + } + + #[tokio::test] + async fn environment_manager_rejects_empty_environment_id() { + let snapshot = EnvironmentProviderSnapshot { + environments: vec![("".to_string(), prepared_websocket_environment())], + default: EnvironmentDefault::Disabled, + include_local: false, + }; + let err = EnvironmentManager::from_snapshot( + snapshot, + Some(test_runtime_paths()), + legacy_http_client_factory(), + ) + .expect_err("empty id should fail"); + + assert_eq!( + err.to_string(), + "exec-server protocol error: environment id cannot be empty" + ); + } + + #[tokio::test] + async fn environment_manager_rejects_provider_supplied_local_environment() { + let snapshot = EnvironmentProviderSnapshot { + environments: vec![( + LOCAL_ENVIRONMENT_ID.to_string(), + prepared_websocket_environment(), + )], + default: EnvironmentDefault::Disabled, + include_local: false, + }; + let err = EnvironmentManager::from_snapshot( + snapshot, + Some(test_runtime_paths()), + legacy_http_client_factory(), + ) + .expect_err("local id should fail"); + + assert_eq!( + err.to_string(), + "exec-server protocol error: environment id `local` is reserved for EnvironmentManager" + ); + } + + #[tokio::test] + async fn environment_manager_uses_explicit_provider_default() { + let snapshot = EnvironmentProviderSnapshot { + environments: vec![("devbox".to_string(), prepared_websocket_environment())], + default: EnvironmentDefault::EnvironmentId("devbox".to_string()), + include_local: true, + }; + let manager = EnvironmentManager::from_snapshot( + snapshot, + Some(test_runtime_paths()), + legacy_http_client_factory(), + ) + .expect("manager"); + + assert_eq!(manager.default_environment_id(), Some("devbox")); + assert_eq!( + manager.default_environment_ids(), + vec!["devbox".to_string(), LOCAL_ENVIRONMENT_ID.to_string()] + ); + assert!(manager.default_environment().expect("default").is_remote()); + } + + #[tokio::test] + async fn environment_manager_disables_provider_default() { + let snapshot = EnvironmentProviderSnapshot { + environments: vec![("devbox".to_string(), prepared_websocket_environment())], + default: EnvironmentDefault::Disabled, + include_local: true, + }; + let manager = EnvironmentManager::from_snapshot( + snapshot, + Some(test_runtime_paths()), + legacy_http_client_factory(), + ) + .expect("manager"); + + assert_eq!(manager.default_environment_id(), None); + assert!(manager.default_environment().is_none()); + assert!(Arc::ptr_eq( + &manager + .get_environment(LOCAL_ENVIRONMENT_ID) + .expect("local environment"), + &manager.try_local_environment().expect("local environment") + )); + } + + #[tokio::test] + async fn environment_manager_rejects_unknown_provider_default() { + let snapshot = EnvironmentProviderSnapshot { + environments: vec![("devbox".to_string(), prepared_websocket_environment())], + default: EnvironmentDefault::EnvironmentId("missing".to_string()), + include_local: true, + }; + let err = EnvironmentManager::from_snapshot( + snapshot, + Some(test_runtime_paths()), + legacy_http_client_factory(), + ) + .expect_err("unknown default should fail"); + + assert_eq!( + err.to_string(), + "exec-server protocol error: default environment `missing` is not configured" + ); + } + + #[tokio::test] + async fn environment_manager_includes_local_for_default_provider_without_url() { + let manager = EnvironmentManager::create_for_tests( + /*exec_server_url*/ None, + Some(test_runtime_paths()), + ) + .await; + + let environment = manager.default_environment().expect("default environment"); + assert_eq!(manager.default_environment_id(), Some(LOCAL_ENVIRONMENT_ID)); + assert!(Arc::ptr_eq( + &environment, + &manager + .get_environment(LOCAL_ENVIRONMENT_ID) + .expect("local environment") + )); + assert!(Arc::ptr_eq( + &environment, + &manager.try_local_environment().expect("local environment") + )); + assert!(!environment.is_remote()); + } + + #[tokio::test] + async fn environment_manager_carries_local_runtime_paths() { + let runtime_paths = test_runtime_paths(); + let manager = EnvironmentManager::create_for_tests( + /*exec_server_url*/ None, + Some(runtime_paths.clone()), + ) + .await; + + let environment = manager.try_local_environment().expect("local environment"); + + assert_eq!(environment.local_runtime_paths(), Some(&runtime_paths)); + let manager = EnvironmentManager::create_for_tests( + /*exec_server_url*/ None, + Some( + environment + .local_runtime_paths() + .expect("local runtime paths") + .clone(), + ), + ) + .await; + let environment = manager.try_local_environment().expect("local environment"); + assert_eq!(environment.local_runtime_paths(), Some(&runtime_paths)); + } + + #[tokio::test] + async fn environment_manager_omits_default_provider_local_lookup_when_default_disabled() { + let manager = EnvironmentManager::create_for_tests( + Some("none".to_string()), + Some(test_runtime_paths()), + ) + .await; + + assert!(manager.default_environment().is_none()); + assert_eq!(manager.default_environment_id(), None); + assert!(manager.get_environment(LOCAL_ENVIRONMENT_ID).is_none()); + assert!(manager.get_environment(REMOTE_ENVIRONMENT_ID).is_none()); + assert_local_environment_unavailable(&manager); + } + + #[tokio::test] + async fn environment_manager_snapshot_without_local_environment_disables_local_default() { + let mut snapshot = EnvironmentProviderSnapshot { + environments: Vec::new(), + default: EnvironmentDefault::EnvironmentId(LOCAL_ENVIRONMENT_ID.to_string()), + include_local: true, + }; + snapshot.include_local = false; + snapshot.default = EnvironmentDefault::Disabled; + let manager = EnvironmentManager::from_snapshot( + snapshot, + /*local_runtime_paths*/ None, + legacy_http_client_factory(), + ) + .expect("environment manager"); + + assert!(manager.default_environment().is_none()); + assert_eq!(manager.default_environment_id(), None); + assert!(manager.get_environment(LOCAL_ENVIRONMENT_ID).is_none()); + assert_local_environment_unavailable(&manager); + } + + #[tokio::test] + async fn get_environment_returns_none_for_unknown_id() { + let manager = EnvironmentManager::default_for_tests(); + + assert!(manager.get_environment("does-not-exist").is_none()); + } + + #[tokio::test] + async fn environment_manager_upserts_named_remote_environment() { + let manager = EnvironmentManager::without_environments(legacy_http_client_factory()); + + manager + .upsert_environment( + "executor-a".to_string(), + "ws://127.0.0.1:8765".to_string(), + /*connect_timeout*/ None, + ) + .expect("remote environment"); + let first = manager + .get_environment("executor-a") + .expect("first remote environment"); + assert!(first.is_remote()); + assert_eq!(manager.default_environment_id(), None); + + manager + .upsert_environment( + "executor-a".to_string(), + "ws://127.0.0.1:9876".to_string(), + /*connect_timeout*/ None, + ) + .expect("updated remote environment"); + let second = manager + .get_environment("executor-a") + .expect("second remote environment"); + assert!(second.is_remote()); + assert!(!Arc::ptr_eq(&first, &second)); + } + + #[tokio::test] + async fn environment_manager_starts_remote_environment_when_upserted() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind websocket listener"); + let manager = EnvironmentManager::without_environments(legacy_http_client_factory()); + + manager + .upsert_environment( + "executor-a".to_string(), + format!("ws://{}", listener.local_addr().expect("listener address")), + /*connect_timeout*/ None, + ) + .expect("remote environment"); + + timeout(Duration::from_secs(5), listener.accept()) + .await + .expect("environment should start connecting when registered") + .expect("accept connection"); + } + + #[tokio::test] + async fn environment_status_keeps_stdio_environment_pending() { + let environment = Environment::remote_with_transport( + ExecServerTransportParams::StdioCommand { + command: StdioExecServerCommand { + program: "codex-missing-exec-server-for-test".to_string(), + args: Vec::new(), + env: HashMap::new(), + cwd: None, + }, + initialize_timeout: Duration::from_secs(1), + }, + /*local_runtime_paths*/ None, + legacy_http_client_factory(), + ); + + assert_eq!( + environment.status().await, + EnvironmentObservedStatus::Pending + ); + assert!(!environment.startup_finished()); + } + + #[tokio::test] + async fn environment_manager_leaves_stdio_environment_lazy() { + let transport = ExecServerTransportParams::StdioCommand { + command: StdioExecServerCommand { + program: "codex-missing-exec-server-for-test".to_string(), + args: Vec::new(), + env: HashMap::new(), + cwd: None, + }, + initialize_timeout: Duration::from_secs(1), + }; + let manager = EnvironmentManager::from_snapshot( + EnvironmentProviderSnapshot { + environments: vec![("stdio".to_string(), transport)], + default: EnvironmentDefault::Disabled, + include_local: false, + }, + /*local_runtime_paths*/ None, + legacy_http_client_factory(), + ) + .expect("environment manager"); + let environment = manager.get_environment("stdio").expect("stdio environment"); + + assert!(!environment.startup_finished()); + assert!(environment.wait_until_ready().await.is_err()); + assert!(environment.startup_finished()); + } + + #[tokio::test] + async fn selected_capability_inspection_keeps_stdio_environment_lazy() { + use codex_protocol::capabilities::CapabilityRootLocation; + use codex_protocol::capabilities::SelectedCapabilityRoot; + + let transport = ExecServerTransportParams::StdioCommand { + command: StdioExecServerCommand { + program: "codex-missing-exec-server-for-test".to_string(), + args: Vec::new(), + env: HashMap::new(), + cwd: None, + }, + initialize_timeout: Duration::from_secs(1), + }; + let manager = EnvironmentManager::from_snapshot( + EnvironmentProviderSnapshot { + environments: vec![("stdio".to_string(), transport)], + default: EnvironmentDefault::Disabled, + include_local: false, + }, + /*local_runtime_paths*/ None, + legacy_http_client_factory(), + ) + .expect("environment manager"); + let environment = manager.get_environment("stdio").expect("stdio environment"); + let selected_root = SelectedCapabilityRoot { + id: "demo@1".to_string(), + location: CapabilityRootLocation::Environment { + environment_id: "stdio".to_string(), + path: PathUri::parse("file:///plugins/demo").expect("plugin path URI"), + }, + }; + + let status = + manager.inspect_selected_capability_roots(std::slice::from_ref(&selected_root)); + assert!(status.ready_roots.is_empty()); + assert_eq!(status.warnings, Vec::::new()); + assert!(!environment.startup_finished()); + + let missing_root = SelectedCapabilityRoot { + id: "missing@1".to_string(), + location: CapabilityRootLocation::Environment { + environment_id: "missing".to_string(), + path: PathUri::parse("file:///plugins/missing").expect("missing plugin path URI"), + }, + }; + let status = manager.inspect_selected_capability_roots(&[missing_root]); + assert!(status.ready_roots.is_empty()); + assert_eq!( + status.warnings, + vec![ + "selected capability root `missing@1` references unavailable environment `missing`" + .to_string() + ] + ); + + assert!(environment.wait_until_ready().await.is_err()); + + let status = manager.inspect_selected_capability_roots(&[selected_root]); + assert!(status.ready_roots.is_empty()); + assert_eq!(status.warnings.len(), 1); + assert!(status.warnings[0].contains("environment `stdio` is unavailable")); + } + + #[tokio::test] + async fn replacing_environment_stops_its_startup_task() { + let first_listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind first websocket listener"); + let second_listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind second websocket listener"); + let manager = EnvironmentManager::without_environments(legacy_http_client_factory()); + manager + .upsert_environment( + "executor-a".to_string(), + format!( + "ws://{}", + first_listener.local_addr().expect("first listener address") + ), + /*connect_timeout*/ None, + ) + .expect("first remote environment"); + let environment = manager + .get_environment("executor-a") + .expect("first remote environment"); + let startup_abort = environment + .startup_task + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .as_ref() + .expect("startup task") + .abort_handle(); + assert!(!startup_abort.is_finished()); + drop(environment); + + manager + .upsert_environment( + "executor-a".to_string(), + format!( + "ws://{}", + second_listener + .local_addr() + .expect("second listener address") + ), + /*connect_timeout*/ None, + ) + .expect("replacement remote environment"); + + timeout(Duration::from_secs(1), async { + while !startup_abort.is_finished() { + tokio::task::yield_now().await; + } + }) + .await + .expect("replacing the environment should cancel its startup task"); + } + + #[tokio::test] + async fn environment_manager_rejects_empty_remote_environment_url() { + let manager = EnvironmentManager::without_environments(legacy_http_client_factory()); + + let err = manager + .upsert_environment( + "executor-a".to_string(), + String::new(), + /*connect_timeout*/ None, + ) + .expect_err("empty URL should fail"); + + assert_eq!( + err.to_string(), + "exec-server protocol error: remote environment requires an exec-server url" + ); + } + + #[tokio::test] + async fn default_environment_has_ready_local_executor() { + let environment = Environment::default_for_tests(); + + let response = environment + .get_exec_backend() + .start(crate::ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("default-env-proc"), + argv: vec!["true".to_string()], + cwd: PathUri::from_host_native_path( + std::env::current_dir().expect("read current dir"), + ) + .expect("cwd URI"), + shell_snapshot: None, + env_policy: None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await + .expect("start process"); + + assert_eq!(response.process.process_id().as_str(), "default-env-proc"); + } + + #[tokio::test] + async fn local_environment_passes_runtime_paths_to_exec_backend() { + let environment = Environment::local(test_runtime_paths(), legacy_http_client_factory()); + #[cfg(unix)] + let uri = "file://server/share/checkout"; + #[cfg(windows)] + let uri = "file:///usr/local/checkout"; + let sandbox_cwd = PathUri::parse(uri).expect("non-native sandbox cwd URI"); + let source = sandbox_cwd + .to_abs_path() + .expect_err("sandbox cwd should not be native to this host"); + let sandbox = crate::FileSystemSandboxContext::from_permission_profile_with_cwd( + codex_protocol::models::PermissionProfile::workspace_write(), + sandbox_cwd.clone(), + ); + + let result = environment + .get_exec_backend() + .start(crate::ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("local-sandbox-proc"), + argv: vec!["true".to_string()], + cwd: PathUri::from_host_native_path( + std::env::current_dir().expect("read current dir"), + ) + .expect("cwd URI"), + shell_snapshot: None, + env_policy: None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: Some(sandbox), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await; + let Err(err) = result else { + panic!("sandbox cwd should be rejected after resolving runtime paths"); + }; + + assert_eq!( + err.to_string(), + format!( + "exec-server rejected request (-32602): sandbox cwd URI `{sandbox_cwd}` is not valid on this exec-server host: {source}" + ) + ); + } + + #[tokio::test] + async fn test_environment_rejects_sandboxed_filesystem_without_runtime_paths() { + let environment = Environment::default_for_tests(); + let path = codex_utils_absolute_path::AbsolutePathBuf::from_absolute_path( + std::env::current_exe().expect("current exe").as_path(), + ) + .expect("absolute current exe"); + let path = codex_utils_path_uri::PathUri::from_abs_path(&path); + let sandbox = crate::FileSystemSandboxContext::from_permission_profile( + codex_protocol::models::PermissionProfile::from_runtime_permissions( + &codex_protocol::permissions::FileSystemSandboxPolicy::restricted(Vec::new()), + codex_protocol::permissions::NetworkSandboxPolicy::Restricted, + ), + ); + + let err = environment + .get_filesystem() + .read_file(&path, Default::default(), Some(&sandbox)) + .await + .expect_err("sandboxed read should require runtime paths"); + + assert_eq!( + err.to_string(), + "sandboxed filesystem operations require configured runtime paths" + ); + } +} diff --git a/codex-rs/exec-server/src/environment/accepted.rs b/codex-rs/exec-server/src/environment/accepted.rs new file mode 100644 index 0000000000000000000000000000000000000000..600512b0803d060e4bca4cc636841541407f0fa6 --- /dev/null +++ b/codex-rs/exec-server/src/environment/accepted.rs @@ -0,0 +1,64 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::RwLock; + +use super::Environment; +use super::EnvironmentManager; +use super::validate_environment_id; +use crate::ExecServerClient; +use crate::ExecServerClientConnectOptions; +use crate::ExecServerError; +use crate::client::LazyRemoteExecServerClient; +use axum::extract::ws::WebSocket; +use codex_http_client::HttpClientFactory; + +impl EnvironmentManager { + /// Builds a manager around a WebSocket already accepted and authenticated by its host. + /// + /// The manager owns client construction, session initialization, and later + /// recovery. The host only supplies the initial socket and authenticated + /// replacement sockets through [`Self::replace_accepted_websocket`]. + pub async fn from_accepted_websocket( + environment_id: String, + websocket: WebSocket, + options: ExecServerClientConnectOptions, + http_client_factory: HttpClientFactory, + ) -> Result { + validate_environment_id(&environment_id)?; + let client = ExecServerClient::connect_accepted_websocket(websocket, options).await?; + let client = + LazyRemoteExecServerClient::from_connected(client, http_client_factory.clone()); + let environment = Arc::new(Environment::remote_with_client( + client, /*local_runtime_paths*/ None, + )); + Ok(Self { + default_environment: Some(environment_id.clone()), + environments: RwLock::new(HashMap::from([(environment_id, environment)])), + local_environment: None, + local_runtime_paths: None, + http_client_factory, + }) + } + + /// Hands a replacement WebSocket to an existing accepted environment. + /// Returns after handoff; recovery continues asynchronously. + pub async fn replace_accepted_websocket( + &self, + environment_id: &str, + websocket: WebSocket, + ) -> Result<(), ExecServerError> { + let environment = self.get_environment(environment_id).ok_or_else(|| { + ExecServerError::Protocol(format!("environment `{environment_id}` is not configured")) + })?; + environment + .remote_client + .as_ref() + .ok_or_else(|| { + ExecServerError::Protocol( + "local environment does not have a replaceable exec-server client".to_string(), + ) + })? + .replace_accepted_websocket(websocket) + .await + } +} diff --git a/codex-rs/exec-server/src/environment/connect_options.rs b/codex-rs/exec-server/src/environment/connect_options.rs new file mode 100644 index 0000000000000000000000000000000000000000..2fc3a15e6e68f1581ccc837d8ac708660d91d630 --- /dev/null +++ b/codex-rs/exec-server/src/environment/connect_options.rs @@ -0,0 +1,115 @@ +//! Host-owned WebSocket connection options whose trusted headers remain private and redacted. + +use std::collections::HashMap; +use std::time::Duration; + +use http::HeaderMap; +use http::HeaderName; +use http::HeaderValue; + +use crate::ExecServerError; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT; +use crate::client_api::ExecServerTransportParams; +use crate::environment_provider::normalize_exec_server_url; + +/// Host-owned connection settings for a named remote execution environment. +/// +/// Headers are sent only on the direct WebSocket upgrade request and subsequent +/// reconnects. The embedding host must derive them from trusted request or +/// session context; they are not exposed through the app-server protocol. +#[derive(Clone, Eq, PartialEq)] +pub struct RemoteEnvironmentOptions { + /// Direct `ws://` or `wss://` endpoint for the remote environment. + pub exec_server_url: String, + /// Optional connection timeout; the standard exec-server default applies when omitted. + pub connect_timeout: Option, + /// Additional trusted headers for every physical WebSocket connection. + pub http_headers: HashMap, +} + +impl std::fmt::Debug for RemoteEnvironmentOptions { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RemoteEnvironmentOptions") + .field("exec_server_url", &self.exec_server_url) + .field("connect_timeout", &self.connect_timeout) + .field("http_headers", &"") + .finish() + } +} + +impl RemoteEnvironmentOptions { + pub(super) fn into_transport_params( + self, + ) -> Result { + let (exec_server_url, disabled) = normalize_exec_server_url(Some(self.exec_server_url)); + if disabled { + return Err(ExecServerError::Protocol( + "remote environment cannot use disabled exec-server url".to_string(), + )); + } + let exec_server_url = exec_server_url.ok_or_else(|| { + ExecServerError::Protocol("remote environment requires an exec-server url".to_string()) + })?; + + if !self.http_headers.is_empty() { + let url = url::Url::parse(&exec_server_url).map_err(|error| { + ExecServerError::Protocol(format!("invalid exec-server WebSocket URL: {error}")) + })?; + let is_loopback = match url.host() { + Some(url::Host::Domain(host)) => host.eq_ignore_ascii_case("localhost"), + Some(url::Host::Ipv4(address)) => address.is_loopback(), + Some(url::Host::Ipv6(address)) => address.is_loopback(), + None => false, + }; + if url.scheme() != "wss" && !is_loopback { + return Err(ExecServerError::Protocol( + "exec-server WebSocket headers require wss:// or a loopback destination" + .to_string(), + )); + } + } + + let mut http_headers = HeaderMap::with_capacity(self.http_headers.len()); + for (name, value) in self.http_headers { + let header_name = HeaderName::from_bytes(name.as_bytes()).map_err(|_| { + ExecServerError::Protocol(format!( + "invalid exec-server WebSocket header name `{name}`" + )) + })?; + if matches!( + header_name.as_str(), + "connection" | "content-length" | "host" | "transfer-encoding" | "upgrade" + ) || header_name.as_str().starts_with("sec-websocket-") + { + return Err(ExecServerError::Protocol(format!( + "exec-server WebSocket header `{header_name}` is controlled by the connection" + ))); + } + let header_value = HeaderValue::from_str(&value).map_err(|_| { + ExecServerError::Protocol(format!( + "invalid value for exec-server WebSocket header `{header_name}`" + )) + })?; + if http_headers.contains_key(&header_name) { + return Err(ExecServerError::Protocol(format!( + "duplicate exec-server WebSocket header `{header_name}`" + ))); + } + http_headers.insert(header_name, header_value); + } + + Ok(ExecServerTransportParams::WebSocketUrl { + websocket_url: exec_server_url, + connect_timeout: self + .connect_timeout + .unwrap_or(DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT), + initialize_timeout: DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT, + http_headers, + }) + } +} + +#[cfg(test)] +#[path = "connect_options_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/environment/connect_options_tests.rs b/codex-rs/exec-server/src/environment/connect_options_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..c1df5b2054288d671ae3adaea77634713483921e --- /dev/null +++ b/codex-rs/exec-server/src/environment/connect_options_tests.rs @@ -0,0 +1,328 @@ +use std::collections::HashMap; +use std::time::Duration; + +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use futures::SinkExt; +use futures::StreamExt; +use http::HeaderValue; +use pretty_assertions::assert_eq; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::sync::oneshot; +use tokio::time::timeout; +use tokio_tungstenite::WebSocketStream; +use tokio_tungstenite::accept_hdr_async; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::handshake::server::Request; +use tokio_tungstenite::tungstenite::handshake::server::Response; + +use super::RemoteEnvironmentOptions; +use crate::EnvironmentManager; +use crate::InitializeParams; +use crate::InitializeResponse; +use crate::protocol::INITIALIZE_METHOD; +use crate::protocol::INITIALIZED_METHOD; +use crate::protocol::JSONRPCMessage; +use crate::protocol::JSONRPCResponse; + +#[test] +fn remote_environment_options_redact_header_values() { + let options = RemoteEnvironmentOptions { + exec_server_url: "wss://relay.example/environment".to_string(), + connect_timeout: Some(Duration::from_secs(5)), + http_headers: HashMap::from([( + "authorization".to_string(), + "Bearer secret-customer-token".to_string(), + )]), + }; + + let debug = format!("{options:?}"); + assert!(debug.contains("")); + assert!(!debug.contains("secret-customer-token")); +} + +#[test] +fn trusted_headers_require_tls_for_non_loopback_destinations() { + let manager = EnvironmentManager::without_environments(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + )); + + let error = manager + .upsert_environment_with_options( + "customer-environment".to_string(), + RemoteEnvironmentOptions { + exec_server_url: "ws://relay.example/environment".to_string(), + connect_timeout: None, + http_headers: HashMap::from([( + "x-session-id".to_string(), + "customer-session".to_string(), + )]), + }, + ) + .expect_err("trusted headers must not be sent over an insecure remote connection"); + + assert_eq!( + error.to_string(), + "exec-server protocol error: exec-server WebSocket headers require wss:// or a loopback destination" + ); + assert!(manager.get_environment("customer-environment").is_none()); +} + +#[test] +fn duplicate_case_insensitive_websocket_headers_fail_before_registration() { + let manager = EnvironmentManager::without_environments(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + )); + + let error = manager + .upsert_environment_with_options( + "customer-environment".to_string(), + RemoteEnvironmentOptions { + exec_server_url: "ws://127.0.0.1:8765".to_string(), + connect_timeout: None, + http_headers: HashMap::from([ + ("X-Session-Id".to_string(), "first-session".to_string()), + ("x-session-id".to_string(), "second-session".to_string()), + ]), + }, + ) + .expect_err("duplicate header names must fail regardless of case"); + + assert_eq!( + error.to_string(), + "exec-server protocol error: duplicate exec-server WebSocket header `x-session-id`" + ); + assert!(manager.get_environment("customer-environment").is_none()); +} + +#[test] +fn invalid_websocket_header_names_fail_before_registration() { + let manager = EnvironmentManager::without_environments(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + )); + + let error = manager + .upsert_environment_with_options( + "customer-environment".to_string(), + RemoteEnvironmentOptions { + exec_server_url: "ws://127.0.0.1:8765".to_string(), + connect_timeout: None, + http_headers: HashMap::from([("bad header".to_string(), "value".to_string())]), + }, + ) + .expect_err("invalid header name should fail"); + + assert_eq!( + error.to_string(), + "exec-server protocol error: invalid exec-server WebSocket header name `bad header`" + ); + assert!(manager.get_environment("customer-environment").is_none()); +} + +#[test] +fn invalid_websocket_header_values_fail_before_registration() { + let manager = EnvironmentManager::without_environments(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + )); + + let error = manager + .upsert_environment_with_options( + "customer-environment".to_string(), + RemoteEnvironmentOptions { + exec_server_url: "ws://127.0.0.1:8765".to_string(), + connect_timeout: None, + http_headers: HashMap::from([( + "x-session-id".to_string(), + "customer\nspoofed".to_string(), + )]), + }, + ) + .expect_err("invalid header value should fail"); + + assert_eq!( + error.to_string(), + "exec-server protocol error: invalid value for exec-server WebSocket header `x-session-id`" + ); + assert!(manager.get_environment("customer-environment").is_none()); +} + +#[test] +fn websocket_controlled_headers_fail_before_registration() { + let manager = EnvironmentManager::without_environments(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + )); + + for header in ["host", "connection", "upgrade", "sec-websocket-key"] { + let error = manager + .upsert_environment_with_options( + "customer-environment".to_string(), + RemoteEnvironmentOptions { + exec_server_url: "ws://127.0.0.1:8765".to_string(), + connect_timeout: None, + http_headers: HashMap::from([(header.to_string(), "overridden".to_string())]), + }, + ) + .expect_err("connection-controlled header should fail"); + + assert_eq!( + error.to_string(), + format!( + "exec-server protocol error: exec-server WebSocket header `{header}` is controlled by the connection" + ) + ); + assert!(manager.get_environment("customer-environment").is_none()); + } +} + +#[tokio::test] +async fn trusted_headers_are_sent_on_initial_websocket_and_session_reconnect() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!( + "ws://{}", + listener.local_addr().expect("listener should have address") + ); + let (resumed_tx, resumed_rx) = oneshot::channel(); + let (finish_tx, finish_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let mut first = accept_routing_headers_websocket(&listener).await; + complete_websocket_initialize(&mut first, /*expected_resume_session_id*/ None).await; + first + .close(None) + .await + .expect("first websocket should close"); + + let mut resumed = accept_routing_headers_websocket(&listener).await; + complete_websocket_initialize(&mut resumed, Some("session-1")).await; + resumed_tx.send(()).expect("resume should signal"); + finish_rx.await.expect("test should finish"); + }); + + let manager = EnvironmentManager::without_environments(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + )); + manager + .upsert_environment_with_options( + "customer-environment".to_string(), + RemoteEnvironmentOptions { + exec_server_url: websocket_url, + connect_timeout: Some(Duration::from_secs(1)), + http_headers: HashMap::from([ + ( + "x-account-id".to_string(), + "customer-account-456".to_string(), + ), + ( + "x-session-id".to_string(), + "customer-session-123".to_string(), + ), + ( + "x-environment-id".to_string(), + "customer-environment-789".to_string(), + ), + ]), + }, + ) + .expect("environment with routing headers should register"); + + let environment = manager + .get_environment("customer-environment") + .expect("environment with routing headers should exist"); + environment + .wait_until_ready() + .await + .expect("environment with routing headers should initialize"); + timeout(Duration::from_secs(3), resumed_rx) + .await + .expect("routing-header session resume should not time out") + .expect("routing-header session resume should signal"); + + finish_tx.send(()).expect("test should finish"); + server.await.expect("server task should finish"); +} + +async fn accept_routing_headers_websocket(listener: &TcpListener) -> WebSocketStream { + let (stream, _) = listener.accept().await.expect("listener should accept"); + accept_hdr_async(stream, |request: &Request, response: Response| { + assert_eq!( + request.headers().get("x-account-id"), + Some(&HeaderValue::from_static("customer-account-456")) + ); + assert_eq!( + request.headers().get("x-session-id"), + Some(&HeaderValue::from_static("customer-session-123")) + ); + assert_eq!( + request.headers().get("x-environment-id"), + Some(&HeaderValue::from_static("customer-environment-789")) + ); + Ok(response) + }) + .await + .expect("routing-header websocket handshake should succeed") +} + +async fn complete_websocket_initialize( + websocket: &mut WebSocketStream, + expected_resume_session_id: Option<&str>, +) { + let message = websocket + .next() + .await + .expect("initialize request should arrive") + .expect("initialize websocket message should succeed"); + let Message::Text(encoded) = message else { + panic!("expected initialize text message"); + }; + let JSONRPCMessage::Request(request) = + serde_json::from_str::(&encoded).expect("initialize request should parse") + else { + panic!("expected initialize request"); + }; + assert_eq!(request.method, INITIALIZE_METHOD); + let params: InitializeParams = serde_json::from_value( + request + .params + .expect("initialize request should contain parameters"), + ) + .expect("initialize parameters should parse"); + assert_eq!( + params.resume_session_id.as_deref(), + expected_resume_session_id + ); + + let response = JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(InitializeResponse { + session_id: "session-1".to_string(), + environment_info: None, + }) + .expect("initialize response should serialize"), + }); + websocket + .send(Message::Text( + serde_json::to_string(&response) + .expect("initialize response should encode") + .into(), + )) + .await + .expect("initialize response should send"); + + let message = websocket + .next() + .await + .expect("initialized notification should arrive") + .expect("initialized websocket message should succeed"); + let Message::Text(encoded) = message else { + panic!("expected initialized text message"); + }; + let JSONRPCMessage::Notification(notification) = + serde_json::from_str::(&encoded) + .expect("initialized notification should parse") + else { + panic!("expected initialized notification"); + }; + assert_eq!(notification.method, INITIALIZED_METHOD); +} diff --git a/codex-rs/exec-server/src/environment_bootstrap.rs b/codex-rs/exec-server/src/environment_bootstrap.rs new file mode 100644 index 0000000000000000000000000000000000000000..ad557c3b292d32702444da91f1d67535bf810a0c --- /dev/null +++ b/codex-rs/exec-server/src/environment_bootstrap.rs @@ -0,0 +1,66 @@ +use codex_http_client::HttpClientFactory; + +use crate::EnvironmentManager; +use crate::ExecServerError; +use crate::ExecServerRuntimePaths; +use crate::environment_provider::EnvironmentDefault; +use crate::environment_provider::EnvironmentProviderSnapshot; +use crate::remote::NoiseRendezvousEnvironmentConfig; + +#[derive(Debug)] +pub(crate) enum PreparedEnvironmentSource { + Noise(NoiseRendezvousEnvironmentConfig), + Snapshot(EnvironmentProviderSnapshot), +} + +/// Holds discovered execution environments before their HTTP policy is resolved. +/// +/// Preparing environments does not start remote connections. Callers can inspect +/// the default environment to choose config-loading behavior and then build the +/// manager with the effective outbound HTTP policy. +#[derive(Debug)] +pub struct PreparedEnvironmentManager { + pub(crate) source: PreparedEnvironmentSource, +} + +impl PreparedEnvironmentManager { + /// Returns whether the discovered default environment is remote. + pub fn default_environment_is_remote(&self) -> bool { + match &self.source { + PreparedEnvironmentSource::Noise(_) => true, + PreparedEnvironmentSource::Snapshot(snapshot) => match &snapshot.default { + EnvironmentDefault::Disabled => false, + EnvironmentDefault::EnvironmentId(default_id) => snapshot + .environments + .iter() + .any(|(environment_id, _)| environment_id == default_id), + }, + } + } + + /// Builds the manager and starts remote connections using the supplied policy. + pub fn build( + self, + local_runtime_paths: Option, + http_client_factory: HttpClientFactory, + ) -> Result { + match self.source { + PreparedEnvironmentSource::Noise(config) => { + EnvironmentManager::from_noise_environment_config( + config, + local_runtime_paths, + http_client_factory, + ) + } + PreparedEnvironmentSource::Snapshot(snapshot) => EnvironmentManager::from_snapshot( + snapshot, + local_runtime_paths, + http_client_factory, + ), + } + } +} + +#[cfg(test)] +#[path = "environment_bootstrap_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/environment_bootstrap_tests.rs b/codex-rs/exec-server/src/environment_bootstrap_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..d877871b227e5a4a9df313f2107bc5bdc4053a88 --- /dev/null +++ b/codex-rs/exec-server/src/environment_bootstrap_tests.rs @@ -0,0 +1,160 @@ +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use pretty_assertions::assert_eq; + +use super::PreparedEnvironmentManager; +use super::PreparedEnvironmentSource; +use crate::DefaultEnvironmentProvider; +use crate::LOCAL_ENVIRONMENT_ID; +use crate::REMOTE_ENVIRONMENT_ID; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::client_api::ExecServerTransportParams; +use crate::environment_provider::EnvironmentDefault; +use crate::environment_provider::EnvironmentProviderSnapshot; +use crate::remote::NoiseRendezvousEnvironmentConfig; + +#[test] +fn prepared_remote_environment_is_detected_without_constructing_a_connection() { + let prepared = PreparedEnvironmentManager { + source: PreparedEnvironmentSource::Snapshot(EnvironmentProviderSnapshot { + environments: vec![( + REMOTE_ENVIRONMENT_ID.to_string(), + ExecServerTransportParams::websocket_url( + "ws://username:password@executor.example/private?token=secret".to_string(), + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + ), + )], + default: EnvironmentDefault::EnvironmentId(REMOTE_ENVIRONMENT_ID.to_string()), + include_local: false, + }), + }; + + assert!(prepared.default_environment_is_remote()); + + let debug = format!("{prepared:?}"); + assert!(!debug.contains("username")); + assert!(!debug.contains("password")); + assert!(!debug.contains("executor.example")); + assert!(!debug.contains("secret")); +} + +#[test] +fn prepared_noise_environment_is_detected_before_http_policy_is_resolved() { + let config = NoiseRendezvousEnvironmentConfig::new( + "https://registry-user:registry-password@registry.example/api?access_token=query-secret#fragment-secret" + .to_string(), + "environment-requested".to_string(), + "registry-token".to_string(), + Some("workspace-123".to_string()), + ) + .expect("Noise environment configuration"); + let prepared = PreparedEnvironmentManager { + source: PreparedEnvironmentSource::Noise(config), + }; + + assert!(prepared.default_environment_is_remote()); + + let debug = format!("{prepared:?}"); + assert!(debug.contains("")); + assert!(!debug.contains("registry-token")); + assert!(!debug.contains("workspace-123")); + assert!(!debug.contains("registry-user")); + assert!(!debug.contains("registry-password")); + assert!(!debug.contains("registry.example")); + assert!(!debug.contains("query-secret")); + assert!(!debug.contains("fragment-secret")); +} + +#[test] +fn prepared_noise_environment_rejects_invalid_configuration() { + let invalid_configs = [ + ("", "environment-requested", "registry-token", None), + ("https://registry.example", "", "registry-token", None), + ( + "https://registry.example", + "environment-requested", + "", + None, + ), + ( + "https://registry.example", + "environment-requested", + "registry\ntoken", + None, + ), + ( + "https://registry.example", + "environment-requested", + "registry-token", + Some("workspace\n123"), + ), + ]; + + for (registry_url, environment_id, auth_token, chatgpt_account_id) in invalid_configs { + let result = NoiseRendezvousEnvironmentConfig::new( + registry_url.to_string(), + environment_id.to_string(), + auth_token.to_string(), + chatgpt_account_id.map(str::to_string), + ); + + assert!(result.is_err()); + } +} + +#[test] +fn prepared_local_and_disabled_environments_are_not_remote() { + let local = PreparedEnvironmentManager { + source: PreparedEnvironmentSource::Snapshot(EnvironmentProviderSnapshot { + environments: Vec::new(), + default: EnvironmentDefault::EnvironmentId(LOCAL_ENVIRONMENT_ID.to_string()), + include_local: true, + }), + }; + let disabled = PreparedEnvironmentManager { + source: PreparedEnvironmentSource::Snapshot(EnvironmentProviderSnapshot { + environments: Vec::new(), + default: EnvironmentDefault::Disabled, + include_local: false, + }), + }; + + assert_eq!( + [ + local.default_environment_is_remote(), + disabled.default_environment_is_remote() + ], + [false, false] + ); +} + +#[tokio::test] +async fn prepared_environment_manager_builds_with_the_explicit_http_policy() { + let prepared = PreparedEnvironmentManager { + source: PreparedEnvironmentSource::Snapshot( + DefaultEnvironmentProvider::new(Some("ws://127.0.0.1:8765".to_string())) + .snapshot_inner(), + ), + }; + let manager = prepared + .build( + /*local_runtime_paths*/ None, + HttpClientFactory::new(OutboundProxyPolicy::RespectSystemProxy), + ) + .expect("environment manager"); + + assert_eq!( + manager.default_environment_id(), + Some(REMOTE_ENVIRONMENT_ID) + ); + assert!( + manager + .default_environment() + .expect("remote environment") + .is_remote() + ); + assert_eq!( + manager.http_client_factory().outbound_proxy_policy(), + OutboundProxyPolicy::RespectSystemProxy + ); +} diff --git a/codex-rs/exec-server/src/environment_config.rs b/codex-rs/exec-server/src/environment_config.rs new file mode 100644 index 0000000000000000000000000000000000000000..b53208b63a87d10210a0be1b12c2a756270c0e81 --- /dev/null +++ b/codex-rs/exec-server/src/environment_config.rs @@ -0,0 +1,188 @@ +use std::collections::HashMap; + +use codex_config::CONFIG_TOML_FILE; +use codex_config::McpServerConfig; +use codex_config::McpServerDisabledReason; +use codex_config::RequirementSource; +use codex_config::RequirementsLayerEntry; +use codex_config::compose_requirements_for_hostname; +use codex_config::format_config_layer_source; +use codex_config::host_name; +use codex_config::loader::LocalTomlLayerStack; +use codex_config::loader::load_local_config_layers; +use codex_exec_server_protocol::EnvironmentConfigLayer; +use codex_exec_server_protocol::EnvironmentConfigLayerStack; +use codex_exec_server_protocol::EnvironmentConfigReadParams; +use codex_exec_server_protocol::EnvironmentConfigReadResponse; +use codex_file_system::ExecutorFileSystem; +use codex_utils_home_dir::find_codex_home; +use codex_utils_path_uri::PathUri; + +use crate::Environment; +use crate::ExecServerError; + +#[derive(Debug, thiserror::Error)] +pub(crate) enum ReadEnvironmentConfigError { + #[error("{0}")] + InvalidParams(String), + #[error("{0}")] + Internal(String), +} + +pub(crate) async fn read_environment_config( + file_system: &dyn ExecutorFileSystem, + params: EnvironmentConfigReadParams, +) -> Result { + validate_paths(¶ms)?; + let cwd = params + .cwd + .to_abs_path() + .map_err(|error| ReadEnvironmentConfigError::InvalidParams(error.to_string()))?; + let codex_home = find_codex_home().map_err(|error| { + ReadEnvironmentConfigError::Internal(format!("failed to find Codex home: {error}")) + })?; + let layers = load_local_config_layers(file_system, codex_home.as_path(), &cwd) + .await + .map_err(|error| { + ReadEnvironmentConfigError::Internal(format!( + "failed to load executor-local config: {error}" + )) + })? + .project(¶ms.config_paths, ¶ms.requirements_paths); + + Ok(EnvironmentConfigReadResponse { + user_home_dir: dirs::home_dir() + .and_then(|home_dir| PathUri::from_host_native_path(home_dir).ok()), + codex_home_dir: PathUri::from_abs_path(&codex_home), + hostname: host_name(), + config: serialize_layer_stack(layers.config, |source| { + format_config_layer_source(source, CONFIG_TOML_FILE) + })?, + requirements: serialize_layer_stack(layers.requirements, ToString::to_string)?, + }) +} + +impl Environment { + /// Reads executor-owned HTTP MCP servers from the selected project's current config. + pub async fn discover_http_mcp_servers( + &self, + cwd: PathUri, + ) -> Result, ExecServerError> { + let (response, http_header_env_vars) = + tokio::time::timeout(std::time::Duration::from_secs(10), async { + let capabilities = self.info().await?.capabilities; + if !capabilities.environment_config_read { + return Ok((None, false)); + } + self.read_environment_config(EnvironmentConfigReadParams { + cwd, + config_paths: vec![vec!["mcp_servers".to_string()]], + requirements_paths: vec![vec!["mcp_servers".to_string()]], + }) + .await + .map(|response| (Some(response), capabilities.http_header_env_vars)) + }) + .await + .map_err(|_| { + ExecServerError::Protocol("executor MCP discovery timed out".to_string()) + })??; + let Some(response) = response else { + return Ok(Vec::new()); + }; + let requirements = compose_requirements_for_hostname( + response.requirements.layers.into_iter().map(|layer| { + RequirementsLayerEntry::from_toml(RequirementSource::Unknown, layer.toml) + }), + response.hostname.as_deref(), + ) + .map_err(|error| { + ExecServerError::Protocol(format!("invalid executor-local MCP requirements: {error}")) + })? + .and_then(|requirements| requirements.mcp_servers); + let mut merged = toml::Value::Table(toml::map::Map::new()); + for layer in response.config.layers { + let config = toml::from_str::(&layer.toml).map_err(|error| { + ExecServerError::Protocol(format!("invalid executor-local MCP config: {error}")) + })?; + codex_config::merge_toml_values(&mut merged, &config); + } + let mut servers = merged + .get("mcp_servers") + .cloned() + .unwrap_or_else(|| toml::Value::Table(toml::map::Map::new())); + if let Some(servers) = servers.as_table_mut() { + servers.retain(|_, server| { + server.get("url").is_some() + && (http_header_env_vars || server.get("bearer_token_env_var").is_none()) + }); + } + let servers = servers + .try_into::>() + .map_err(|error| { + ExecServerError::Protocol(format!("invalid executor-local MCP servers: {error}")) + })?; + + Ok(servers + .into_iter() + .map(|(name, mut server)| { + if let Some(requirements) = requirements.as_ref() + && !requirements + .value + .get(&name) + .is_some_and(|requirement| server.matches_requirement(requirement)) + { + server.enabled = false; + server.disabled_reason = Some(McpServerDisabledReason::Requirements { + source: requirements.source.clone(), + }); + } + (name, server) + }) + .collect()) + } +} + +fn validate_paths(params: &EnvironmentConfigReadParams) -> Result<(), ReadEnvironmentConfigError> { + if params.config_paths.is_empty() && params.requirements_paths.is_empty() { + return Err(ReadEnvironmentConfigError::InvalidParams( + "at least one config or requirements path is required".to_string(), + )); + } + if params + .config_paths + .iter() + .chain(¶ms.requirements_paths) + .any(Vec::is_empty) + { + return Err(ReadEnvironmentConfigError::InvalidParams( + "TOML paths must contain at least one key segment".to_string(), + )); + } + Ok(()) +} + +fn serialize_layer_stack( + stack: LocalTomlLayerStack, + source_name: impl Fn(&S) -> String, +) -> Result { + let layers = stack + .layers + .into_iter() + .map(|layer| { + let toml = toml::to_string(&layer.toml).map_err(|error| { + ReadEnvironmentConfigError::Internal(format!( + "failed to serialize executor-local config: {error}" + )) + })?; + Ok(EnvironmentConfigLayer { + source: source_name(&layer.source), + base_dir: PathUri::from_abs_path(&layer.base_dir), + toml, + }) + }) + .collect::, ReadEnvironmentConfigError>>()?; + Ok(EnvironmentConfigLayerStack { + layers, + cloud_insertion_index: stack.cloud_insertion_index, + }) +} diff --git a/codex-rs/exec-server/src/environment_provider.rs b/codex-rs/exec-server/src/environment_provider.rs new file mode 100644 index 0000000000000000000000000000000000000000..2366eb8be24b2d472dbc6245bc876fcc02704e6b --- /dev/null +++ b/codex-rs/exec-server/src/environment_provider.rs @@ -0,0 +1,211 @@ +use std::future::Future; +use std::pin::Pin; + +use crate::ExecServerError; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::client_api::ExecServerTransportParams; +use crate::environment::CODEX_EXEC_SERVER_URL_ENV_VAR; +use crate::environment::LOCAL_ENVIRONMENT_ID; +use crate::environment::REMOTE_ENVIRONMENT_ID; + +/// Lists the remote environment transports available to Codex. +/// +/// Implementations own a startup snapshot containing both the available +/// environment transport list in configured order and the default environment +/// selection. Providers return transport descriptions before the effective HTTP +/// policy is available; `include_local` controls whether `EnvironmentManager` +/// should add the local environment when the snapshot is built. +pub trait EnvironmentProvider: Send + Sync { + /// Returns the provider-owned environment startup snapshot. + fn snapshot(&self) -> EnvironmentProviderFuture<'_>; +} + +pub type EnvironmentProviderFuture<'a> = + Pin> + Send + 'a>>; + +#[derive(Clone)] +pub struct EnvironmentProviderSnapshot { + pub(crate) environments: Vec<(String, ExecServerTransportParams)>, + pub default: EnvironmentDefault, + pub include_local: bool, +} + +impl std::fmt::Debug for EnvironmentProviderSnapshot { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let environment_ids: Vec<_> = self.environments.iter().map(|(id, _)| id).collect(); + f.debug_struct("EnvironmentProviderSnapshot") + .field("environments", &environment_ids) + .field("default", &self.default) + .field("include_local", &self.include_local) + .finish() + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum EnvironmentDefault { + Disabled, + EnvironmentId(String), +} + +/// Default provider backed by `CODEX_EXEC_SERVER_URL`. +#[derive(Clone, Debug)] +pub struct DefaultEnvironmentProvider { + exec_server_url: Option, +} + +impl DefaultEnvironmentProvider { + /// Builds a provider from an already-read raw `CODEX_EXEC_SERVER_URL` value. + pub fn new(exec_server_url: Option) -> Self { + Self { exec_server_url } + } + + /// Builds a provider by reading `CODEX_EXEC_SERVER_URL`. + pub fn from_env() -> Self { + Self::new(std::env::var(CODEX_EXEC_SERVER_URL_ENV_VAR).ok()) + } + + pub(crate) fn snapshot_inner(&self) -> EnvironmentProviderSnapshot { + let mut environments = Vec::new(); + let (exec_server_url, disabled) = normalize_exec_server_url(self.exec_server_url.clone()); + + if let Some(exec_server_url) = exec_server_url { + environments.push(( + REMOTE_ENVIRONMENT_ID.to_string(), + ExecServerTransportParams::websocket_url( + exec_server_url, + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + ), + )); + } + + let has_remote = environments + .iter() + .any(|(id, _environment)| id == REMOTE_ENVIRONMENT_ID); + let include_local = !disabled && !has_remote; + let default = if disabled { + EnvironmentDefault::Disabled + } else if has_remote { + EnvironmentDefault::EnvironmentId(REMOTE_ENVIRONMENT_ID.to_string()) + } else { + EnvironmentDefault::EnvironmentId(LOCAL_ENVIRONMENT_ID.to_string()) + }; + + EnvironmentProviderSnapshot { + environments, + default, + include_local, + } + } +} + +impl EnvironmentProvider for DefaultEnvironmentProvider { + fn snapshot(&self) -> EnvironmentProviderFuture<'_> { + Box::pin(async { Ok(self.snapshot_inner()) }) + } +} + +pub(crate) fn normalize_exec_server_url(exec_server_url: Option) -> (Option, bool) { + match exec_server_url.as_deref().map(str::trim) { + None | Some("") => (None, false), + Some(url) if url.eq_ignore_ascii_case("none") => (None, true), + Some(url) => (Some(url.to_string()), false), + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use pretty_assertions::assert_eq; + + use super::*; + + #[tokio::test] + async fn default_provider_requests_local_environment_when_url_is_missing() { + let provider = DefaultEnvironmentProvider::new(/*exec_server_url*/ None); + let snapshot = provider.snapshot().await.expect("environments"); + let EnvironmentProviderSnapshot { + environments, + default, + include_local, + } = snapshot; + let environments: HashMap<_, _> = environments.into_iter().collect(); + + assert!(include_local); + assert!(!environments.contains_key(LOCAL_ENVIRONMENT_ID)); + assert!(!environments.contains_key(REMOTE_ENVIRONMENT_ID)); + assert_eq!( + default, + EnvironmentDefault::EnvironmentId(LOCAL_ENVIRONMENT_ID.to_string()) + ); + } + + #[tokio::test] + async fn default_provider_requests_local_environment_when_url_is_empty() { + let provider = DefaultEnvironmentProvider::new(Some(String::new())); + let snapshot = provider.snapshot().await.expect("environments"); + let EnvironmentProviderSnapshot { + environments, + default, + include_local, + } = snapshot; + let environments: HashMap<_, _> = environments.into_iter().collect(); + + assert!(include_local); + assert!(!environments.contains_key(LOCAL_ENVIRONMENT_ID)); + assert!(!environments.contains_key(REMOTE_ENVIRONMENT_ID)); + assert_eq!( + default, + EnvironmentDefault::EnvironmentId(LOCAL_ENVIRONMENT_ID.to_string()) + ); + } + + #[tokio::test] + async fn default_provider_omits_local_environment_for_none_value() { + let provider = DefaultEnvironmentProvider::new(Some("none".to_string())); + let snapshot = provider.snapshot().await.expect("environments"); + let EnvironmentProviderSnapshot { + environments, + default, + include_local, + } = snapshot; + let environments: HashMap<_, _> = environments.into_iter().collect(); + + assert!(!include_local); + assert!(!environments.contains_key(LOCAL_ENVIRONMENT_ID)); + assert!(!environments.contains_key(REMOTE_ENVIRONMENT_ID)); + assert_eq!(default, EnvironmentDefault::Disabled); + } + + #[tokio::test] + async fn default_provider_adds_remote_environment_for_websocket_url() { + let provider = DefaultEnvironmentProvider::new(Some("ws://127.0.0.1:8765".to_string())); + let snapshot = provider.snapshot().await.expect("environments"); + let EnvironmentProviderSnapshot { + environments, + default, + include_local, + } = snapshot; + let environments: HashMap<_, _> = environments.into_iter().collect(); + + assert!(!include_local); + assert!(!environments.contains_key(LOCAL_ENVIRONMENT_ID)); + assert!(matches!( + &environments[REMOTE_ENVIRONMENT_ID], + ExecServerTransportParams::WebSocketUrl { websocket_url, .. } + if websocket_url == "ws://127.0.0.1:8765" + )); + assert_eq!( + default, + EnvironmentDefault::EnvironmentId(REMOTE_ENVIRONMENT_ID.to_string()) + ); + } + + #[test] + fn normalizes_exec_server_url() { + assert_eq!( + normalize_exec_server_url(Some(" ws://127.0.0.1:8765 ".to_string())), + (Some("ws://127.0.0.1:8765".to_string()), false) + ); + } +} diff --git a/codex-rs/exec-server/src/environment_registry.rs b/codex-rs/exec-server/src/environment_registry.rs new file mode 100644 index 0000000000000000000000000000000000000000..bc0755fd68b4d5634d3728fff7d5201150b29a9e --- /dev/null +++ b/codex-rs/exec-server/src/environment_registry.rs @@ -0,0 +1,68 @@ +use serde::Deserialize; +use serde::Serialize; + +use crate::NoiseChannelPublicKey; + +/// Request body for registering an executor with the environment registry. +#[derive(Clone, Deserialize, Eq, PartialEq, Serialize)] +pub struct EnvironmentRegistryRegistrationRequest { + pub security_profile: String, + pub executor_public_key: NoiseChannelPublicKey, +} + +/// Environment registry response returned after executor registration. +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +pub struct EnvironmentRegistryRegistrationResponse { + pub environment_id: String, + pub url: String, + pub security_profile: String, + pub executor_registration_id: String, +} + +/// Request body for connecting a harness key with the environment registry. +#[derive(Clone, Deserialize, Eq, PartialEq, Serialize)] +pub struct EnvironmentRegistryConnectRequest { + pub harness_public_key: NoiseChannelPublicKey, +} + +/// Environment registry response returned after connecting a harness key. +#[derive(Clone, Deserialize, Eq, PartialEq, Serialize)] +pub struct EnvironmentRegistryConnectResponse { + pub environment_id: String, + pub url: String, + pub security_profile: String, + pub executor_registration_id: String, + pub executor_public_key: NoiseChannelPublicKey, + pub harness_key_authorization: String, +} + +impl std::fmt::Debug for EnvironmentRegistryConnectResponse { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("EnvironmentRegistryConnectResponse") + .field("environment_id", &self.environment_id) + .field("url", &"") + .field("security_profile", &self.security_profile) + .field("executor_registration_id", &self.executor_registration_id) + .field("executor_public_key", &self.executor_public_key) + .field("harness_key_authorization", &"") + .finish() + } +} + +/// Request body for authorizing a harness key with the environment registry. +#[derive(Clone, Deserialize, Eq, PartialEq, Serialize)] +pub struct EnvironmentRegistryHarnessKeyValidationRequest { + pub executor_registration_id: String, + pub harness_public_key: NoiseChannelPublicKey, + pub harness_key_authorization: String, +} + +/// Environment registry response returned after harness key validation. +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +pub struct EnvironmentRegistryHarnessKeyValidationResponse { + pub valid: bool, +} + +#[cfg(test)] +#[path = "environment_registry_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/environment_registry_tests.rs b/codex-rs/exec-server/src/environment_registry_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..8c32e0fb45af19d862509bb8625001f1ee9840aa --- /dev/null +++ b/codex-rs/exec-server/src/environment_registry_tests.rs @@ -0,0 +1,22 @@ +use crate::EnvironmentRegistryConnectResponse; +use crate::NoiseChannelIdentity; + +#[test] +fn connect_response_debug_redacts_authorizations() { + let response = EnvironmentRegistryConnectResponse { + environment_id: "environment-1".to_string(), + url: "wss://rendezvous.test?sig=secret-url-authorization".to_string(), + security_profile: "noise_hybrid_ik_v1".to_string(), + executor_registration_id: "registration-1".to_string(), + executor_public_key: NoiseChannelIdentity::generate() + .expect("identity") + .public_key(), + harness_key_authorization: "secret-harness-authorization".to_string(), + }; + + let debug = format!("{response:?}"); + + assert!(debug.contains("")); + assert!(!debug.contains("secret-url-authorization")); + assert!(!debug.contains("secret-harness-authorization")); +} diff --git a/codex-rs/exec-server/src/environment_toml.rs b/codex-rs/exec-server/src/environment_toml.rs new file mode 100644 index 0000000000000000000000000000000000000000..e86c956fdfc4824de62d21564af82be973196fa7 --- /dev/null +++ b/codex-rs/exec-server/src/environment_toml.rs @@ -0,0 +1,896 @@ +use std::collections::HashMap; +use std::collections::HashSet; +use std::path::Path; +use std::path::PathBuf; +use std::time::Duration; + +use http::HeaderMap; +use serde::Deserialize; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; + +use crate::DefaultEnvironmentProvider; +use crate::EnvironmentProvider; +use crate::EnvironmentProviderFuture; +use crate::ExecServerError; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT; +use crate::client_api::ExecServerTransportParams; +use crate::client_api::StdioExecServerCommand; +use crate::environment::LOCAL_ENVIRONMENT_ID; +use crate::environment_provider::EnvironmentDefault; +use crate::environment_provider::EnvironmentProviderSnapshot; + +const ENVIRONMENTS_TOML_FILE: &str = "environments.toml"; +const MAX_ENVIRONMENT_ID_LEN: usize = 64; + +#[derive(Deserialize, Debug, Default)] +#[serde(deny_unknown_fields)] +struct EnvironmentsToml { + default: Option, + include_local: Option, + + #[serde(default)] + environments: Vec, +} + +#[derive(Deserialize, Debug, Default, PartialEq, Eq)] +#[serde(deny_unknown_fields)] +struct EnvironmentToml { + id: String, + url: Option, + program: Option, + args: Option>, + env: Option>, + cwd: Option, + #[serde(default, with = "option_duration_secs")] + connect_timeout_sec: Option, + #[serde(default, with = "option_duration_secs")] + initialize_timeout_sec: Option, +} + +#[derive(Clone, Debug)] +struct TomlEnvironmentProvider { + default: EnvironmentDefault, + include_local: bool, + environments: Vec<(String, ExecServerTransportParams)>, +} + +impl TomlEnvironmentProvider { + #[cfg(test)] + fn new(config: EnvironmentsToml) -> Result { + Self::new_with_config_dir(config, /*config_dir*/ None) + } + + fn new_with_config_dir( + config: EnvironmentsToml, + config_dir: Option<&Path>, + ) -> Result { + let EnvironmentsToml { + default, + include_local, + environments, + } = config; + let include_local = include_local.unwrap_or(true); + let mut ids = HashSet::new(); + if include_local { + ids.insert(LOCAL_ENVIRONMENT_ID.to_string()); + } + let mut parsed_environments = Vec::with_capacity(environments.len()); + for item in environments { + let (id, transport) = parse_environment_toml(item, config_dir)?; + if !ids.insert(id.clone()) { + return Err(ExecServerError::Protocol(format!( + "environment id `{id}` is duplicated" + ))); + } + parsed_environments.push((id, transport)); + } + let default = normalize_default_environment_id(default.as_deref(), include_local, &ids)?; + Ok(Self { + default, + include_local, + environments: parsed_environments, + }) + } + + async fn snapshot(&self) -> Result { + Ok(EnvironmentProviderSnapshot { + environments: self.environments.clone(), + default: self.default.clone(), + include_local: self.include_local, + }) + } +} + +impl EnvironmentProvider for TomlEnvironmentProvider { + fn snapshot(&self) -> EnvironmentProviderFuture<'_> { + Box::pin(TomlEnvironmentProvider::snapshot(self)) + } +} + +fn parse_environment_toml( + item: EnvironmentToml, + config_dir: Option<&Path>, +) -> Result<(String, ExecServerTransportParams), ExecServerError> { + let EnvironmentToml { + id, + url, + program, + args, + env, + cwd, + connect_timeout_sec, + initialize_timeout_sec, + } = item; + validate_environment_id(&id)?; + if program.is_none() && (args.is_some() || env.is_some() || cwd.is_some()) { + return Err(ExecServerError::Protocol(format!( + "environment `{id}` args, env, and cwd require program" + ))); + } + if url.is_none() && connect_timeout_sec.is_some() { + return Err(ExecServerError::Protocol(format!( + "environment `{id}` connect_timeout_sec requires url" + ))); + } + + let connect_timeout = connect_timeout_sec.unwrap_or(DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT); + let initialize_timeout = + initialize_timeout_sec.unwrap_or(DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT); + + let transport_params = match (url, program) { + (Some(url), None) => { + let url = validate_websocket_url(url)?; + ExecServerTransportParams::WebSocketUrl { + websocket_url: url, + connect_timeout, + initialize_timeout, + http_headers: HeaderMap::new(), + } + } + (None, Some(program)) => { + let program = program.trim().to_string(); + if program.is_empty() { + return Err(ExecServerError::Protocol(format!( + "environment `{id}` program cannot be empty" + ))); + } + let cwd = normalize_stdio_cwd(&id, cwd, config_dir)?; + ExecServerTransportParams::StdioCommand { + command: StdioExecServerCommand { + program, + args: args.unwrap_or_default(), + env: env.unwrap_or_default(), + cwd, + }, + initialize_timeout, + } + } + (None, None) | (Some(_), Some(_)) => { + return Err(ExecServerError::Protocol(format!( + "environment `{id}` must set exactly one of url or program" + ))); + } + }; + + Ok((id, transport_params)) +} + +fn normalize_stdio_cwd( + id: &str, + cwd: Option, + config_dir: Option<&Path>, +) -> Result, ExecServerError> { + let Some(cwd) = cwd else { + return Ok(None); + }; + if cwd.is_absolute() { + return Ok(Some(cwd)); + } + let Some(config_dir) = config_dir else { + return Err(ExecServerError::Protocol(format!( + "environment `{id}` cwd must be absolute" + ))); + }; + Ok(Some(config_dir.join(cwd))) +} + +pub(crate) fn environment_provider_from_codex_home( + codex_home: &Path, +) -> Result, ExecServerError> { + let path = codex_home.join(ENVIRONMENTS_TOML_FILE); + let Some(environments) = load_environments_toml(&path)? else { + return Ok(Box::new(DefaultEnvironmentProvider::from_env())); + }; + + Ok(Box::new(TomlEnvironmentProvider::new_with_config_dir( + environments, + Some(codex_home), + )?)) +} + +fn normalize_default_environment_id( + default: Option<&str>, + include_local: bool, + ids: &HashSet, +) -> Result { + let Some(default) = default.map(str::trim) else { + return if include_local { + Ok(EnvironmentDefault::EnvironmentId( + LOCAL_ENVIRONMENT_ID.to_string(), + )) + } else { + Ok(EnvironmentDefault::Disabled) + }; + }; + if default.is_empty() { + return Err(ExecServerError::Protocol( + "default environment id cannot be empty".to_string(), + )); + } + if !default.eq_ignore_ascii_case("none") && !ids.contains(default) { + return Err(ExecServerError::Protocol(format!( + "default environment `{default}` is not configured" + ))); + } + if default.eq_ignore_ascii_case("none") { + Ok(EnvironmentDefault::Disabled) + } else { + Ok(EnvironmentDefault::EnvironmentId(default.to_string())) + } +} + +fn validate_environment_id(id: &str) -> Result<(), ExecServerError> { + let trimmed_id = id.trim(); + if trimmed_id.is_empty() { + return Err(ExecServerError::Protocol( + "environment id cannot be empty".to_string(), + )); + } + if trimmed_id != id { + return Err(ExecServerError::Protocol(format!( + "environment id `{id}` must not contain surrounding whitespace" + ))); + } + if id == LOCAL_ENVIRONMENT_ID || id.eq_ignore_ascii_case("none") { + return Err(ExecServerError::Protocol(format!( + "environment id `{id}` is reserved" + ))); + } + if id.len() > MAX_ENVIRONMENT_ID_LEN { + return Err(ExecServerError::Protocol(format!( + "environment id `{id}` cannot be longer than {MAX_ENVIRONMENT_ID_LEN} characters" + ))); + } + if !id + .chars() + .all(|ch| ch.is_ascii_alphanumeric() || ch == '-' || ch == '_') + { + return Err(ExecServerError::Protocol(format!( + "environment id `{id}` must contain only ASCII letters, numbers, '-' or '_'" + ))); + } + Ok(()) +} + +fn validate_websocket_url(url: String) -> Result { + let url = url.trim(); + if url.is_empty() { + return Err(ExecServerError::Protocol( + "environment url cannot be empty".to_string(), + )); + } + if !url.starts_with("ws://") && !url.starts_with("wss://") { + return Err(ExecServerError::Protocol(format!( + "environment url `{url}` must use ws:// or wss://" + ))); + } + url.into_client_request().map_err(|err| { + ExecServerError::Protocol(format!("environment url `{url}` is invalid: {err}")) + })?; + Ok(url.to_string()) +} + +/// Returns `None` when the config is missing; other I/O and parse failures remain errors. +fn load_environments_toml(path: &Path) -> Result, ExecServerError> { + let contents = match std::fs::read_to_string(path) { + Ok(contents) => contents, + Err(err) if err.kind() == std::io::ErrorKind::NotFound => return Ok(None), + Err(err) => { + return Err(ExecServerError::Protocol(format!( + "failed to read environment config `{}`: {err}", + path.display() + ))); + } + }; + + toml::from_str(&contents) + .map_err(|err| { + ExecServerError::Protocol(format!( + "failed to parse environment config `{}`: {err}", + path.display() + )) + }) + .map(Some) +} + +mod option_duration_secs { + use std::time::Duration; + + use serde::Deserialize; + use serde::Deserializer; + + pub fn deserialize<'de, D>(deserializer: D) -> Result, D::Error> + where + D: Deserializer<'de>, + { + let secs = Option::::deserialize(deserializer)?; + secs.map(|secs| Duration::try_from_secs_f64(secs).map_err(serde::de::Error::custom)) + .transpose() + } +} + +#[cfg(test)] +mod tests { + use pretty_assertions::assert_eq; + use tempfile::tempdir; + + use super::*; + + #[tokio::test] + async fn toml_provider_includes_local_and_adds_configured_environments() { + let provider = TomlEnvironmentProvider::new(EnvironmentsToml { + default: Some("ssh-dev".to_string()), + include_local: None, + environments: vec![ + EnvironmentToml { + id: "devbox".to_string(), + url: Some(" ws://127.0.0.1:8765 ".to_string()), + ..Default::default() + }, + EnvironmentToml { + id: "ssh-dev".to_string(), + program: Some(" ssh ".to_string()), + args: Some(vec![ + "dev".to_string(), + "codex exec-server --listen stdio".to_string(), + ]), + env: Some(HashMap::from([( + "CODEX_LOG".to_string(), + "debug".to_string(), + )])), + ..Default::default() + }, + ], + }) + .expect("provider"); + + let snapshot = provider.snapshot().await.expect("environments"); + let EnvironmentProviderSnapshot { + environments, + default, + include_local, + } = snapshot; + let environment_ids: Vec<_> = environments + .iter() + .map(|(id, _environment)| id.as_str()) + .collect(); + assert_eq!(environment_ids, vec!["devbox", "ssh-dev"]); + let environments: HashMap<_, _> = environments.into_iter().collect(); + + assert!(include_local); + assert!(!environments.contains_key(LOCAL_ENVIRONMENT_ID)); + assert!(matches!( + &environments["devbox"], + ExecServerTransportParams::WebSocketUrl { .. } + )); + assert!(matches!( + &environments["ssh-dev"], + ExecServerTransportParams::StdioCommand { .. } + )); + assert_eq!( + default, + EnvironmentDefault::EnvironmentId("ssh-dev".to_string()) + ); + } + + #[tokio::test] + async fn toml_provider_default_omitted_selects_local() { + let provider = TomlEnvironmentProvider::new(EnvironmentsToml::default()).expect("provider"); + let snapshot = provider.snapshot().await.expect("environments"); + + assert!(snapshot.include_local); + assert_eq!( + snapshot.default, + EnvironmentDefault::EnvironmentId(LOCAL_ENVIRONMENT_ID.to_string()) + ); + } + + #[tokio::test] + async fn toml_provider_default_none_disables_default() { + let provider = TomlEnvironmentProvider::new(EnvironmentsToml { + default: Some("none".to_string()), + include_local: None, + environments: Vec::new(), + }) + .expect("provider"); + let snapshot = provider.snapshot().await.expect("environments"); + + assert!(snapshot.include_local); + assert_eq!(snapshot.default, EnvironmentDefault::Disabled); + } + + #[tokio::test] + async fn toml_provider_can_disable_local_environment() { + let provider = TomlEnvironmentProvider::new(EnvironmentsToml { + default: Some("ssh-dev".to_string()), + include_local: Some(false), + environments: vec![EnvironmentToml { + id: "ssh-dev".to_string(), + program: Some("ssh".to_string()), + ..Default::default() + }], + }) + .expect("provider"); + let snapshot = provider.snapshot().await.expect("environments"); + + assert!(!snapshot.include_local); + assert_eq!( + snapshot.default, + EnvironmentDefault::EnvironmentId("ssh-dev".to_string()) + ); + } + + #[tokio::test] + async fn toml_provider_without_local_and_default_omitted_disables_default() { + let provider = TomlEnvironmentProvider::new(EnvironmentsToml { + include_local: Some(false), + ..Default::default() + }) + .expect("provider"); + let snapshot = provider.snapshot().await.expect("environments"); + + assert!(!snapshot.include_local); + assert_eq!(snapshot.default, EnvironmentDefault::Disabled); + } + + #[test] + fn toml_provider_rejects_local_default_when_local_is_disabled() { + let err = TomlEnvironmentProvider::new(EnvironmentsToml { + default: Some(LOCAL_ENVIRONMENT_ID.to_string()), + include_local: Some(false), + environments: Vec::new(), + }) + .expect_err("local default without local environment should fail"); + + assert_eq!( + err.to_string(), + "exec-server protocol error: default environment `local` is not configured" + ); + } + + #[test] + fn toml_provider_rejects_invalid_environments() { + let cases = [ + ( + EnvironmentToml { + id: "local".to_string(), + url: Some("ws://127.0.0.1:8765".to_string()), + ..Default::default() + }, + "environment id `local` is reserved", + ), + ( + EnvironmentToml { + id: " devbox ".to_string(), + url: Some("ws://127.0.0.1:8765".to_string()), + ..Default::default() + }, + "environment id ` devbox ` must not contain surrounding whitespace", + ), + ( + EnvironmentToml { + id: "dev box".to_string(), + url: Some("ws://127.0.0.1:8765".to_string()), + ..Default::default() + }, + "environment id `dev box` must contain only ASCII letters, numbers, '-' or '_'", + ), + ( + EnvironmentToml { + id: "devbox".to_string(), + url: Some("http://127.0.0.1:8765".to_string()), + ..Default::default() + }, + "environment url `http://127.0.0.1:8765` must use ws:// or wss://", + ), + ( + EnvironmentToml { + id: "devbox".to_string(), + url: Some("ws://127.0.0.1:8765".to_string()), + program: Some("codex".to_string()), + ..Default::default() + }, + "environment `devbox` must set exactly one of url or program", + ), + ( + EnvironmentToml { + id: "devbox".to_string(), + program: Some(" ".to_string()), + ..Default::default() + }, + "environment `devbox` program cannot be empty", + ), + ( + EnvironmentToml { + id: "devbox".to_string(), + args: Some(Vec::new()), + ..Default::default() + }, + "environment `devbox` args, env, and cwd require program", + ), + ( + EnvironmentToml { + id: "ssh-dev".to_string(), + program: Some("ssh".to_string()), + connect_timeout_sec: Some(Duration::from_secs(1)), + ..Default::default() + }, + "environment `ssh-dev` connect_timeout_sec requires url", + ), + ]; + + for (item, expected) in cases { + let err = TomlEnvironmentProvider::new(EnvironmentsToml { + default: None, + include_local: None, + environments: vec![item], + }) + .expect_err("invalid item should fail"); + + assert_eq!( + err.to_string(), + format!("exec-server protocol error: {expected}") + ); + } + } + + #[test] + fn toml_provider_resolves_relative_stdio_cwd_from_config_dir() { + let config_dir = tempdir().expect("tempdir"); + let provider = TomlEnvironmentProvider::new_with_config_dir( + EnvironmentsToml { + default: None, + include_local: None, + environments: vec![EnvironmentToml { + id: "ssh-dev".to_string(), + program: Some("ssh".to_string()), + cwd: Some(PathBuf::from("workspace")), + ..Default::default() + }], + }, + Some(config_dir.path()), + ) + .expect("provider"); + + let ExecServerTransportParams::StdioCommand { + command, + initialize_timeout, + } = &provider.environments[0].1 + else { + panic!("expected stdio transport"); + }; + assert_eq!( + command, + &StdioExecServerCommand { + program: "ssh".to_string(), + args: Vec::new(), + env: HashMap::new(), + cwd: Some(config_dir.path().join("workspace")), + } + ); + assert_eq!( + *initialize_timeout, + DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT + ); + } + + #[test] + fn toml_provider_parses_configured_transport_timeouts() { + let provider = TomlEnvironmentProvider::new(EnvironmentsToml { + default: None, + include_local: None, + environments: vec![ + EnvironmentToml { + id: "devbox".to_string(), + url: Some("ws://127.0.0.1:8765".to_string()), + connect_timeout_sec: Some(Duration::from_secs(12)), + initialize_timeout_sec: Some(Duration::from_secs(34)), + ..Default::default() + }, + EnvironmentToml { + id: "ssh-dev".to_string(), + program: Some("ssh".to_string()), + initialize_timeout_sec: Some(Duration::from_secs(56)), + ..Default::default() + }, + ], + }) + .expect("provider"); + + let ExecServerTransportParams::WebSocketUrl { + websocket_url, + connect_timeout, + initialize_timeout, + .. + } = &provider.environments[0].1 + else { + panic!("expected websocket transport"); + }; + assert_eq!(websocket_url, "ws://127.0.0.1:8765"); + assert_eq!(*connect_timeout, Duration::from_secs(12)); + assert_eq!(*initialize_timeout, Duration::from_secs(34)); + + let ExecServerTransportParams::StdioCommand { + command, + initialize_timeout, + } = &provider.environments[1].1 + else { + panic!("expected stdio transport"); + }; + assert_eq!( + command, + &StdioExecServerCommand { + program: "ssh".to_string(), + args: Vec::new(), + env: HashMap::new(), + cwd: None, + } + ); + assert_eq!(*initialize_timeout, Duration::from_secs(56)); + } + + #[test] + fn toml_provider_rejects_relative_stdio_cwd_without_config_dir() { + let err = TomlEnvironmentProvider::new(EnvironmentsToml { + default: None, + include_local: None, + environments: vec![EnvironmentToml { + id: "ssh-dev".to_string(), + program: Some("ssh".to_string()), + cwd: Some(PathBuf::from("workspace")), + ..Default::default() + }], + }) + .expect_err("relative cwd without config dir should fail"); + + assert_eq!( + err.to_string(), + "exec-server protocol error: environment `ssh-dev` cwd must be absolute" + ); + } + + #[test] + fn toml_provider_rejects_duplicate_ids() { + let err = TomlEnvironmentProvider::new(EnvironmentsToml { + default: None, + include_local: None, + environments: vec![ + EnvironmentToml { + id: "devbox".to_string(), + url: Some("ws://127.0.0.1:8765".to_string()), + ..Default::default() + }, + EnvironmentToml { + id: "devbox".to_string(), + program: Some("codex".to_string()), + ..Default::default() + }, + ], + }) + .expect_err("duplicate id should fail"); + + assert_eq!( + err.to_string(), + "exec-server protocol error: environment id `devbox` is duplicated" + ); + } + + #[test] + fn toml_provider_rejects_overlong_id() { + let id = "a".repeat(MAX_ENVIRONMENT_ID_LEN + 1); + let err = TomlEnvironmentProvider::new(EnvironmentsToml { + default: None, + include_local: None, + environments: vec![EnvironmentToml { + id: id.clone(), + url: Some("ws://127.0.0.1:8765".to_string()), + ..Default::default() + }], + }) + .expect_err("overlong id should fail"); + + assert_eq!( + err.to_string(), + format!( + "exec-server protocol error: environment id `{id}` cannot be longer than {MAX_ENVIRONMENT_ID_LEN} characters" + ) + ); + } + + #[test] + fn toml_provider_rejects_unknown_default() { + let err = TomlEnvironmentProvider::new(EnvironmentsToml { + default: Some("missing".to_string()), + include_local: None, + environments: Vec::new(), + }) + .expect_err("unknown default should fail"); + + assert_eq!( + err.to_string(), + "exec-server protocol error: default environment `missing` is not configured" + ); + } + + #[test] + fn load_environments_toml_reads_root_environment_list() { + let codex_home = tempdir().expect("tempdir"); + let path = codex_home.path().join(ENVIRONMENTS_TOML_FILE); + std::fs::write( + &path, + r#" +default = "ssh-dev" +include_local = false + +[[environments]] +id = "devbox" +url = "ws://127.0.0.1:4512" +connect_timeout_sec = 12.0 +initialize_timeout_sec = 34.0 + +[[environments]] +id = "ssh-dev" +program = "ssh" +args = ["dev", "codex exec-server --listen stdio"] +cwd = "/tmp" +[environments.env] +CODEX_LOG = "debug" +"#, + ) + .expect("write environments.toml"); + + let environments = load_environments_toml(&path) + .expect("environments.toml") + .expect("environments.toml should exist"); + + assert_eq!(environments.default.as_deref(), Some("ssh-dev")); + assert_eq!(environments.include_local, Some(false)); + assert_eq!(environments.environments.len(), 2); + assert_eq!( + environments.environments[0], + EnvironmentToml { + id: "devbox".to_string(), + url: Some("ws://127.0.0.1:4512".to_string()), + connect_timeout_sec: Some(Duration::from_secs(12)), + initialize_timeout_sec: Some(Duration::from_secs(34)), + ..Default::default() + } + ); + assert_eq!( + environments.environments[1], + EnvironmentToml { + id: "ssh-dev".to_string(), + program: Some("ssh".to_string()), + args: Some(vec![ + "dev".to_string(), + "codex exec-server --listen stdio".to_string(), + ]), + env: Some(HashMap::from([( + "CODEX_LOG".to_string(), + "debug".to_string(), + )])), + cwd: Some(PathBuf::from("/tmp")), + ..Default::default() + } + ); + } + + #[test] + fn load_environments_toml_rejects_unknown_fields() { + let codex_home = tempdir().expect("tempdir"); + let cases = [ + ("unknown = true\n", "unknown field `unknown`"), + ( + r#" +[[environments]] +id = "devbox" +url = "ws://127.0.0.1:4512" +unknown = true +"#, + "unknown field `unknown`", + ), + ]; + + for (index, (contents, expected)) in cases.into_iter().enumerate() { + let path = codex_home.path().join(format!("environments-{index}.toml")); + std::fs::write(&path, contents).expect("write environments.toml"); + + let err = load_environments_toml(&path).expect_err("unknown field should fail"); + + assert!( + err.to_string().contains(expected), + "expected `{err}` to contain `{expected}`" + ); + } + } + + #[test] + fn toml_provider_rejects_malformed_websocket_url() { + let err = TomlEnvironmentProvider::new(EnvironmentsToml { + default: None, + include_local: None, + environments: vec![EnvironmentToml { + id: "devbox".to_string(), + url: Some("ws://".to_string()), + ..Default::default() + }], + }) + .expect_err("malformed websocket url should fail"); + + assert!( + err.to_string() + .contains("environment url `ws://` is invalid"), + "expected malformed URL error, got `{err}`" + ); + } + + #[tokio::test] + async fn environment_provider_from_codex_home_uses_present_environments_file() { + let codex_home = tempdir().expect("tempdir"); + std::fs::write( + codex_home.path().join(ENVIRONMENTS_TOML_FILE), + r#" +default = "none" +include_local = false +"#, + ) + .expect("write environments.toml"); + + let provider = + environment_provider_from_codex_home(codex_home.path()).expect("environment provider"); + + let snapshot = provider.snapshot().await.expect("environments"); + let environment_ids: Vec<_> = snapshot + .environments + .into_iter() + .map(|(id, _environment)| id) + .collect(); + + assert!(!snapshot.include_local); + assert!(!environment_ids.contains(&LOCAL_ENVIRONMENT_ID.to_string())); + assert_eq!(snapshot.default, EnvironmentDefault::Disabled); + } + + #[tokio::test] + async fn environment_provider_from_codex_home_falls_back_when_file_is_missing() { + let codex_home = tempdir().expect("tempdir"); + + let provider = + environment_provider_from_codex_home(codex_home.path()).expect("environment provider"); + + let snapshot = provider.snapshot().await.expect("environments"); + let environment_ids: Vec<_> = snapshot + .environments + .into_iter() + .map(|(id, _environment)| id) + .collect(); + + assert!(snapshot.include_local); + assert!(!environment_ids.contains(&LOCAL_ENVIRONMENT_ID.to_string())); + assert_eq!( + snapshot.default, + EnvironmentDefault::EnvironmentId(LOCAL_ENVIRONMENT_ID.to_string()) + ); + } +} diff --git a/codex-rs/exec-server/src/file_read.rs b/codex-rs/exec-server/src/file_read.rs new file mode 100644 index 0000000000000000000000000000000000000000..8cc84b67300df2913d25af4458c4ac39c0231bfd --- /dev/null +++ b/codex-rs/exec-server/src/file_read.rs @@ -0,0 +1,128 @@ +use std::collections::HashMap; +use std::fs::File; +use std::io; +use std::sync::Arc; + +use codex_file_system::FILE_READ_CHUNK_SIZE; +use tokio::sync::Mutex; + +const MAX_OPEN_FILE_READS: usize = 128; + +#[derive(Debug, Eq, PartialEq)] +pub(crate) struct FileReadBlock { + pub(crate) bytes: Vec, + pub(crate) eof: bool, +} + +#[derive(Clone, Default)] +pub(crate) struct FileReadHandleManager { + handles: Arc>>>, +} + +impl FileReadHandleManager { + pub(crate) async fn open( + &self, + handle_id: String, + file: tokio::fs::File, + ) -> io::Result { + let file = Arc::new(file.into_std().await); + let mut handles = self.handles.lock().await; + if handles.contains_key(&handle_id) { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("file read handle `{handle_id}` already exists"), + )); + } + if handles.len() >= MAX_OPEN_FILE_READS { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("at most {MAX_OPEN_FILE_READS} file reads may be open per connection"), + )); + } + handles.insert(handle_id.clone(), file); + Ok(handle_id) + } + + pub(crate) async fn read_block( + &self, + handle_id: &str, + offset: u64, + len: usize, + ) -> io::Result { + validate_read_block_len(len)?; + let file = { + let handles = self.handles.lock().await; + handles + .get(handle_id) + .cloned() + .ok_or_else(|| unknown_handle_error(handle_id))? + }; + let result = + match tokio::task::spawn_blocking(move || read_block_at(&file, offset, len)).await { + Ok(result) => result, + Err(error) => Err(io::Error::other(format!( + "file read task stopped unexpectedly: {error}" + ))), + }; + if result.is_err() { + self.close(handle_id).await; + } + result + } + + pub(crate) async fn close(&self, handle_id: &str) { + self.handles.lock().await.remove(handle_id); + } + + pub(crate) async fn close_all(&self) { + self.handles.lock().await.clear(); + } +} + +fn read_block_at(file: &File, offset: u64, len: usize) -> io::Result { + let mut bytes = vec![0; len]; + let mut bytes_read = 0; + while bytes_read < len { + let read_offset = offset.checked_add(bytes_read as u64).ok_or_else(|| { + io::Error::new(io::ErrorKind::InvalidInput, "file read offset overflowed") + })?; + match read_file_at(file, &mut bytes[bytes_read..], read_offset) { + Ok(0) => break, + Ok(read) => bytes_read += read, + Err(error) if error.kind() == io::ErrorKind::Interrupted => {} + Err(error) => return Err(error), + } + } + bytes.truncate(bytes_read); + Ok(FileReadBlock { + eof: bytes_read < len, + bytes, + }) +} + +#[cfg(unix)] +fn read_file_at(file: &File, bytes: &mut [u8], offset: u64) -> io::Result { + std::os::unix::fs::FileExt::read_at(file, bytes, offset) +} + +#[cfg(windows)] +fn read_file_at(file: &File, bytes: &mut [u8], offset: u64) -> io::Result { + std::os::windows::fs::FileExt::seek_read(file, bytes, offset) +} + +fn validate_read_block_len(len: usize) -> io::Result<()> { + if !(1..=FILE_READ_CHUNK_SIZE).contains(&len) { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("file read block length must be between 1 and {FILE_READ_CHUNK_SIZE}"), + )); + } + Ok(()) +} + +fn unknown_handle_error(handle_id: &str) -> io::Error { + io::Error::new( + io::ErrorKind::NotFound, + format!("unknown file read handle `{handle_id}`"), + ) +} diff --git a/codex-rs/exec-server/src/forward.rs b/codex-rs/exec-server/src/forward.rs new file mode 100644 index 0000000000000000000000000000000000000000..09ccc5d586fbe05e582dbbed24d39907671cb579 --- /dev/null +++ b/codex-rs/exec-server/src/forward.rs @@ -0,0 +1,195 @@ +use std::time::Duration; + +use bytes::Bytes; +use codex_http_client::HttpClientFactory; +use codex_websocket_client::WebSocketConnector; +use futures::Sink; +use futures::SinkExt; +use futures::StreamExt; +use tokio::time::timeout; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::protocol::WebSocketConfig; +use tokio_tungstenite::tungstenite::protocol::frame::Frame; +use tokio_tungstenite::tungstenite::protocol::frame::coding::Data; +use tokio_tungstenite::tungstenite::protocol::frame::coding::OpCode; +use tokio_util::task::AbortOnDropHandle; +use tracing::warn; + +use crate::ExecServerError; +use crate::ExecServerTelemetry; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::noise_relay::message_framing::MAX_NOISE_JSONRPC_MESSAGE_LEN; +use crate::noise_relay::message_framing::frame_message; +use crate::noise_relay::stream_handler::NoiseOutboundMessage; +use crate::noise_relay::stream_handler::NoiseStreamConnection; +use crate::noise_relay::stream_handler::NoiseStreamHandler; +use crate::telemetry::ConnectionTransport; + +// Existing exec-server listeners accept 64 MiB messages but only 16 MiB frames. +const WEBSOCKET_FRAGMENT_LEN: usize = 8 * 1024 * 1024; +const WEBSOCKET_CLOSE_TIMEOUT: Duration = Duration::from_secs(1); + +async fn send_websocket_message(websocket: &mut S, mut payload: Bytes) -> Result<(), S::Error> +where + S: Sink + Unpin, +{ + if payload.len() <= WEBSOCKET_FRAGMENT_LEN { + return websocket.send(Message::Binary(payload)).await; + } + + let mut opcode = OpCode::Data(Data::Binary); + while !payload.is_empty() { + let chunk = payload.split_to(payload.len().min(WEBSOCKET_FRAGMENT_LEN)); + let is_final = payload.is_empty(); + websocket + .send(Message::Frame(Frame::message(chunk, opcode, is_final))) + .await?; + opcode = OpCode::Data(Data::Continue); + } + Ok(()) +} + +/// Copies authenticated remote messages to an independently owned executor. +#[derive(Clone)] +pub(crate) struct Forwarder { + websocket_url: String, + connector: WebSocketConnector, + telemetry: ExecServerTelemetry, +} + +impl Forwarder { + pub(crate) fn new( + websocket_url: String, + http_client_factory: &HttpClientFactory, + telemetry: ExecServerTelemetry, + ) -> Result { + let url = url::Url::parse(&websocket_url) + .map_err(|error| ExecServerError::WebSocketConfiguration(error.to_string()))?; + if !matches!(url.scheme(), "ws" | "wss") || url.host_str().is_none() { + return Err(ExecServerError::WebSocketConfiguration( + "forward destination must be a ws:// or wss:// URL".to_string(), + )); + } + let connector = WebSocketConnector::new(http_client_factory) + .map_err(|error| ExecServerError::WebSocketConfiguration(error.to_string()))? + .with_tcp_nodelay(); + Ok(Self { + websocket_url, + connector, + telemetry, + }) + } + + pub(crate) async fn run_connection(self, mut remote: NoiseStreamConnection) { + let mut writer_task = AbortOnDropHandle::new(remote.writer_task); + let _metrics = self + .telemetry + .connection_started(ConnectionTransport::Relay); + let connect = async { + let request = self.websocket_url.as_str().into_client_request()?; + self.connector + .connect( + request, + WebSocketConfig::default() + .max_message_size(Some(MAX_NOISE_JSONRPC_MESSAGE_LEN)) + .max_frame_size(Some(MAX_NOISE_JSONRPC_MESSAGE_LEN)), + ) + .await + }; + let connected = tokio::select! { + biased; + _ = remote.disconnected_rx.wait_for(|disconnected| *disconnected) => None, + result = timeout(DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, connect) => { + match result { + Ok(Ok((websocket, _))) => Some(websocket), + Ok(Err(_)) => { + warn!("failed to connect to forwarded exec-server"); + None + } + Err(_) => { + warn!("timed out connecting to forwarded exec-server"); + None + } + } + } + }; + let drain_outgoing = if let Some(websocket) = connected { + let (mut destination_tx, mut destination_rx) = websocket.split(); + let to_destination = async { + while let Some(payload) = remote.incoming_rx.recv().await { + if send_websocket_message(&mut destination_tx, payload) + .await + .is_err() + { + break; + } + } + }; + let from_destination = async { + while let Some(Ok(message)) = destination_rx.next().await { + let payload = match message { + Message::Text(_) | Message::Binary(_) => message.into_data(), + Message::Close(_) => return true, + Message::Ping(_) | Message::Pong(_) | Message::Frame(_) => continue, + }; + if remote.outgoing_tx.send(payload).await.is_err() { + break; + } + } + false + }; + let (drain_outgoing, received_close) = tokio::select! { + _ = remote.disconnected_rx.wait_for(|disconnected| *disconnected) => (false, false), + _ = to_destination => (false, false), + received_close = from_destination => (true, received_close), + }; + if received_close && let Ok(mut websocket) = destination_tx.reunite(destination_rx) { + // Reuniting drops any canceled application send before flushing + // Tungstenite's automatically queued Close acknowledgement. + tokio::select! { + _ = timeout(WEBSOCKET_CLOSE_TIMEOUT, websocket.flush()) => {}, + _ = remote.disconnected_rx.wait_for(|disconnected| *disconnected) => {}, + } + } + drain_outgoing + } else { + false + }; + // Preserve messages received before the destination's Close. Never wait + // for a dead remote, and let owner shutdown interrupt a blocked drain. + drop(remote.outgoing_tx); + if drain_outgoing { + tokio::select! { + _ = &mut writer_task => return, + _ = remote.disconnected_rx.wait_for(|disconnected| *disconnected) => {}, + } + } + writer_task.abort(); + let _ = writer_task.await; + } +} + +impl NoiseStreamHandler for Forwarder { + type Incoming = Bytes; + type Outgoing = Bytes; + + fn decode(payload: Bytes) -> Result { + Ok(payload) + } + + fn encode(payload: Bytes) -> Result { + Ok(NoiseOutboundMessage { + framed: frame_message(&payload)?, + trace: None, + }) + } + + async fn run_connection(self, connection: NoiseStreamConnection) { + Forwarder::run_connection(self, connection).await; + } +} + +#[cfg(test)] +#[path = "forward_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/forward_tests.rs b/codex-rs/exec-server/src/forward_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..069e5d764cda2a7d2b9aecc2274bd168157c1697 --- /dev/null +++ b/codex-rs/exec-server/src/forward_tests.rs @@ -0,0 +1,89 @@ +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use bytes::Bytes; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use futures::SinkExt; +use futures::StreamExt; +use pretty_assertions::assert_eq; +use tokio::io::AsyncReadExt; +use tokio::net::TcpListener; +use tokio::sync::mpsc; +use tokio::sync::watch; +use tokio::time::timeout; +use tokio_tungstenite::accept_async; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::protocol::CloseFrame; +use tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode; + +use super::Forwarder; +use crate::ExecServerTelemetry; +use crate::noise_relay::stream_handler::NoiseStreamConnection; + +#[tokio::test] +async fn transport_disconnect_cancels_an_unfinished_websocket_handshake() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let forwarder = Forwarder::new( + format!("ws://{}", listener.local_addr()?), + &HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ExecServerTelemetry::default(), + )?; + let (_incoming, incoming_rx) = mpsc::channel(1); + let (outgoing_tx, _outgoing) = mpsc::channel(1); + let (disconnected, disconnected_rx) = watch::channel(false); + let connection = NoiseStreamConnection { + incoming_rx, + outgoing_tx, + disconnected_rx, + writer_task: tokio::spawn(async {}), + executor_registration: None, + }; + let task = tokio::spawn(forwarder.run_connection(connection)); + let deadline = Duration::from_secs(5); + let (mut socket, _) = timeout(deadline, listener.accept()).await??; + disconnected.send(true)?; + timeout(deadline, task).await??; + timeout(deadline, socket.read_to_end(&mut Vec::new())).await??; + Ok(()) +} + +#[tokio::test] +async fn destination_close_is_acknowledged() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let forwarder = Forwarder::new( + format!("ws://{}", listener.local_addr()?), + &HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ExecServerTelemetry::default(), + )?; + let (_incoming, incoming_rx) = mpsc::channel(1); + let (outgoing_tx, mut outgoing) = mpsc::channel(1); + let (_disconnected, disconnected_rx) = watch::channel(false); + let task = tokio::spawn(forwarder.run_connection(NoiseStreamConnection { + incoming_rx, + outgoing_tx, + disconnected_rx, + writer_task: tokio::spawn(async {}), + executor_registration: None, + })); + let deadline = Duration::from_secs(5); + let (socket, _) = timeout(deadline, listener.accept()).await??; + let mut socket = timeout(deadline, accept_async(socket)).await??; + let response = Bytes::from_static(b"final response"); + socket.send(Message::Binary(response.clone())).await?; + let close = CloseFrame { + code: CloseCode::Normal, + reason: "finished".into(), + }; + socket.close(Some(close.clone())).await?; + + let reply = timeout(deadline, socket.next()) + .await? + .context("destination should receive the Close acknowledgement")??; + assert_eq!(reply, Message::Close(Some(close))); + timeout(deadline, task).await??; + assert_eq!(outgoing.recv().await, Some(response)); + assert_eq!(outgoing.recv().await, None); + Ok(()) +} diff --git a/codex-rs/exec-server/src/fs_helper.rs b/codex-rs/exec-server/src/fs_helper.rs new file mode 100644 index 0000000000000000000000000000000000000000..ad6eacf93cfda03b223c197f7c79c2dac1ec0fe4 --- /dev/null +++ b/codex-rs/exec-server/src/fs_helper.rs @@ -0,0 +1,467 @@ +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD; +use codex_exec_server_protocol::JSONRPCErrorError; +use serde::Deserialize; +use serde::Serialize; +use tokio::io; + +use crate::CapabilityRootsDiscoverParams; +use crate::CapabilityRootsDiscoverResponse; +use crate::CopyOptions; +use crate::CreateDirectoryOptions; +use crate::ExecutorFileSystem; +use crate::GetMetadataOptions; +use crate::ReadFileOptions; +use crate::RemoveOptions; +use crate::WriteFileOptions; +use crate::local_file_system::DirectFileSystem; +use crate::protocol::CAPABILITY_ROOTS_DISCOVER_METHOD; +use crate::protocol::FS_CANONICALIZE_METHOD; +use crate::protocol::FS_COPY_METHOD; +use crate::protocol::FS_CREATE_DIRECTORY_METHOD; +use crate::protocol::FS_GET_METADATA_METHOD; +use crate::protocol::FS_OPEN_METHOD; +use crate::protocol::FS_READ_DIRECTORY_METHOD; +use crate::protocol::FS_READ_FILE_METHOD; +use crate::protocol::FS_REMOVE_METHOD; +use crate::protocol::FS_WALK_METHOD; +use crate::protocol::FS_WRITE_FILE_METHOD; +use crate::protocol::FsCanonicalizeParams; +use crate::protocol::FsCanonicalizeResponse; +use crate::protocol::FsCopyParams; +use crate::protocol::FsCopyResponse; +use crate::protocol::FsCreateDirectoryParams; +use crate::protocol::FsCreateDirectoryResponse; +use crate::protocol::FsGetMetadataParams; +use crate::protocol::FsGetMetadataResponse; +use crate::protocol::FsReadDirectoryEntry; +use crate::protocol::FsReadDirectoryParams; +use crate::protocol::FsReadDirectoryResponse; +use crate::protocol::FsReadFileParams; +use crate::protocol::FsReadFileResponse; +use crate::protocol::FsRemoveParams; +use crate::protocol::FsRemoveResponse; +use crate::protocol::FsWalkParams; +use crate::protocol::FsWalkResponse; +use crate::protocol::FsWriteFileParams; +use crate::protocol::FsWriteFileResponse; +use crate::rpc::internal_error; +use crate::rpc::invalid_request; +use crate::rpc::not_found; + +pub const CODEX_FS_HELPER_ARG1: &str = "--codex-run-as-fs-helper"; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "operation", content = "params")] +pub(crate) enum FsHelperRequest { + #[serde(rename = "capabilityRoots/discoverV1")] + DiscoverCapabilityRoots(CapabilityRootsDiscoverParams), + #[serde(rename = "fs/open")] + Open(FsReadFileParams), + #[serde(rename = "fs/readFile")] + ReadFile(FsReadFileParams), + #[serde(rename = "fs/writeFile")] + WriteFile(FsWriteFileParams), + #[serde(rename = "fs/createDirectory")] + CreateDirectory(FsCreateDirectoryParams), + #[serde(rename = "fs/getMetadata")] + GetMetadata(FsGetMetadataParams), + #[serde(rename = "fs/canonicalize")] + Canonicalize(FsCanonicalizeParams), + #[serde(rename = "fs/readDirectory")] + ReadDirectory(FsReadDirectoryParams), + #[serde(rename = "fs/walk")] + Walk(FsWalkParams), + #[serde(rename = "fs/remove")] + Remove(FsRemoveParams), + #[serde(rename = "fs/copy")] + Copy(FsCopyParams), +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(tag = "status", content = "payload", rename_all = "camelCase")] +pub(crate) enum FsHelperResponse { + Ok(FsHelperPayload), + Error(JSONRPCErrorError), +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct FsHelperOpenResponse { + // Windows duplicates the handle from the helper process. + #[cfg(windows)] + pub(crate) process_id: u32, + // Unix passes the fd directly instead. + #[cfg(windows)] + pub(crate) file_handle: u64, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "operation", content = "response")] +pub(crate) enum FsHelperPayload { + #[serde(rename = "capabilityRoots/discoverV1")] + DiscoverCapabilityRoots(CapabilityRootsDiscoverResponse), + #[serde(rename = "fs/open")] + Open(FsHelperOpenResponse), + #[serde(rename = "fs/readFile")] + ReadFile(FsReadFileResponse), + #[serde(rename = "fs/writeFile")] + WriteFile(FsWriteFileResponse), + #[serde(rename = "fs/createDirectory")] + CreateDirectory(FsCreateDirectoryResponse), + #[serde(rename = "fs/getMetadata")] + GetMetadata(FsGetMetadataResponse), + #[serde(rename = "fs/canonicalize")] + Canonicalize(FsCanonicalizeResponse), + #[serde(rename = "fs/readDirectory")] + ReadDirectory(FsReadDirectoryResponse), + #[serde(rename = "fs/walk")] + Walk(FsWalkResponse), + #[serde(rename = "fs/remove")] + Remove(FsRemoveResponse), + #[serde(rename = "fs/copy")] + Copy(FsCopyResponse), +} + +impl FsHelperPayload { + fn operation(&self) -> &'static str { + match self { + Self::DiscoverCapabilityRoots(_) => CAPABILITY_ROOTS_DISCOVER_METHOD, + Self::Open(_) => FS_OPEN_METHOD, + Self::ReadFile(_) => FS_READ_FILE_METHOD, + Self::WriteFile(_) => FS_WRITE_FILE_METHOD, + Self::CreateDirectory(_) => FS_CREATE_DIRECTORY_METHOD, + Self::GetMetadata(_) => FS_GET_METADATA_METHOD, + Self::Canonicalize(_) => FS_CANONICALIZE_METHOD, + Self::ReadDirectory(_) => FS_READ_DIRECTORY_METHOD, + Self::Walk(_) => FS_WALK_METHOD, + Self::Remove(_) => FS_REMOVE_METHOD, + Self::Copy(_) => FS_COPY_METHOD, + } + } + + pub(crate) fn expect_capability_roots_discover( + self, + ) -> Result { + match self { + Self::DiscoverCapabilityRoots(response) => Ok(response), + other => Err(unexpected_response( + CAPABILITY_ROOTS_DISCOVER_METHOD, + other.operation(), + )), + } + } + + pub(crate) fn expect_read_file(self) -> Result { + match self { + Self::ReadFile(response) => Ok(response), + other => Err(unexpected_response(FS_READ_FILE_METHOD, other.operation())), + } + } + + pub(crate) fn expect_write_file(self) -> Result { + match self { + Self::WriteFile(response) => Ok(response), + other => Err(unexpected_response(FS_WRITE_FILE_METHOD, other.operation())), + } + } + + pub(crate) fn expect_create_directory( + self, + ) -> Result { + match self { + Self::CreateDirectory(response) => Ok(response), + other => Err(unexpected_response( + FS_CREATE_DIRECTORY_METHOD, + other.operation(), + )), + } + } + + pub(crate) fn expect_get_metadata(self) -> Result { + match self { + Self::GetMetadata(response) => Ok(response), + other => Err(unexpected_response( + FS_GET_METADATA_METHOD, + other.operation(), + )), + } + } + + pub(crate) fn expect_canonicalize(self) -> Result { + match self { + Self::Canonicalize(response) => Ok(response), + other => Err(unexpected_response( + FS_CANONICALIZE_METHOD, + other.operation(), + )), + } + } + + pub(crate) fn expect_read_directory( + self, + ) -> Result { + match self { + Self::ReadDirectory(response) => Ok(response), + other => Err(unexpected_response( + FS_READ_DIRECTORY_METHOD, + other.operation(), + )), + } + } + + pub(crate) fn expect_walk(self) -> Result { + match self { + Self::Walk(response) => Ok(response), + other => Err(unexpected_response(FS_WALK_METHOD, other.operation())), + } + } + + pub(crate) fn expect_remove(self) -> Result { + match self { + Self::Remove(response) => Ok(response), + other => Err(unexpected_response(FS_REMOVE_METHOD, other.operation())), + } + } + + pub(crate) fn expect_copy(self) -> Result { + match self { + Self::Copy(response) => Ok(response), + other => Err(unexpected_response(FS_COPY_METHOD, other.operation())), + } + } +} + +fn unexpected_response(expected: &str, actual: &str) -> JSONRPCErrorError { + internal_error(format!( + "unexpected fs sandbox helper response: expected {expected}, got {actual}" + )) +} + +pub(crate) async fn run_direct_request( + request: FsHelperRequest, +) -> Result { + let file_system = DirectFileSystem; + match request { + FsHelperRequest::DiscoverCapabilityRoots(params) => { + let response = crate::discover_capability_roots(&file_system, params) + .await + .map_err(|error| invalid_request(error.to_string()))?; + Ok(FsHelperPayload::DiscoverCapabilityRoots(response)) + } + FsHelperRequest::Open(_) => Err(invalid_request( + "opening a file requires descriptor handoff".to_string(), + )), + FsHelperRequest::ReadFile(params) => { + let data = file_system + .read_file( + ¶ms.path, + ReadFileOptions { + follow_symlinks: params.follow_symlinks.unwrap_or(true), + }, + /*sandbox*/ None, + ) + .await + .map_err(map_fs_error)?; + Ok(FsHelperPayload::ReadFile(FsReadFileResponse { + data_base64: STANDARD.encode(data), + })) + } + FsHelperRequest::WriteFile(params) => { + let bytes = STANDARD.decode(params.data_base64).map_err(|err| { + invalid_request(format!( + "{FS_WRITE_FILE_METHOD} requires valid base64 dataBase64: {err}" + )) + })?; + file_system + .write_file( + ¶ms.path, + bytes, + WriteFileOptions { + follow_symlinks: params.follow_symlinks.unwrap_or(true), + }, + /*sandbox*/ None, + ) + .await + .map_err(map_fs_error)?; + Ok(FsHelperPayload::WriteFile(FsWriteFileResponse {})) + } + FsHelperRequest::CreateDirectory(params) => { + file_system + .create_directory( + ¶ms.path, + CreateDirectoryOptions { + recursive: params.recursive.unwrap_or(true), + follow_symlinks: params.follow_symlinks.unwrap_or(true), + }, + /*sandbox*/ None, + ) + .await + .map_err(map_fs_error)?; + Ok(FsHelperPayload::CreateDirectory( + FsCreateDirectoryResponse {}, + )) + } + FsHelperRequest::GetMetadata(params) => { + let metadata = file_system + .get_metadata( + ¶ms.path, + GetMetadataOptions { + follow_symlinks: params.follow_symlinks.unwrap_or(true), + }, + /*sandbox*/ None, + ) + .await + .map_err(map_fs_error)?; + Ok(FsHelperPayload::GetMetadata(FsGetMetadataResponse { + is_directory: metadata.is_directory, + is_file: metadata.is_file, + is_symlink: metadata.is_symlink, + size: metadata.size, + created_at_ms: metadata.created_at_ms, + modified_at_ms: metadata.modified_at_ms, + })) + } + FsHelperRequest::Canonicalize(params) => { + let path = file_system + .canonicalize(¶ms.path, /*sandbox*/ None) + .await + .map_err(map_fs_error)?; + Ok(FsHelperPayload::Canonicalize(FsCanonicalizeResponse { + path, + })) + } + FsHelperRequest::ReadDirectory(params) => { + let entries = file_system + .read_directory(¶ms.path, /*sandbox*/ None) + .await + .map_err(map_fs_error)? + .into_iter() + .map(|entry| FsReadDirectoryEntry { + file_name: entry.file_name, + is_directory: entry.is_directory, + is_file: entry.is_file, + }) + .collect(); + Ok(FsHelperPayload::ReadDirectory(FsReadDirectoryResponse { + entries, + })) + } + FsHelperRequest::Walk(params) => { + let outcome = file_system + .walk(¶ms.path, params.options, /*sandbox*/ None) + .await + .map_err(map_fs_error)?; + Ok(FsHelperPayload::Walk(outcome)) + } + FsHelperRequest::Remove(params) => { + file_system + .remove( + ¶ms.path, + RemoveOptions { + recursive: params.recursive.unwrap_or(true), + force: params.force.unwrap_or(true), + follow_symlinks: params.follow_symlinks.unwrap_or(true), + }, + /*sandbox*/ None, + ) + .await + .map_err(map_fs_error)?; + Ok(FsHelperPayload::Remove(FsRemoveResponse {})) + } + FsHelperRequest::Copy(params) => { + file_system + .copy( + ¶ms.source_path, + ¶ms.destination_path, + CopyOptions { + recursive: params.recursive, + }, + /*sandbox*/ None, + ) + .await + .map_err(map_fs_error)?; + Ok(FsHelperPayload::Copy(FsCopyResponse {})) + } + } +} + +pub(crate) fn map_fs_error(err: io::Error) -> JSONRPCErrorError { + match err.kind() { + io::ErrorKind::NotFound => not_found(err.to_string()), + io::ErrorKind::InvalidInput | io::ErrorKind::PermissionDenied => { + invalid_request(err.to_string()) + } + _ => internal_error(err.to_string()), + } +} + +#[cfg(test)] +mod tests { + use codex_utils_path_uri::PathUri; + use pretty_assertions::assert_eq; + use serde_json::json; + + use super::*; + + #[test] + fn helper_protocol_uses_path_uris() -> serde_json::Result<()> { + let local_path = + PathUri::from_host_native_path(std::env::current_dir().expect("cwd").join("file")) + .expect("path URI"); + let paths = [ + local_path, + PathUri::parse("file://server/share/file").expect("path URI"), + ]; + + for path in paths { + let expected_path = path.to_string(); + + let request = serde_json::to_value(FsHelperRequest::WriteFile(FsWriteFileParams { + path: path.clone(), + data_base64: String::new(), + follow_symlinks: None, + sandbox: None, + }))?; + assert_eq!( + request, + json!({ + "operation": FS_WRITE_FILE_METHOD, + "params": { + "path": expected_path.as_str(), + "dataBase64": "", + "sandbox": null, + }, + }), + ); + let request_path = request["params"]["path"] + .as_str() + .expect("request path should be a string"); + assert_eq!(request_path, expected_path); + assert!(request_path.starts_with("file:")); + + let response = serde_json::to_value(FsHelperResponse::Ok( + FsHelperPayload::Canonicalize(FsCanonicalizeResponse { path }), + ))?; + assert_eq!( + response, + json!({ + "status": "ok", + "payload": { + "operation": FS_CANONICALIZE_METHOD, + "response": { + "path": expected_path.as_str(), + }, + }, + }), + ); + let response_path = response["payload"]["response"]["path"] + .as_str() + .expect("canonicalize response path should be a string"); + assert_eq!(response_path, expected_path); + assert!(response_path.starts_with("file:")); + } + + Ok(()) + } +} diff --git a/codex-rs/exec-server/src/fs_helper_main.rs b/codex-rs/exec-server/src/fs_helper_main.rs new file mode 100644 index 0000000000000000000000000000000000000000..3ca9ecc21e3d9e525a807680f91d8ebe04059a6d --- /dev/null +++ b/codex-rs/exec-server/src/fs_helper_main.rs @@ -0,0 +1,91 @@ +use std::error::Error; + +use tokio::io; +use tokio::io::AsyncBufReadExt; +use tokio::io::AsyncWriteExt; +use tokio::io::BufReader; + +use crate::fs_helper::FsHelperOpenResponse; +use crate::fs_helper::FsHelperPayload; +use crate::fs_helper::FsHelperRequest; +use crate::fs_helper::FsHelperResponse; +use crate::fs_helper::map_fs_error; +use crate::fs_helper::run_direct_request; +use crate::regular_file; + +pub fn main() -> ! { + let exit_code = match tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + { + Ok(runtime) => match runtime.block_on(run_main()) { + Ok(()) => 0, + Err(err) => { + eprintln!("fs sandbox helper failed: {err}"); + 1 + } + }, + Err(err) => { + eprintln!("failed to start fs sandbox helper runtime: {err}"); + 1 + } + }; + std::process::exit(exit_code); +} + +async fn run_main() -> Result<(), Box> { + let mut stdin = BufReader::new(io::stdin()); + let mut input = String::new(); + stdin.read_line(&mut input).await?; + let request: FsHelperRequest = serde_json::from_str(&input)?; + let mut opened_file = None; + let result = match request { + FsHelperRequest::Open(params) => { + let result: io::Result<_> = async { + let path = params.path.to_abs_path()?; + let file = regular_file::open(path.as_path()).await?; + // Unix can hand the opened fd directly to the parent. + #[cfg(unix)] + crate::sandboxed_file_open::transfer_file(&file)?; + let response = FsHelperOpenResponse { + // Windows duplicates from the helper process instead. + #[cfg(windows)] + process_id: std::process::id(), + // The parent needs the raw handle to duplicate it. + #[cfg(windows)] + file_handle: { + use std::os::windows::io::AsRawHandle; + + file.as_raw_handle() as usize as u64 + }, + }; + opened_file = Some(file); + Ok(FsHelperPayload::Open(response)) + } + .await; + result.map_err(map_fs_error) + } + request => run_direct_request(request).await, + }; + let response = match result { + Ok(payload) => FsHelperResponse::Ok(payload), + Err(error) => FsHelperResponse::Error(error), + }; + let mut stdout = io::stdout(); + stdout + .write_all(serde_json::to_string(&response)?.as_bytes()) + .await?; + stdout.write_all(b"\n").await?; + stdout.flush().await?; + + // Keep the Windows handle alive until the parent duplicates it. + #[cfg(windows)] + if opened_file.is_some() { + use tokio::io::AsyncReadExt; + + let mut acknowledgement = Vec::new(); + stdin.read_to_end(&mut acknowledgement).await?; + } + drop(opened_file); + Ok(()) +} diff --git a/codex-rs/exec-server/src/fs_sandbox.rs b/codex-rs/exec-server/src/fs_sandbox.rs new file mode 100644 index 0000000000000000000000000000000000000000..74a20d1a5c0c38b4e8bbf3e5146ff6d213a32d19 --- /dev/null +++ b/codex-rs/exec-server/src/fs_sandbox.rs @@ -0,0 +1,911 @@ +use std::collections::HashMap; +#[cfg(any(windows, test))] +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_protocol::config_types::WindowsSandboxLevel; +use codex_protocol::models::PermissionProfile; +use codex_protocol::permissions::FileSystemAccessMode; +use codex_protocol::permissions::FileSystemPath; +use codex_protocol::permissions::FileSystemSandboxEntry; +use codex_protocol::permissions::FileSystemSandboxPolicy; +use codex_protocol::permissions::FileSystemSpecialPath; +use codex_protocol::permissions::NetworkSandboxPolicy; +use codex_sandboxing::SandboxCommand; +use codex_sandboxing::SandboxDirectSpawnTransformRequest; +use codex_sandboxing::SandboxExecRequest; +use codex_sandboxing::SandboxManager; +use codex_sandboxing::SandboxTransformRequest; +use codex_sandboxing::SandboxType; +use codex_utils_absolute_path::AbsolutePathBuf; +#[cfg(not(target_os = "linux"))] +use codex_utils_absolute_path::canonicalize_preserving_symlinks; +use codex_utils_path_uri::PathUri; +#[cfg(any(windows, test))] +use tokio::io::AsyncBufReadExt; +#[cfg(any(windows, test))] +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::process::Command; + +use crate::ExecServerRuntimePaths; +use crate::FileSystemSandboxContext; +use crate::fs_helper::CODEX_FS_HELPER_ARG1; +use crate::fs_helper::FsHelperPayload; +use crate::fs_helper::FsHelperRequest; +use crate::fs_helper::FsHelperResponse; +use crate::local_file_system::current_sandbox_cwd; +use crate::rpc::internal_error; +use crate::rpc::invalid_request; + +const FS_HELPER_ENV_ALLOWLIST: &[&str] = &["PATH", "TMPDIR", "TMP", "TEMP"]; +#[cfg(any(windows, test))] +const FS_HELPER_EXIT_TIMEOUT: Duration = Duration::from_secs(/*secs*/ 2); +#[cfg(any(windows, test))] +const MAX_FS_HELPER_STDERR_BYTES: u64 = 4096; +#[cfg(debug_assertions)] +const FS_HELPER_BAZEL_BWRAP_ENV_ALLOWLIST: &[&str] = &[ + "CARGO_BIN_EXE_bwrap", + "RUNFILES_DIR", + "RUNFILES_MANIFEST_FILE", + "RUNFILES_MANIFEST_ONLY", + "TEST_SRCDIR", + "TEST_WORKSPACE", +]; + +#[derive(Debug, PartialEq, Eq)] +struct SandboxCwd { + uri: PathUri, + native: AbsolutePathBuf, +} + +#[derive(Clone, Debug)] +pub(crate) struct FileSystemSandboxRunner { + runtime_paths: ExecServerRuntimePaths, + helper_env: HashMap, +} + +impl FileSystemSandboxRunner { + pub(crate) fn new(runtime_paths: ExecServerRuntimePaths) -> Self { + Self { + runtime_paths, + helper_env: helper_env(), + } + } + + #[tracing::instrument(name = "fs.sandbox_request", skip_all)] + pub(crate) async fn run( + &self, + sandbox: &FileSystemSandboxContext, + request: FsHelperRequest, + ) -> Result { + let command = self.sandbox_command(sandbox)?; + let request_json = serde_json::to_vec(&request).map_err(json_error)?; + run_command(command, request_json).await + } + + #[tracing::instrument( + name = "fs.sandbox_prepare", + skip_all, + fields(permission_entries = tracing::field::Empty) + )] + pub(crate) fn sandbox_command( + &self, + sandbox: &FileSystemSandboxContext, + ) -> Result { + let cwd = sandbox_cwd(sandbox)?; + let native_workspace_roots = sandbox + .workspace_roots + .iter() + .map(native_workspace_root) + .collect::, _>>()?; + let workspace_roots = native_workspace_roots.as_slice(); + let native_permissions: PermissionProfile = + sandbox.permissions.clone().try_into().map_err(|err| { + invalid_request(format!("invalid sandbox permission path URI: {err}")) + })?; + let native_permissions = + native_permissions.materialize_project_roots_with_workspace_roots(workspace_roots); + let mut file_system_policy = native_permissions.file_system_sandbox_policy(); + tracing::Span::current().record("permission_entries", file_system_policy.entries.len()); + let helper_read_roots = if sandbox.use_legacy_landlock { + Vec::new() + } else { + helper_read_roots(&self.runtime_paths) + }; + add_helper_runtime_permissions( + &mut file_system_policy, + &helper_read_roots, + cwd.native.as_path(), + ); + // Linux resolves aliases in the sandbox helper. Doing it here also probes + // unrelated permission roots synchronously on the executor's runtime thread. + #[cfg(not(target_os = "linux"))] + normalize_file_system_policy_root_aliases(&mut file_system_policy); + let network_policy = NetworkSandboxPolicy::Restricted; + let permission_profile = PermissionProfile::from_runtime_permissions_with_enforcement( + native_permissions.enforcement(), + &file_system_policy, + network_policy, + ); + self.sandbox_exec_request(&permission_profile, &cwd, workspace_roots, sandbox) + } + + fn sandbox_exec_request( + &self, + permission_profile: &PermissionProfile, + cwd: &SandboxCwd, + workspace_roots: &[AbsolutePathBuf], + sandbox_context: &FileSystemSandboxContext, + ) -> Result { + let helper = &self.runtime_paths.codex_self_exe; + let sandbox_manager = SandboxManager::for_file_system_helpers(); + #[cfg(target_os = "macos")] + let sandbox_manager = sandbox_manager.with_allowed_symlinked_codex_home( + self.runtime_paths.allowed_symlinked_codex_home.clone(), + ); + let (sandbox, windows_sandbox_level) = crate::sandbox_selection::select_sandbox( + &sandbox_manager, + permission_profile, + sandbox_context, + /*has_managed_network_requirements*/ false, + ); + if sandbox == SandboxType::None { + return Err(invalid_request( + "filesystem sandbox cannot be enforced on this executor".to_string(), + )); + } + let command = SandboxCommand { + program: helper.as_path().as_os_str().to_owned(), + args: vec![CODEX_FS_HELPER_ARG1.to_string()], + cwd: cwd.uri.clone(), + env: self.helper_env.clone(), + managed_network: None, + additional_permissions: None, + }; + sandbox_manager + .transform_for_direct_spawn(SandboxDirectSpawnTransformRequest { + workspace_roots, + windows_sandbox_proxy_settings_mode: + codex_sandboxing::WindowsSandboxProxySettingsMode::Preserve, + transform: SandboxTransformRequest { + command, + permissions: permission_profile, + sandbox, + enforce_managed_network: false, + environment_id: None, + network: None, + sandbox_policy_cwd: &cwd.uri, + sandbox_exe: if cfg!(windows) { + Some(self.runtime_paths.codex_self_exe.as_path()) + } else { + self.runtime_paths.codex_linux_sandbox_exe.as_deref() + }, + use_legacy_landlock: sandbox_context.use_legacy_landlock, + windows_sandbox_level: windows_sandbox_level + .unwrap_or(WindowsSandboxLevel::Disabled), + windows_sandbox_private_desktop: sandbox_context + .windows_sandbox_private_desktop, + }, + }) + .map_err(|err| invalid_request(format!("failed to prepare fs sandbox: {err}"))) + } +} + +fn sandbox_cwd(sandbox: &FileSystemSandboxContext) -> Result { + if let Some(uri) = &sandbox.cwd { + return Ok(SandboxCwd { + native: native_sandbox_cwd(uri)?, + uri: uri.clone(), + }); + } + + if sandbox.has_cwd_dependent_permissions() { + return Err(invalid_request( + "file system sandbox context with dynamic permissions requires cwd".to_string(), + )); + } + + let native = AbsolutePathBuf::from_absolute_path(current_sandbox_cwd().map_err(io_error)?) + .map_err(|err| invalid_request(format!("current directory is not absolute: {err}")))?; + let uri = PathUri::from_abs_path(&native); + Ok(SandboxCwd { uri, native }) +} + +fn native_sandbox_cwd(cwd: &PathUri) -> Result { + cwd.to_abs_path() + .map_err(|err| invalid_request(err.to_string())) +} + +fn native_workspace_root(root: &PathUri) -> Result { + root.to_abs_path().map_err(|err| { + invalid_request(format!( + "file system sandbox workspace root is not native to this exec-server host: {err}" + )) + }) +} + +fn helper_read_roots(runtime_paths: &ExecServerRuntimePaths) -> Vec { + let mut roots = vec![runtime_paths.codex_self_exe.clone()]; + if let Some(path) = &runtime_paths.codex_linux_sandbox_exe + && !roots.contains(path) + { + roots.push(path.clone()); + } + roots +} + +fn add_helper_runtime_permissions( + file_system_policy: &mut FileSystemSandboxPolicy, + helper_read_roots: &[AbsolutePathBuf], + cwd: &std::path::Path, +) { + if !file_system_policy.has_full_disk_read_access() { + let minimal_read_entry = FileSystemSandboxEntry::new( + FileSystemPath::Special { + value: FileSystemSpecialPath::Minimal, + }, + FileSystemAccessMode::Read, + ); + if !file_system_policy.entries.contains(&minimal_read_entry) { + file_system_policy.entries.push(minimal_read_entry); + } + } + + for helper_read_root in helper_read_roots { + if file_system_policy.can_read_local_path_with_cwd(helper_read_root.as_path(), cwd) { + continue; + } + + file_system_policy.entries.push(FileSystemSandboxEntry::new( + helper_read_root.clone().into(), + FileSystemAccessMode::Read, + )); + } +} + +#[cfg(not(target_os = "linux"))] +fn normalize_file_system_policy_root_aliases(file_system_policy: &mut FileSystemSandboxPolicy) { + for entry in &mut file_system_policy.entries { + // Alias normalization uses this executor's filesystem; leave foreign + // or opaque PathUris unchanged. + if let FileSystemPath::Path { path } = &mut entry.path + && let Ok(native_path) = path.to_abs_path() + { + *path = normalize_top_level_alias(native_path).into(); + } + } +} + +#[cfg(not(target_os = "linux"))] +fn normalize_top_level_alias(path: AbsolutePathBuf) -> AbsolutePathBuf { + let raw_path = path.to_path_buf(); + for ancestor in raw_path.ancestors() { + if std::fs::symlink_metadata(ancestor).is_err() { + continue; + } + let Ok(normalized_ancestor) = canonicalize_preserving_symlinks(ancestor) else { + continue; + }; + if normalized_ancestor == ancestor { + continue; + } + let Ok(suffix) = raw_path.strip_prefix(ancestor) else { + continue; + }; + if let Ok(normalized_path) = + AbsolutePathBuf::from_absolute_path(normalized_ancestor.join(suffix)) + { + return normalized_path; + } + } + path +} + +fn helper_env() -> HashMap { + helper_env_from_vars(std::env::vars_os()) +} + +fn helper_env_from_vars( + vars: impl IntoIterator, +) -> HashMap { + vars.into_iter() + .filter_map(|(key, value)| { + let key = key.to_string_lossy(); + helper_env_key_is_allowed(&key) + .then(|| (key.into_owned(), value.to_string_lossy().into_owned())) + }) + .collect() +} + +fn helper_env_key_is_allowed(key: &str) -> bool { + FS_HELPER_ENV_ALLOWLIST.contains(&key) + // CoreFoundation consults this before falling back to user lookup during helper startup. + || (cfg!(target_os = "macos") && key == "__CF_USER_TEXT_ENCODING") + || bazel_bwrap_env_key_is_allowed(key) + || (cfg!(windows) && key.eq_ignore_ascii_case("PATH")) +} + +#[cfg(debug_assertions)] +fn bazel_bwrap_env_key_is_allowed(key: &str) -> bool { + option_env!("BAZEL_PACKAGE").is_some() && FS_HELPER_BAZEL_BWRAP_ENV_ALLOWLIST.contains(&key) +} + +#[cfg(not(debug_assertions))] +fn bazel_bwrap_env_key_is_allowed(_key: &str) -> bool { + false +} + +#[tracing::instrument(name = "fs.sandbox_execute", skip_all)] +async fn run_command( + command: SandboxExecRequest, + request_json: Vec, +) -> Result { + let mut child = spawn_command(command, std::process::Stdio::piped())?; + let mut stdin = child + .stdin + .take() + .ok_or_else(|| internal_error("failed to open fs sandbox helper stdin".to_string()))?; + + #[cfg(windows)] + let mut request_json = request_json; + #[cfg(windows)] + request_json.push(b'\n'); + stdin.write_all(&request_json).await.map_err(io_error)?; + + #[cfg(windows)] + let response = { + stdin.flush().await.map_err(io_error)?; + let stdout = child + .stdout + .take() + .ok_or_else(|| internal_error("failed to open fs sandbox helper stdout".to_string()))?; + let stderr = drain_helper_stderr(&mut child); + let response = read_helper_response(stdout).await; + drop(stdin); + reap_helper_after_response(child, stderr).await?; + response? + }; + + #[cfg(not(windows))] + let response = { + stdin.shutdown().await.map_err(io_error)?; + drop(stdin); + wait_for_helper_output(child).await?.stdout + }; + + let response = serde_json::from_slice(&response).map_err(json_error)?; + match response { + FsHelperResponse::Ok(payload) => Ok(payload), + FsHelperResponse::Error(error) => Err(error), + } +} + +#[cfg(any(windows, test))] +pub(crate) async fn read_helper_response( + stdout: impl tokio::io::AsyncRead + Unpin, +) -> Result, JSONRPCErrorError> { + let mut response = Vec::new(); + let bytes_read = tokio::io::BufReader::new(stdout) + .read_until(b'\n', &mut response) + .await + .map_err(io_error)?; + if bytes_read == 0 { + return Err(internal_error( + "fs sandbox helper closed stdout without responding".to_string(), + )); + } + Ok(response) +} + +#[cfg(any(windows, test))] +pub(crate) fn drain_helper_stderr( + child: &mut tokio::process::Child, +) -> tokio::task::JoinHandle, std::io::Error>> { + let stderr_pipe = child.stderr.take(); + tokio::spawn(async move { + let mut stderr = Vec::new(); + if let Some(mut stderr_pipe) = stderr_pipe { + (&mut stderr_pipe) + .take(MAX_FS_HELPER_STDERR_BYTES) + .read_to_end(&mut stderr) + .await?; + tokio::io::copy(&mut stderr_pipe, &mut tokio::io::sink()).await?; + } + Ok::<_, std::io::Error>(stderr) + }) +} + +#[cfg(any(windows, test))] +pub(crate) async fn reap_helper_after_response( + mut child: tokio::process::Child, + stderr: tokio::task::JoinHandle, std::io::Error>>, +) -> Result<(), JSONRPCErrorError> { + let (status, stderr) = match tokio::time::timeout(FS_HELPER_EXIT_TIMEOUT, async { + tokio::try_join!(child.wait(), async { + stderr.await.map_err(std::io::Error::other)? + }) + }) + .await + { + Ok(result) => result.map_err(io_error)?, + Err(_) => { + tokio::time::timeout(FS_HELPER_EXIT_TIMEOUT, child.kill()) + .await + .map_err(|_| { + internal_error("fs sandbox helper did not stop after its response".to_string()) + })? + .map_err(io_error)?; + return Ok(()); + } + }; + if status.success() { + return Ok(()); + } + + Err(internal_error(format!( + "fs sandbox helper failed with status {status}: {stderr}", + stderr = String::from_utf8_lossy(&stderr).trim() + ))) +} + +#[cfg(not(windows))] +pub(crate) async fn wait_for_helper_output( + child: tokio::process::Child, +) -> Result { + let output = child.wait_with_output().await.map_err(io_error)?; + if !output.status.success() { + return Err(internal_error(format!( + "fs sandbox helper failed with status {status}: {stderr}", + status = output.status, + stderr = String::from_utf8_lossy(&output.stderr).trim() + ))); + } + Ok(output) +} + +pub(crate) fn spawn_command( + SandboxExecRequest { + command: argv, + cwd, + mut env, + arg0, + .. + }: SandboxExecRequest, + stdin: std::process::Stdio, +) -> Result { + let Some((program, args)) = argv.split_first() else { + return Err(invalid_request("fs sandbox command was empty".to_string())); + }; + let mut command = Command::new(program); + #[cfg(unix)] + if let Some(arg0) = arg0 { + command.arg0(arg0); + } + #[cfg(not(unix))] + let _ = arg0; + command.args(args); + // TODO(anp): Keep PathUri through the filesystem helper launch boundary. + let cwd = cwd.to_abs_path().map_err(io_error)?; + command.current_dir(cwd.as_path()); + env.retain(|name, _| !codex_protocol::shell_environment::is_non_inheritable_env_var(name)); + command.env_clear(); + command.envs(env); + command.stdin(stdin); + command.stdout(std::process::Stdio::piped()); + command.stderr(std::process::Stdio::piped()); + command.kill_on_drop(true); + // macOS cannot receive passed fds with close-on-exec set atomically. + #[cfg(target_os = "macos")] + // SAFETY: Descriptor cleanup only uses fork-safe system calls. + unsafe { + command.pre_exec(|| { + codex_utils_pty::pty::close_inherited_fds_except(&[]); + Ok(()) + }); + } + command.spawn().map_err(io_error) +} + +pub(crate) fn io_error(err: std::io::Error) -> JSONRPCErrorError { + internal_error(err.to_string()) +} + +fn json_error(err: serde_json::Error) -> JSONRPCErrorError { + internal_error(format!( + "failed to encode or decode fs sandbox helper message: {err}" + )) +} + +#[cfg(test)] +#[path = "fs_sandbox_windows_tests.rs"] +mod windows_tests; + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + use std::ffi::OsString; + + use codex_protocol::models::PermissionProfile; + use codex_protocol::permissions::FileSystemAccessMode; + use codex_protocol::permissions::FileSystemPath; + use codex_protocol::permissions::FileSystemSandboxEntry; + use codex_protocol::permissions::FileSystemSandboxPolicy; + use codex_protocol::permissions::FileSystemSpecialPath; + use codex_protocol::permissions::NetworkSandboxPolicy; + use codex_utils_absolute_path::AbsolutePathBuf; + use codex_utils_path_uri::PathUri; + use pretty_assertions::assert_eq; + + use crate::ExecServerRuntimePaths; + + use super::FileSystemSandboxRunner; + use super::SandboxCwd; + use super::add_helper_runtime_permissions; + use super::helper_env; + use super::helper_env_from_vars; + use super::helper_env_key_is_allowed; + use super::helper_read_roots; + use super::sandbox_cwd; + + #[test] + fn helper_permissions_enable_minimal_reads_for_restricted_profile() { + let cwd = AbsolutePathBuf::from_absolute_path(std::env::temp_dir().as_path()) + .expect("absolute cwd"); + let mut policy = restricted_policy(Vec::new()); + + add_helper_runtime_permissions(&mut policy, /*helper_read_roots*/ &[], cwd.as_path()); + + assert!(policy.include_platform_defaults()); + } + + #[test] + fn helper_permissions_enable_minimal_reads_for_restricted_profile_with_writes() { + let cwd = AbsolutePathBuf::from_absolute_path(std::env::temp_dir().as_path()) + .expect("absolute cwd"); + let mut policy = restricted_policy(vec![path_entry( + cwd.join("writable"), + FileSystemAccessMode::Write, + )]); + + add_helper_runtime_permissions(&mut policy, /*helper_read_roots*/ &[], cwd.as_path()); + + assert!(policy.include_platform_defaults()); + } + + #[test] + fn helper_permissions_preserve_existing_writes() { + let codex_self_exe = std::env::current_exe().expect("current exe"); + let runtime_paths = + ExecServerRuntimePaths::new(codex_self_exe, /*codex_linux_sandbox_exe*/ None) + .expect("runtime paths"); + let cwd = AbsolutePathBuf::from_absolute_path(std::env::temp_dir().as_path()) + .expect("absolute cwd"); + let writable = cwd.join("writable"); + let mut policy = restricted_policy(vec![path_entry( + writable.clone(), + FileSystemAccessMode::Write, + )]); + let readable = runtime_paths.codex_self_exe.clone(); + + add_helper_runtime_permissions( + &mut policy, + &helper_read_roots(&runtime_paths), + cwd.as_path(), + ); + + assert!(policy.can_read_local_path_with_cwd(readable.as_path(), cwd.as_path())); + assert!(policy.can_write_local_path_with_cwd(writable.as_path(), cwd.as_path())); + } + + #[test] + fn helper_env_carries_only_allowlisted_runtime_vars() { + let env = helper_env(); + + let expected = std::env::vars_os() + .filter_map(|(key, value)| { + let key = key.to_string_lossy(); + helper_env_key_is_allowed(&key) + .then(|| (key.into_owned(), value.to_string_lossy().into_owned())) + }) + .collect::>(); + + assert_eq!(env, expected); + } + + #[test] + fn helper_env_preserves_path_for_system_bwrap_discovery_without_leaking_secrets() { + let env = helper_env_from_vars( + [ + ("PATH", "/usr/bin:/bin"), + ("TMPDIR", "/tmp/codex"), + ("TMP", "/tmp"), + ("TEMP", "/tmp"), + ("HOME", "/home/user"), + ("OPENAI_API_KEY", "secret"), + ("HTTPS_PROXY", "http://proxy.example"), + ] + .map(|(key, value)| (OsString::from(key), OsString::from(value))), + ); + + assert_eq!( + env, + HashMap::from([ + ("PATH".to_string(), "/usr/bin:/bin".to_string()), + ("TMPDIR".to_string(), "/tmp/codex".to_string()), + ("TMP".to_string(), "/tmp".to_string()), + ("TEMP".to_string(), "/tmp".to_string()), + ]) + ); + } + + #[cfg(target_os = "macos")] + #[test] + fn helper_env_preserves_corefoundation_text_encoding() { + let env = helper_env_from_vars( + [ + ("__CF_USER_TEXT_ENCODING", "0x1F6:0x0:0x0"), + ("HOME", "/Users/test"), + ] + .map(|(key, value)| (OsString::from(key), OsString::from(value))), + ); + + assert_eq!( + env, + HashMap::from([( + "__CF_USER_TEXT_ENCODING".to_string(), + "0x1F6:0x0:0x0".to_string(), + )]) + ); + } + + #[cfg(windows)] + #[test] + fn helper_env_preserves_windows_path_key_for_system_bwrap_discovery() { + let env = helper_env_from_vars( + [ + ("Path", r"C:\Windows\System32"), + ("PATH_INJECTION", "bad"), + ("OPENAI_API_KEY", "secret"), + ] + .map(|(key, value)| (OsString::from(key), OsString::from(value))), + ); + + assert_eq!( + env, + HashMap::from([("Path".to_string(), r"C:\Windows\System32".to_string())]) + ); + } + + #[test] + fn sandbox_exec_request_carries_helper_env() { + let Some((path_key, path)) = std::env::vars_os().find(|(key, _)| { + let key = key.to_string_lossy(); + key == "PATH" || (cfg!(windows) && key.eq_ignore_ascii_case("PATH")) + }) else { + return; + }; + let path_key = path_key.to_string_lossy().into_owned(); + let path = path.to_string_lossy().into_owned(); + let codex_self_exe = std::env::current_exe().expect("current exe"); + let runtime_paths = + ExecServerRuntimePaths::new(codex_self_exe.clone(), Some(codex_self_exe)) + .expect("runtime paths"); + let runner = FileSystemSandboxRunner::new(runtime_paths); + let native_cwd = AbsolutePathBuf::current_dir().expect("cwd"); + let cwd = PathUri::from_abs_path(&native_cwd); + let file_system_policy = restricted_policy(vec![ + #[cfg(windows)] + special_entry(FileSystemSpecialPath::Root, FileSystemAccessMode::Read), + path_entry(native_cwd.clone(), FileSystemAccessMode::Write), + ]); + let network_policy = NetworkSandboxPolicy::Restricted; + let permission_profile = + PermissionProfile::from_runtime_permissions(&file_system_policy, network_policy); + let sandbox_context = sandbox_context_with_cwd(&file_system_policy, cwd.clone()); + let sandbox_cwd = SandboxCwd { + uri: cwd, + native: native_cwd, + }; + #[cfg(windows)] + let sandbox_context = { + let error = runner + .sandbox_exec_request( + &permission_profile, + &sandbox_cwd, + std::slice::from_ref(&sandbox_cwd.native), + &sandbox_context, + ) + .expect_err("disabled Windows sandbox must not run the helper unsandboxed"); + assert_eq!( + error.message, + "filesystem sandbox cannot be enforced on this executor" + ); + crate::FileSystemSandboxContext { + windows_sandbox_selection: + codex_file_system::WindowsSandboxSelection::RestrictedToken, + ..sandbox_context + } + }; + + let request = runner + .sandbox_exec_request( + &permission_profile, + &sandbox_cwd, + std::slice::from_ref(&sandbox_cwd.native), + &sandbox_context, + ) + .expect("sandbox exec request"); + + assert_eq!(request.env.get(&path_key), Some(&path)); + } + + #[test] + fn sandbox_cwd_uses_context_cwd() { + let native_cwd = AbsolutePathBuf::from_absolute_path(std::env::temp_dir().as_path()) + .expect("absolute cwd"); + let cwd = PathUri::from_abs_path(&native_cwd); + let policy = restricted_policy(vec![special_entry( + FileSystemSpecialPath::project_roots(/*subpath*/ None), + FileSystemAccessMode::Write, + )]); + let sandbox_context = sandbox_context_with_cwd(&policy, cwd.clone()); + + assert_eq!( + sandbox_cwd(&sandbox_context).expect("sandbox cwd"), + SandboxCwd { + uri: cwd, + native: native_cwd + } + ); + } + + #[test] + fn sandbox_cwd_rejects_non_native_context_cwd_without_fallback() { + let cwd = non_native_cwd(); + let policy = restricted_policy(vec![special_entry( + FileSystemSpecialPath::project_roots(/*subpath*/ None), + FileSystemAccessMode::Write, + )]); + let sandbox_context = sandbox_context_with_cwd(&policy, cwd.clone()); + + let err = sandbox_cwd(&sandbox_context).expect_err("non-native cwd should be rejected"); + + assert_eq!( + err, + crate::rpc::invalid_request(format!( + "'{cwd}' is invalid on '{}'", + std::env::consts::OS + )) + ); + } + + #[test] + fn sandbox_cwd_rejects_cwd_dependent_profile_without_context_cwd() { + let policy = FileSystemSandboxPolicy::restricted(vec![FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + access: FileSystemAccessMode::Write, + missing_path_behavior: None, + }]); + let sandbox_context = codex_file_system::FileSystemSandboxContext::from_permission_profile( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + ); + + let err = sandbox_cwd(&sandbox_context).expect_err("missing cwd should be rejected"); + + assert_eq!( + err.message, + "file system sandbox context with dynamic permissions requires cwd" + ); + } + + #[test] + fn helper_permissions_include_only_the_helper_executable() { + let codex_self_exe = std::env::current_exe().expect("current exe"); + let runtime_paths = + ExecServerRuntimePaths::new(codex_self_exe, /*codex_linux_sandbox_exe*/ None) + .expect("runtime paths"); + let cwd = AbsolutePathBuf::from_absolute_path(std::env::temp_dir().as_path()) + .expect("absolute cwd"); + let mut policy = restricted_policy(Vec::new()); + let parent = runtime_paths + .codex_self_exe + .parent() + .expect("current exe parent"); + let sibling = parent.join("credentials.json"); + + add_helper_runtime_permissions( + &mut policy, + &helper_read_roots(&runtime_paths), + cwd.as_path(), + ); + + assert!( + policy.can_read_local_path_with_cwd( + runtime_paths.codex_self_exe.as_path(), + cwd.as_path(), + ) + ); + assert!(!policy.can_read_local_path_with_cwd(parent.as_path(), cwd.as_path())); + assert!(!policy.can_read_local_path_with_cwd(sibling.as_path(), cwd.as_path())); + } + + #[test] + fn helper_permissions_include_only_linux_sandbox_alias_executable() { + let root = tempfile::tempdir().expect("temp dir"); + let codex_self_exe = root.path().join("bin").join("codex"); + let codex_linux_sandbox_exe = root.path().join("aliases").join("codex-linux-sandbox"); + let runtime_paths = + ExecServerRuntimePaths::new(codex_self_exe, Some(codex_linux_sandbox_exe)) + .expect("runtime paths"); + let cwd = AbsolutePathBuf::from_absolute_path(std::env::temp_dir().as_path()) + .expect("absolute cwd"); + let mut policy = restricted_policy(Vec::new()); + let codex_parent = runtime_paths.codex_self_exe.parent().expect("codex parent"); + let alias = runtime_paths + .codex_linux_sandbox_exe + .as_ref() + .expect("linux sandbox alias"); + let alias_parent = alias.parent().expect("alias parent"); + + add_helper_runtime_permissions( + &mut policy, + &helper_read_roots(&runtime_paths), + cwd.as_path(), + ); + + assert!( + policy.can_read_local_path_with_cwd( + runtime_paths.codex_self_exe.as_path(), + cwd.as_path(), + ) + ); + assert!(policy.can_read_local_path_with_cwd(alias.as_path(), cwd.as_path())); + assert!(!policy.can_read_local_path_with_cwd(codex_parent.as_path(), cwd.as_path())); + assert!(!policy.can_read_local_path_with_cwd(alias_parent.as_path(), cwd.as_path())); + } + + fn restricted_policy(entries: Vec) -> FileSystemSandboxPolicy { + FileSystemSandboxPolicy::restricted(entries) + } + + fn sandbox_context_with_cwd( + policy: &FileSystemSandboxPolicy, + cwd: PathUri, + ) -> crate::FileSystemSandboxContext { + codex_file_system::FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(policy, NetworkSandboxPolicy::Restricted), + cwd, + ) + } + + fn non_native_cwd() -> PathUri { + #[cfg(unix)] + let uri = "file://server/share/checkout"; + #[cfg(windows)] + let uri = "file:///usr/local/checkout"; + + PathUri::parse(uri).expect("non-native cwd URI") + } + + fn path_entry(path: AbsolutePathBuf, access: FileSystemAccessMode) -> FileSystemSandboxEntry { + FileSystemSandboxEntry { + path: path.into(), + access, + missing_path_behavior: None, + } + } + + fn special_entry( + value: FileSystemSpecialPath, + access: FileSystemAccessMode, + ) -> FileSystemSandboxEntry { + FileSystemSandboxEntry { + path: FileSystemPath::Special { value }, + access, + missing_path_behavior: None, + } + } +} diff --git a/codex-rs/exec-server/src/fs_sandbox_windows_tests.rs b/codex-rs/exec-server/src/fs_sandbox_windows_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..477f3e5b67154212c109245db01b45d10da28ec7 --- /dev/null +++ b/codex-rs/exec-server/src/fs_sandbox_windows_tests.rs @@ -0,0 +1,261 @@ +#[cfg(windows)] +use std::collections::HashMap; +#[cfg(windows)] +use std::path::Path; +use std::process::Stdio; +use std::time::Duration; + +#[cfg(windows)] +use codex_protocol::config_types::WindowsSandboxLevel; +#[cfg(windows)] +use codex_protocol::models::PermissionProfile; +#[cfg(windows)] +use codex_sandboxing::SandboxExecRequest; +#[cfg(windows)] +use codex_sandboxing::SandboxType; +#[cfg(windows)] +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +#[cfg(windows)] +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::process::Command; + +use super::drain_helper_stderr; +use super::read_helper_response; +use super::reap_helper_after_response; +#[cfg(windows)] +use super::run_command; +#[cfg(windows)] +use crate::fs_helper::FsHelperPayload; +#[cfg(windows)] +use crate::protocol::FsReadFileResponse; +#[cfg(windows)] +use crate::protocol::FsWriteFileResponse; + +#[tokio::test(start_paused = true)] +async fn filesystem_operation_is_not_limited_by_helper_response_deadline() { + let (reader, mut writer) = tokio::io::duplex(/*max_buf_size*/ 256); + let response = tokio::spawn(async move { read_helper_response(reader).await }); + + tokio::task::yield_now().await; + tokio::time::advance(Duration::from_secs(/*secs*/ 31)).await; + tokio::task::yield_now().await; + assert!( + !response.is_finished(), + "a filesystem operation must not time out before the helper responds" + ); + + writer + .write_all(b"completed after the old deadline\n") + .await + .expect("helper response"); + + assert_eq!( + response + .await + .expect("response task") + .expect("unbounded filesystem operation"), + b"completed after the old deadline\n" + ); +} + +#[tokio::test] +async fn noisy_failing_helper_preserves_exit_status_and_bounded_stderr() { + #[cfg(unix)] + let mut command = { + let mut command = Command::new("sh"); + command.arg("-c").arg( + "printf 'expected helper diagnostic' >&2; i=0; while [ \"$i\" -lt 1024 ]; do printf '%0128d' 0 >&2; i=$((i + 1)); done; exit 7", + ); + command + }; + #[cfg(windows)] + let mut command = { + let system_root = std::env::var("SystemRoot").expect("Windows system root"); + let powershell = Path::new(&system_root) + .join("System32") + .join("WindowsPowerShell") + .join("v1.0") + .join("powershell.exe"); + let mut command = Command::new(powershell); + command + .arg("-NoProfile") + .arg("-Command") + .arg("[Console]::Error.Write('expected helper diagnostic' + ('x' * 131072)); exit 7"); + command + }; + command.stdout(Stdio::null()); + command.stderr(Stdio::piped()); + command.kill_on_drop(true); + let mut child = command.spawn().expect("noisy helper process"); + let stderr = drain_helper_stderr(&mut child); + + let error = tokio::time::timeout( + Duration::from_secs(/*secs*/ 8), + reap_helper_after_response(child, stderr), + ) + .await + .expect("helper stderr must be drained during bounded cleanup") + .expect_err("nonzero helper exit should fail after its stderr pipe fills"); + + assert!(error.message.contains('7'), "{}", error.message); + assert!( + error.message.contains("expected helper diagnostic"), + "{}", + error.message + ); + assert!(error.message.len() < 4400, "{}", error.message.len()); +} + +#[tokio::test] +async fn helper_stderr_is_drained_before_the_response() { + #[cfg(unix)] + let mut command = { + let mut command = Command::new("sh"); + command.arg("-c").arg( + "printf 'expected pre-response diagnostic' >&2; i=0; while [ \"$i\" -lt 1024 ]; do printf '%0128d' 0 >&2; i=$((i + 1)); done; printf 'completed after noisy stderr\\n'", + ); + command + }; + #[cfg(windows)] + let mut command = { + let system_root = std::env::var("SystemRoot").expect("Windows system root"); + let powershell = Path::new(&system_root) + .join("System32") + .join("WindowsPowerShell") + .join("v1.0") + .join("powershell.exe"); + let mut command = Command::new(powershell); + command.arg("-NoProfile").arg("-Command").arg( + "[Console]::Error.Write('expected pre-response diagnostic' + ('x' * 131072)); [Console]::Out.WriteLine('completed after noisy stderr')", + ); + command + }; + command.stdout(Stdio::piped()); + command.stderr(Stdio::piped()); + command.kill_on_drop(true); + let mut child = command.spawn().expect("noisy helper process"); + let stdout = child.stdout.take().expect("helper stdout"); + let stderr = drain_helper_stderr(&mut child); + + let response = tokio::time::timeout( + Duration::from_secs(/*secs*/ 2), + read_helper_response(stdout), + ) + .await + .expect("helper stderr must be drained before awaiting its response") + .expect("helper response"); + + assert_eq!(response.trim_ascii_end(), b"completed after noisy stderr"); + reap_helper_after_response(child, stderr) + .await + .expect("noisy helper should be cleaned up after its response"); +} + +#[cfg(windows)] +#[tokio::test] +async fn completed_windows_image_read_does_not_wait_for_a_stuck_helper() { + let directory = tempfile::tempdir().expect("temporary directory"); + let path = directory.path().join("image.png"); + let operations = [ + ( + r#"[IO.File]::WriteAllText($env:CODEX_FS_HELPER_TEST_PATH, 'image') +[Console]::Out.WriteLine('{"status":"ok","payload":{"operation":"fs/writeFile","response":{}}}')"#, + FsHelperPayload::WriteFile(FsWriteFileResponse {}), + ), + ( + r#"$data = [Convert]::ToBase64String([IO.File]::ReadAllBytes($env:CODEX_FS_HELPER_TEST_PATH)) +[Console]::Out.WriteLine('{"status":"ok","payload":{"operation":"fs/readFile","response":{"dataBase64":"' + $data + '"}}}')"#, + FsHelperPayload::ReadFile(FsReadFileResponse { + data_base64: "aW1hZ2U=".to_string(), + }), + ), + ]; + + for (operation, expected) in operations { + let script = format!( + "[Console]::In.ReadLine() | Out-Null\n{operation}\n[Console]::Out.Flush()\n[Threading.Thread]::Sleep(30000)" + ); + let command = powershell_command(&script, &path).expect("PowerShell helper command"); + let result = tokio::time::timeout( + Duration::from_secs(/*secs*/ 8), + run_command(command, b"{}".to_vec()), + ) + .await + .expect("the completed operation must not wait for helper termination") + .expect("helper response"); + + assert_eq!(result, expected); + } + assert_eq!(std::fs::read(&path).expect("created file"), b"image"); +} + +#[cfg(windows)] +#[tokio::test] +async fn duplicated_windows_file_handle_survives_bounded_helper_cleanup() { + let directory = tempfile::tempdir().expect("temporary directory"); + let path = directory.path().join("image.png"); + std::fs::write(&path, b"image").expect("image file"); + let command = powershell_command( + r#"[Console]::In.ReadLine() | Out-Null +$file = [IO.File]::OpenRead($env:CODEX_FS_HELPER_TEST_PATH) +$handle = $file.SafeFileHandle.DangerousGetHandle().ToInt64() +[Console]::Out.WriteLine('{"status":"ok","payload":{"operation":"fs/open","response":{"processId":' + $PID + ',"fileHandle":' + $handle + '}}}') +[Console]::Out.Flush() +[Threading.Thread]::Sleep(30000)"#, + &path, + ) + .expect("PowerShell helper command"); + + let mut file = tokio::time::timeout( + Duration::from_secs(/*secs*/ 8), + crate::sandboxed_file_open::open( + command, + PathUri::from_host_native_path(&path).expect("image path URI"), + ), + ) + .await + .expect("the opened file must not wait for helper termination") + .expect("duplicated file handle"); + let mut data = Vec::new(); + file.read_to_end(&mut data).await.expect("image contents"); + + assert_eq!(data, b"image"); +} + +#[cfg(windows)] +fn powershell_command(script: &str, path: &Path) -> anyhow::Result { + let system_root = std::env::var("SystemRoot")?; + let powershell = Path::new(&system_root) + .join("System32") + .join("WindowsPowerShell") + .join("v1.0") + .join("powershell.exe"); + let cwd = PathUri::from_host_native_path(std::env::current_dir()?)?; + + Ok(SandboxExecRequest { + command: vec![ + powershell.to_string_lossy().into_owned(), + "-NoProfile".to_string(), + "-Command".to_string(), + script.to_string(), + ], + cwd: cwd.clone(), + sandbox_policy_cwd: cwd, + env: HashMap::from([ + ("SystemRoot".to_string(), system_root), + ( + "CODEX_FS_HELPER_TEST_PATH".to_string(), + path.to_string_lossy().into_owned(), + ), + ]), + network: None, + network_environment_id: None, + sandbox: SandboxType::None, + windows_sandbox_level: WindowsSandboxLevel::Disabled, + windows_sandbox_private_desktop: false, + permission_profile: PermissionProfile::Disabled, + arg0: None, + }) +} diff --git a/codex-rs/exec-server/src/lib.rs b/codex-rs/exec-server/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..c01d5368fb2a07db8c35fa019e9678a30ca608cf --- /dev/null +++ b/codex-rs/exec-server/src/lib.rs @@ -0,0 +1,220 @@ +mod arg0_exec_helper; +mod capability_discovery; +mod capability_discovery_cache; +mod client; +mod client_api; +mod client_telemetry; +mod client_transport; +mod connection; +mod environment; +mod environment_bootstrap; +mod environment_config; +mod environment_provider; +mod environment_registry; +mod environment_toml; +mod file_read; +mod forward; +mod fs_helper; +mod fs_helper_main; +mod fs_sandbox; +mod local_file_system; +mod local_process; +mod network_policy_decisions; +mod no_follow; +mod noise_channel; +mod noise_relay; +mod process; +mod process_sandbox; +mod process_telemetry; +mod regular_file; +mod relay; +mod relay_proto; +mod remote; +mod remote_file_system; +mod remote_process; +mod resolved_capability; +mod rpc; +mod rpc_server_requests; +mod runtime_paths; +mod sandbox_selection; +mod sandboxed_file_open; +mod sandboxed_file_system; +mod server; +#[cfg(unix)] +mod shell_snapshot; +mod telemetry; +mod trace_context; +mod websocket_pong_watchdog; + +use codex_exec_server_protocol as protocol; + +/// Process-local opt-in for tying a remote executor to its parent's stdin pipe. +pub const CODEX_EXEC_SERVER_EXIT_ON_STDIN_CLOSE_ENV_VAR: &str = + "CODEX_EXEC_SERVER_EXIT_ON_STDIN_CLOSE"; + +pub use arg0_exec_helper::CODEX_ARG0_EXEC_HELPER_ARG1; +pub use arg0_exec_helper::main as run_arg0_exec_helper_main; +pub use capability_discovery::CapabilityDiscoveryError; +pub use capability_discovery::discover_capability_roots; +pub use capability_discovery_cache::ExecutorCapabilityDiscoveryCache; +pub use client::ExecServerClient; +pub use client::ExecServerError; +pub use client::http_client::HttpResponseBodyStream; +pub use client::http_client::RouteAwareHttpClient; +pub use client_api::ExecServerClientConnectOptions; +pub use client_api::HttpClient; +pub use client_api::NoiseRendezvousConnectArgs; +pub use client_api::NoiseRendezvousConnectBundle; +pub use client_api::NoiseRendezvousConnectProvider; +pub use client_api::RemoteExecServerConnectArgs; +pub use codex_exec_server_protocol::ExecutorCapabilityDiscoverySnapshot; +pub use codex_exec_server_protocol::ProcessId; +pub use codex_file_system::CopyOptions; +pub use codex_file_system::CreateDirectoryOptions; +pub use codex_file_system::ExecutorFileSystem; +pub use codex_file_system::ExecutorFileSystemFuture; +pub use codex_file_system::FILE_READ_CHUNK_SIZE; +pub use codex_file_system::FileMetadata; +pub use codex_file_system::FileSystemReadStream; +pub use codex_file_system::FileSystemResult; +pub use codex_file_system::FileSystemSandboxContext; +pub use codex_file_system::GetMetadataOptions; +pub use codex_file_system::ReadDirectoryEntry; +pub use codex_file_system::ReadFileOptions; +pub use codex_file_system::RemoveOptions; +pub use codex_file_system::WalkEntry; +pub use codex_file_system::WalkEntryKind; +pub use codex_file_system::WalkError; +pub use codex_file_system::WalkOptions; +pub use codex_file_system::WalkOutcome; +pub use codex_file_system::WindowsSandboxSelection; +pub use codex_file_system::WriteFileOptions; +pub use codex_protocol::shell_environment::CODEX_EXEC_SERVER_NOISE_AUTH_TOKEN_ENV_VAR; +pub use environment::CODEX_EXEC_SERVER_NOISE_CHATGPT_ACCOUNT_ID_ENV_VAR; +pub use environment::CODEX_EXEC_SERVER_NOISE_ENVIRONMENT_ID_ENV_VAR; +pub use environment::CODEX_EXEC_SERVER_NOISE_REGISTRY_URL_ENV_VAR; +pub use environment::CODEX_EXEC_SERVER_URL_ENV_VAR; +pub use environment::Environment; +pub use environment::EnvironmentConnectionState; +pub use environment::EnvironmentManager; +pub use environment::EnvironmentObservedStatus; +pub use environment::EnvironmentReadyInfo; +pub use environment::LOCAL_ENVIRONMENT_ID; +pub use environment::MAX_SELECTED_CAPABILITY_ROOTS; +pub use environment::REMOTE_ENVIRONMENT_ID; +pub use environment::RemoteEnvironmentOptions; +pub use environment_bootstrap::PreparedEnvironmentManager; +pub use environment_provider::DefaultEnvironmentProvider; +pub use environment_provider::EnvironmentProvider; +pub use environment_provider::EnvironmentProviderFuture; +pub use environment_registry::EnvironmentRegistryConnectRequest; +pub use environment_registry::EnvironmentRegistryConnectResponse; +pub use environment_registry::EnvironmentRegistryHarnessKeyValidationRequest; +pub use environment_registry::EnvironmentRegistryHarnessKeyValidationResponse; +pub use environment_registry::EnvironmentRegistryRegistrationRequest; +pub use environment_registry::EnvironmentRegistryRegistrationResponse; +pub use fs_helper::CODEX_FS_HELPER_ARG1; +pub use fs_helper_main::main as run_fs_helper_main; +pub use local_file_system::LOCAL_FS; +pub use local_file_system::LocalFileSystem; +pub use noise_channel::NoiseChannelError; +pub use noise_channel::NoiseChannelIdentity; +pub use noise_channel::NoiseChannelPublicKey; +pub use process::ExecBackend; +pub use process::ExecBackendFuture; +pub use process::ExecProcess; +pub use process::ExecProcessEvent; +pub use process::ExecProcessEventReceiver; +pub use process::ExecProcessFuture; +pub use process::StartedExecProcess; +pub use protocol::ByteChunk; +pub use protocol::CAPABILITY_ROOTS_DISCOVER_METHOD; +pub use protocol::CapabilityRootDiscoverRequest; +pub use protocol::CapabilityRootDiscovery; +pub use protocol::CapabilityRootsDiscoverParams; +pub use protocol::CapabilityRootsDiscoverResponse; +pub use protocol::CapabilityTextFile; +pub use protocol::DiscoveredPluginFiles; +pub use protocol::DiscoveredSkillFiles; +pub use protocol::EnvironmentCapabilities; +pub use protocol::EnvironmentConfigLayer; +pub use protocol::EnvironmentConfigLayerStack; +pub use protocol::EnvironmentConfigReadParams; +pub use protocol::EnvironmentConfigReadResponse; +pub use protocol::EnvironmentInfo; +pub use protocol::EnvironmentStatus; +pub use protocol::EnvironmentStatusKind; +pub use protocol::ExecClosedNotification; +pub use protocol::ExecEnvPolicy; +pub use protocol::ExecExitedNotification; +pub use protocol::ExecMetadata; +pub use protocol::ExecOutputDeltaNotification; +pub use protocol::ExecOutputStream; +pub use protocol::ExecParams; +pub use protocol::ExecResponse; +pub use protocol::ExecServerNetworkPolicyDecision; +pub use protocol::ExecServerNetworkPolicyRequest; +pub use protocol::ExecServerNetworkProtocol; +pub use protocol::FsCanonicalizeParams; +pub use protocol::FsCanonicalizeResponse; +pub use protocol::FsCloseParams; +pub use protocol::FsCloseResponse; +pub use protocol::FsCopyParams; +pub use protocol::FsCopyResponse; +pub use protocol::FsCreateDirectoryParams; +pub use protocol::FsCreateDirectoryResponse; +pub use protocol::FsGetMetadataParams; +pub use protocol::FsGetMetadataResponse; +pub use protocol::FsOpenParams; +pub use protocol::FsOpenResponse; +pub use protocol::FsReadBlockParams; +pub use protocol::FsReadBlockResponse; +pub use protocol::FsReadDirectoryEntry; +pub use protocol::FsReadDirectoryParams; +pub use protocol::FsReadDirectoryResponse; +pub use protocol::FsReadFileParams; +pub use protocol::FsReadFileResponse; +pub use protocol::FsRemoveParams; +pub use protocol::FsRemoveResponse; +pub use protocol::FsWalkParams; +pub use protocol::FsWalkResponse; +pub use protocol::FsWriteFileParams; +pub use protocol::FsWriteFileResponse; +pub use protocol::HttpHeader; +pub use protocol::HttpRedirectPolicy; +pub use protocol::HttpRequestBodyDeltaNotification; +pub use protocol::HttpRequestParams; +pub use protocol::HttpRequestResponse; +pub use protocol::InitializeParams; +pub use protocol::InitializeResponse; +pub use protocol::NetworkPolicyRequestParams; +pub use protocol::NetworkPolicyRequestResponse; +pub use protocol::ProcessOutputChunk; +pub use protocol::ProcessSignal; +pub use protocol::ReadParams; +pub use protocol::ReadResponse; +pub use protocol::ShellInfo; +pub use protocol::ShellSnapshotRequest; +pub use protocol::SignalParams; +pub use protocol::SignalResponse; +pub use protocol::TerminateParams; +pub use protocol::TerminateResponse; +pub use protocol::WriteParams; +pub use protocol::WriteResponse; +pub use protocol::WriteStatus; +pub use regular_file::read_sensitive_file_to_string; +pub use remote::RemoteEnvironmentConfig; +pub use remote::RemoteEnvironmentTransport; +pub use remote::run_remote_environment; +pub use remote::run_remote_environment_forward_until_shutdown; +pub use remote::run_remote_environment_until_shutdown; +pub use resolved_capability::ResolvedSelectedCapabilityRoot; +pub use resolved_capability::SelectedCapabilityRootsStatus; +pub use runtime_paths::ExecServerRuntimePaths; +pub use server::ConcurrentRequestLimit; +pub use server::DEFAULT_LISTEN_URL; +pub use server::ExecServerListenUrlParseError; +pub use server::RequestDispatchMode; +pub use server::run_main; +pub use server::run_main_with_telemetry; +pub use telemetry::ExecServerTelemetry; diff --git a/codex-rs/exec-server/src/local_file_system.rs b/codex-rs/exec-server/src/local_file_system.rs new file mode 100644 index 0000000000000000000000000000000000000000..63048da72ad3bcbc69b7a1dd8d75b4a1077ded5f --- /dev/null +++ b/codex-rs/exec-server/src/local_file_system.rs @@ -0,0 +1,1433 @@ +use codex_file_system::MAX_WALK_DEPTH; +use codex_file_system::MAX_WALK_DIRECTORIES; +use codex_file_system::MAX_WALK_ENTRIES; +use codex_file_system::MAX_WALK_RESPONSE_BYTES; +use codex_file_system::WALK_RESPONSE_ITEM_OVERHEAD_BYTES; +use codex_utils_absolute_path::AbsolutePathBuf; +use codex_utils_path_uri::PathUri; +use std::collections::HashSet; +use std::collections::VecDeque; +use std::path::Path; +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::LazyLock; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; +use tokio::io; +use tokio::io::AsyncReadExt; +use tokio_util::io::ReaderStream; +use tokio_util::sync::CancellationToken; + +use crate::CopyOptions; +use crate::CreateDirectoryOptions; +use crate::ExecServerRuntimePaths; +use crate::ExecutorFileSystem; +use crate::ExecutorFileSystemFuture; +use crate::FILE_READ_CHUNK_SIZE; +use crate::FileMetadata; +use crate::FileSystemReadStream; +use crate::FileSystemResult; +use crate::FileSystemSandboxContext; +use crate::GetMetadataOptions; +use crate::ReadDirectoryEntry; +use crate::ReadFileOptions; +use crate::RemoveOptions; +use crate::WalkEntry; +use crate::WalkEntryKind; +use crate::WalkError; +use crate::WalkOptions; +use crate::WalkOutcome; +use crate::WriteFileOptions; +use crate::no_follow; +use crate::regular_file; +use crate::sandboxed_file_system::SandboxedFileSystem; + +const MAX_READ_FILE_BYTES: u64 = 512 * 1024 * 1024; + +fn file_too_large_error() -> io::Error { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("file is too large to read: limit is {MAX_READ_FILE_BYTES} bytes"), + ) +} + +pub static LOCAL_FS: LazyLock> = + LazyLock::new(|| -> Arc { Arc::new(LocalFileSystem::unsandboxed()) }); + +#[derive(Clone, Default)] +pub(crate) struct DirectFileSystem; + +#[derive(Clone, Default)] +pub(crate) struct UnsandboxedFileSystem { + file_system: DirectFileSystem, +} + +#[derive(Clone, Default)] +pub struct LocalFileSystem { + unsandboxed: UnsandboxedFileSystem, + sandboxed: Option, +} + +impl LocalFileSystem { + pub fn unsandboxed() -> Self { + Self { + unsandboxed: UnsandboxedFileSystem::default(), + sandboxed: None, + } + } + + pub fn with_runtime_paths(runtime_paths: ExecServerRuntimePaths) -> Self { + Self { + unsandboxed: UnsandboxedFileSystem::default(), + sandboxed: Some(SandboxedFileSystem::new(runtime_paths)), + } + } + + pub(crate) fn sandboxed(&self) -> io::Result<&SandboxedFileSystem> { + self.sandboxed.as_ref().ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidInput, + "sandboxed filesystem operations require configured runtime paths", + ) + }) + } + + fn file_system_for<'a>( + &'a self, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> io::Result<( + &'a dyn ExecutorFileSystem, + Option<&'a FileSystemSandboxContext>, + )> { + if sandbox.is_some_and(FileSystemSandboxContext::should_run_in_sandbox) { + Ok((self.sandboxed()?, sandbox)) + } else { + Ok((&self.unsandboxed, sandbox)) + } + } +} + +impl LocalFileSystem { + pub(crate) async fn open_file_for_read( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + if sandbox.is_some_and(FileSystemSandboxContext::should_run_in_sandbox) { + return self.sandboxed()?.open_file_for_read(path, sandbox).await; + } + self.unsandboxed.open_file_for_read(path, sandbox).await + } + + async fn canonicalize( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + let (file_system, sandbox) = self.file_system_for(sandbox)?; + file_system.canonicalize(path, sandbox).await + } + + #[tracing::instrument( + name = "fs.read_file", + skip_all, + fields(sandboxed = sandbox.is_some_and(FileSystemSandboxContext::should_run_in_sandbox)) + )] + async fn read_file( + &self, + path: &PathUri, + options: ReadFileOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult> { + let (file_system, sandbox) = self.file_system_for(sandbox)?; + file_system.read_file(path, options, sandbox).await + } + + async fn read_file_stream( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + let (file_system, sandbox) = self.file_system_for(sandbox)?; + file_system.read_file_stream(path, sandbox).await + } + + async fn write_file( + &self, + path: &PathUri, + contents: Vec, + options: WriteFileOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + let (file_system, sandbox) = self.file_system_for(sandbox)?; + file_system + .write_file(path, contents, options, sandbox) + .await + } + + async fn create_directory( + &self, + path: &PathUri, + options: CreateDirectoryOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + let (file_system, sandbox) = self.file_system_for(sandbox)?; + file_system.create_directory(path, options, sandbox).await + } + + #[tracing::instrument( + name = "fs.get_metadata", + skip_all, + fields(sandboxed = sandbox.is_some_and(FileSystemSandboxContext::should_run_in_sandbox)) + )] + async fn get_metadata( + &self, + path: &PathUri, + options: GetMetadataOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + let (file_system, sandbox) = self.file_system_for(sandbox)?; + file_system.get_metadata(path, options, sandbox).await + } + + async fn read_directory( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult> { + let (file_system, sandbox) = self.file_system_for(sandbox)?; + file_system.read_directory(path, sandbox).await + } + + async fn walk( + &self, + path: &PathUri, + options: WalkOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + let (file_system, sandbox) = self.file_system_for(sandbox)?; + file_system.walk(path, options, sandbox).await + } + + async fn remove( + &self, + path: &PathUri, + options: RemoveOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + let (file_system, sandbox) = self.file_system_for(sandbox)?; + file_system.remove(path, options, sandbox).await + } + + async fn copy( + &self, + source_path: &PathUri, + destination_path: &PathUri, + options: CopyOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + let (file_system, sandbox) = self.file_system_for(sandbox)?; + file_system + .copy(source_path, destination_path, options, sandbox) + .await + } +} + +impl ExecutorFileSystem for LocalFileSystem { + fn canonicalize<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, PathUri> { + Box::pin(LocalFileSystem::canonicalize(self, path, sandbox)) + } + + fn read_file<'a>( + &'a self, + path: &'a PathUri, + options: ReadFileOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, Vec> { + Box::pin(LocalFileSystem::read_file(self, path, options, sandbox)) + } + + fn read_file_stream<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, FileSystemReadStream> { + Box::pin(LocalFileSystem::read_file_stream(self, path, sandbox)) + } + + fn write_file<'a>( + &'a self, + path: &'a PathUri, + contents: Vec, + options: WriteFileOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(LocalFileSystem::write_file( + self, path, contents, options, sandbox, + )) + } + + fn create_directory<'a>( + &'a self, + path: &'a PathUri, + options: CreateDirectoryOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(LocalFileSystem::create_directory( + self, path, options, sandbox, + )) + } + + fn get_metadata<'a>( + &'a self, + path: &'a PathUri, + options: GetMetadataOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, FileMetadata> { + Box::pin(LocalFileSystem::get_metadata(self, path, options, sandbox)) + } + + fn read_directory<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, Vec> { + Box::pin(LocalFileSystem::read_directory(self, path, sandbox)) + } + + fn walk<'a>( + &'a self, + path: &'a PathUri, + options: WalkOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, WalkOutcome> { + Box::pin(LocalFileSystem::walk(self, path, options, sandbox)) + } + + fn remove<'a>( + &'a self, + path: &'a PathUri, + options: RemoveOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(LocalFileSystem::remove(self, path, options, sandbox)) + } + + fn copy<'a>( + &'a self, + source_path: &'a PathUri, + destination_path: &'a PathUri, + options: CopyOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(LocalFileSystem::copy( + self, + source_path, + destination_path, + options, + sandbox, + )) + } +} + +impl UnsandboxedFileSystem { + async fn open_file_for_read( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + reject_platform_sandbox_context(sandbox)?; + self.file_system + .open_file_for_read(path, /*sandbox*/ None) + .await + } + + async fn canonicalize( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + reject_platform_sandbox_context(sandbox)?; + self.file_system.canonicalize(path, /*sandbox*/ None).await + } + + async fn read_file( + &self, + path: &PathUri, + options: ReadFileOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult> { + reject_platform_sandbox_context(sandbox)?; + self.file_system + .read_file(path, options, /*sandbox*/ None) + .await + } + + async fn read_file_stream( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + reject_platform_sandbox_context(sandbox)?; + self.file_system + .read_file_stream(path, /*sandbox*/ None) + .await + } + + async fn write_file( + &self, + path: &PathUri, + contents: Vec, + options: WriteFileOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + reject_platform_sandbox_context(sandbox)?; + self.file_system + .write_file(path, contents, options, /*sandbox*/ None) + .await + } + + async fn create_directory( + &self, + path: &PathUri, + options: CreateDirectoryOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + reject_platform_sandbox_context(sandbox)?; + self.file_system + .create_directory(path, options, /*sandbox*/ None) + .await + } + + async fn get_metadata( + &self, + path: &PathUri, + options: GetMetadataOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + reject_platform_sandbox_context(sandbox)?; + self.file_system + .get_metadata(path, options, /*sandbox*/ None) + .await + } + + async fn read_directory( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult> { + reject_platform_sandbox_context(sandbox)?; + self.file_system + .read_directory(path, /*sandbox*/ None) + .await + } + + async fn remove( + &self, + path: &PathUri, + options: RemoveOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + reject_platform_sandbox_context(sandbox)?; + self.file_system + .remove(path, options, /*sandbox*/ None) + .await + } + + async fn copy( + &self, + source_path: &PathUri, + destination_path: &PathUri, + options: CopyOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + reject_platform_sandbox_context(sandbox)?; + self.file_system + .copy( + source_path, + destination_path, + options, + /*sandbox*/ None, + ) + .await + } +} + +impl ExecutorFileSystem for UnsandboxedFileSystem { + fn canonicalize<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, PathUri> { + Box::pin(UnsandboxedFileSystem::canonicalize(self, path, sandbox)) + } + + fn read_file<'a>( + &'a self, + path: &'a PathUri, + options: ReadFileOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, Vec> { + Box::pin(UnsandboxedFileSystem::read_file( + self, path, options, sandbox, + )) + } + + fn read_file_stream<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, FileSystemReadStream> { + Box::pin(UnsandboxedFileSystem::read_file_stream(self, path, sandbox)) + } + + fn write_file<'a>( + &'a self, + path: &'a PathUri, + contents: Vec, + options: WriteFileOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(UnsandboxedFileSystem::write_file( + self, path, contents, options, sandbox, + )) + } + + fn create_directory<'a>( + &'a self, + path: &'a PathUri, + options: CreateDirectoryOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(UnsandboxedFileSystem::create_directory( + self, path, options, sandbox, + )) + } + + fn get_metadata<'a>( + &'a self, + path: &'a PathUri, + options: GetMetadataOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, FileMetadata> { + Box::pin(UnsandboxedFileSystem::get_metadata( + self, path, options, sandbox, + )) + } + + fn read_directory<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, Vec> { + Box::pin(UnsandboxedFileSystem::read_directory(self, path, sandbox)) + } + + fn walk<'a>( + &'a self, + path: &'a PathUri, + options: WalkOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, WalkOutcome> { + Box::pin(async move { + reject_platform_sandbox_context(sandbox)?; + self.file_system.walk(path, options, /*sandbox*/ None).await + }) + } + + fn remove<'a>( + &'a self, + path: &'a PathUri, + options: RemoveOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(UnsandboxedFileSystem::remove(self, path, options, sandbox)) + } + + fn copy<'a>( + &'a self, + source_path: &'a PathUri, + destination_path: &'a PathUri, + options: CopyOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(UnsandboxedFileSystem::copy( + self, + source_path, + destination_path, + options, + sandbox, + )) + } +} + +impl DirectFileSystem { + async fn open_file_for_read( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + reject_sandbox_context(sandbox)?; + let path = path.to_abs_path()?; + regular_file::open(path.as_path()).await + } + + async fn canonicalize( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + reject_sandbox_context(sandbox)?; + let path = path.to_abs_path()?; + let canonicalized = + AbsolutePathBuf::from_absolute_path(tokio::fs::canonicalize(path.as_path()).await?)?; + Ok(PathUri::from_abs_path(&canonicalized)) + } + + async fn read_file( + &self, + path: &PathUri, + options: ReadFileOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult> { + reject_sandbox_context(sandbox)?; + let file = if options.follow_symlinks { + self.open_file_for_read(path, /*sandbox*/ None).await? + } else { + no_follow::open_file(path.to_abs_path()?.as_path()).await? + }; + let metadata = file.metadata().await?; + if metadata.len() > MAX_READ_FILE_BYTES { + return Err(file_too_large_error()); + } + let mut bytes = Vec::with_capacity(metadata.len() as usize); + file.take(MAX_READ_FILE_BYTES + 1) + .read_to_end(&mut bytes) + .await?; + if bytes.len() as u64 > MAX_READ_FILE_BYTES { + return Err(file_too_large_error()); + } + Ok(bytes) + } + + async fn read_file_stream( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + let file = self.open_file_for_read(path, sandbox).await?; + Ok(FileSystemReadStream::new(ReaderStream::with_capacity( + file, + FILE_READ_CHUNK_SIZE, + ))) + } + + async fn write_file( + &self, + path: &PathUri, + contents: Vec, + options: WriteFileOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + reject_sandbox_context(sandbox)?; + let path = path.to_abs_path()?; + if options.follow_symlinks { + tokio::fs::write(path.as_path(), contents).await + } else { + no_follow::write_file(path.as_path(), contents).await + } + } + + async fn create_directory( + &self, + path: &PathUri, + options: CreateDirectoryOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + reject_sandbox_context(sandbox)?; + let path = path.to_abs_path()?; + if !options.follow_symlinks { + return no_follow::create_directory(path.as_path(), options.recursive).await; + } + if options.recursive { + tokio::fs::create_dir_all(path.as_path()).await?; + } else { + tokio::fs::create_dir(path.as_path()).await?; + } + Ok(()) + } + + async fn get_metadata( + &self, + path: &PathUri, + options: GetMetadataOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + reject_sandbox_context(sandbox)?; + let path = path.to_abs_path()?; + if !options.follow_symlinks { + return no_follow::metadata(path.as_path()).await; + } + let symlink_metadata = tokio::fs::symlink_metadata(path.as_path()).await?; + let is_symlink = symlink_metadata.is_symlink(); + let metadata = if is_symlink { + tokio::fs::metadata(path.as_path()).await? + } else { + symlink_metadata + }; + Ok(file_metadata(metadata, is_symlink)) + } + + async fn read_directory( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult> { + reject_sandbox_context(sandbox)?; + let path = path.to_abs_path()?; + let mut entries = Vec::new(); + let mut read_dir = tokio::fs::read_dir(path.as_path()).await?; + while let Some(entry) = read_dir.next_entry().await? { + let Ok(mut file_type) = entry.file_type().await else { + continue; + }; + if file_type.is_symlink() { + let Ok(metadata) = tokio::fs::metadata(entry.path()).await else { + continue; + }; + file_type = metadata.file_type(); + } + entries.push(ReadDirectoryEntry { + file_name: entry.file_name().to_string_lossy().into_owned(), + is_directory: file_type.is_dir(), + is_file: file_type.is_file(), + }); + } + Ok(entries) + } + + fn sync_walk( + root: &PathUri, + options: WalkOptions, + cancelled: &CancellationToken, + ) -> io::Result { + if options.max_directories == 0 || options.max_entries == 0 { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "filesystem walk limits must be greater than zero", + )); + } + if options.max_depth > MAX_WALK_DEPTH + || options.max_directories > MAX_WALK_DIRECTORIES + || options.max_entries > MAX_WALK_ENTRIES + { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!( + "filesystem walk limits exceed maximums: depth={MAX_WALK_DEPTH}, directories={MAX_WALK_DIRECTORIES}, entries={MAX_WALK_ENTRIES}" + ), + )); + } + + check_walk_cancelled(cancelled)?; + let (root_metadata, root_is_symlink) = walk_metadata(root)?; + if !root_metadata.is_dir() || (root_is_symlink && !options.follow_directory_symlinks) { + return Ok(WalkOutcome::default()); + } + + let root_identity = if options.follow_directory_symlinks { + check_walk_cancelled(cancelled)?; + walk_canonicalize(root)? + } else { + root.clone() + }; + let mut outcome = WalkOutcome::default(); + let mut queue = VecDeque::from([(root.clone(), 0usize)]); + let mut visited_directories = HashSet::from([root_identity]); + let mut directory_count = 1usize; + let mut entry_count = 0usize; + let mut response_bytes = 0usize; + + while let Some((directory, depth)) = queue.pop_front() { + let entries = walk_read_directory(&directory, cancelled); + check_walk_cancelled(cancelled)?; + let mut entries = match entries { + Ok(entries) => entries, + Err(error) => { + if !push_walk_error( + &mut outcome, + &mut response_bytes, + directory, + error.to_string(), + ) { + return Ok(outcome); + } + continue; + } + }; + entries.sort(); + + for file_name in entries { + check_walk_cancelled(cancelled)?; + if entry_count == options.max_entries { + outcome.truncated = true; + return Ok(outcome); + } + entry_count += 1; + + let path = match directory.join(&file_name) { + Ok(path) => path, + Err(error) => { + if !push_walk_error( + &mut outcome, + &mut response_bytes, + directory.clone(), + error.to_string(), + ) { + return Ok(outcome); + } + continue; + } + }; + let (metadata, is_symlink) = match walk_metadata(&path) { + Ok(metadata) => metadata, + Err(error) => { + if !push_walk_error( + &mut outcome, + &mut response_bytes, + path, + error.to_string(), + ) { + return Ok(outcome); + } + continue; + } + }; + if is_symlink && (!options.follow_directory_symlinks || !metadata.is_dir()) { + continue; + } + + let kind = if metadata.is_dir() { + WalkEntryKind::Directory + } else if metadata.is_file() { + WalkEntryKind::File + } else { + continue; + }; + if !reserve_walk_response_bytes( + &mut outcome, + &mut response_bytes, + path.to_string().len(), + ) { + return Ok(outcome); + } + outcome.entries.push(WalkEntry { + path: path.clone(), + kind, + }); + + if kind == WalkEntryKind::Directory && depth < options.max_depth { + if options.prune_hidden_directories && file_name.starts_with('.') { + continue; + } + let directory_identity = if options.follow_directory_symlinks { + check_walk_cancelled(cancelled)?; + match walk_canonicalize(&path) { + Ok(path) => path, + Err(error) => { + if !push_walk_error( + &mut outcome, + &mut response_bytes, + path, + error.to_string(), + ) { + return Ok(outcome); + } + continue; + } + } + } else { + path.clone() + }; + if !visited_directories.insert(directory_identity) { + continue; + } + if directory_count == options.max_directories { + outcome.truncated = true; + } else { + directory_count += 1; + queue.push_back((path, depth + 1)); + } + } + } + } + + Ok(outcome) + } + + async fn remove( + &self, + path: &PathUri, + options: RemoveOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + reject_sandbox_context(sandbox)?; + let path = path.to_abs_path()?; + if !options.follow_symlinks { + return no_follow::remove(path.as_path(), options.recursive, options.force).await; + } + match tokio::fs::symlink_metadata(path.as_path()).await { + Ok(metadata) => { + let file_type = metadata.file_type(); + if file_type.is_dir() { + if options.recursive { + tokio::fs::remove_dir_all(path.as_path()).await?; + } else { + tokio::fs::remove_dir(path.as_path()).await?; + } + } else { + tokio::fs::remove_file(path.as_path()).await?; + } + Ok(()) + } + Err(err) if err.kind() == io::ErrorKind::NotFound && options.force => Ok(()), + Err(err) => Err(err), + } + } + + async fn copy( + &self, + source_path: &PathUri, + destination_path: &PathUri, + options: CopyOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + reject_sandbox_context(sandbox)?; + let source_path = source_path.to_abs_path()?.into_path_buf(); + let destination_path = destination_path.to_abs_path()?.into_path_buf(); + tokio::task::spawn_blocking(move || -> FileSystemResult<()> { + let metadata = std::fs::symlink_metadata(source_path.as_path())?; + let file_type = metadata.file_type(); + + if file_type.is_dir() { + if !options.recursive { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "fs/copy requires recursive: true when sourcePath is a directory", + )); + } + if destination_is_same_or_descendant_of_source( + source_path.as_path(), + destination_path.as_path(), + )? { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "fs/copy cannot copy a directory to itself or one of its descendants", + )); + } + copy_dir_recursive(source_path.as_path(), destination_path.as_path())?; + return Ok(()); + } + + if file_type.is_symlink() { + copy_symlink(source_path.as_path(), destination_path.as_path())?; + return Ok(()); + } + + if file_type.is_file() { + std::fs::copy(source_path.as_path(), destination_path.as_path())?; + return Ok(()); + } + + Err(io::Error::new( + io::ErrorKind::InvalidInput, + "fs/copy only supports regular files, directories, and symlinks", + )) + }) + .await + .map_err(|err| io::Error::other(format!("filesystem task failed: {err}")))? + } +} + +impl ExecutorFileSystem for DirectFileSystem { + fn canonicalize<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, PathUri> { + Box::pin(DirectFileSystem::canonicalize(self, path, sandbox)) + } + + fn read_file<'a>( + &'a self, + path: &'a PathUri, + options: ReadFileOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, Vec> { + Box::pin(DirectFileSystem::read_file(self, path, options, sandbox)) + } + + fn read_file_stream<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, FileSystemReadStream> { + Box::pin(DirectFileSystem::read_file_stream(self, path, sandbox)) + } + + fn write_file<'a>( + &'a self, + path: &'a PathUri, + contents: Vec, + options: WriteFileOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(DirectFileSystem::write_file( + self, path, contents, options, sandbox, + )) + } + + fn create_directory<'a>( + &'a self, + path: &'a PathUri, + options: CreateDirectoryOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(DirectFileSystem::create_directory( + self, path, options, sandbox, + )) + } + + fn get_metadata<'a>( + &'a self, + path: &'a PathUri, + options: GetMetadataOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, FileMetadata> { + Box::pin(DirectFileSystem::get_metadata(self, path, options, sandbox)) + } + + fn read_directory<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, Vec> { + Box::pin(DirectFileSystem::read_directory(self, path, sandbox)) + } + + fn walk<'a>( + &'a self, + path: &'a PathUri, + options: WalkOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, WalkOutcome> { + Box::pin(async move { + reject_sandbox_context(sandbox)?; + let path = path.clone(); + let cancelled = CancellationToken::new(); + let _cancel_on_drop = cancelled.clone().drop_guard(); + tokio::task::spawn_blocking(move || Self::sync_walk(&path, options, &cancelled)) + .await + .map_err(|err| io::Error::other(format!("filesystem task failed: {err}")))? + }) + } + + fn remove<'a>( + &'a self, + path: &'a PathUri, + options: RemoveOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(DirectFileSystem::remove(self, path, options, sandbox)) + } + + fn copy<'a>( + &'a self, + source_path: &'a PathUri, + destination_path: &'a PathUri, + options: CopyOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(DirectFileSystem::copy( + self, + source_path, + destination_path, + options, + sandbox, + )) + } +} + +fn check_walk_cancelled(cancelled: &CancellationToken) -> io::Result<()> { + if cancelled.is_cancelled() { + return Err(io::Error::new( + io::ErrorKind::Interrupted, + "filesystem walk cancelled", + )); + } + Ok(()) +} + +fn walk_metadata(path: &PathUri) -> io::Result<(std::fs::Metadata, bool)> { + let path = path.to_abs_path()?; + let metadata = std::fs::symlink_metadata(path.as_path())?; + let is_symlink = metadata.is_symlink(); + let metadata = if is_symlink { + std::fs::metadata(path.as_path())? + } else { + metadata + }; + Ok((metadata, is_symlink)) +} + +fn walk_canonicalize(path: &PathUri) -> io::Result { + let path = path.to_abs_path()?; + let canonicalized = + AbsolutePathBuf::from_absolute_path(std::fs::canonicalize(path.as_path())?)?; + Ok(PathUri::from_abs_path(&canonicalized)) +} + +fn walk_read_directory(path: &PathUri, cancelled: &CancellationToken) -> io::Result> { + check_walk_cancelled(cancelled)?; + let path = path.to_abs_path()?; + let mut entries = Vec::new(); + for entry in std::fs::read_dir(path.as_path())? { + check_walk_cancelled(cancelled)?; + let entry = entry?; + let Ok(file_type) = entry.file_type() else { + continue; + }; + // Match DirectFileSystem::read_directory: omit broken or inaccessible links. + if file_type.is_symlink() { + check_walk_cancelled(cancelled)?; + if std::fs::metadata(entry.path()).is_err() { + continue; + } + } + entries.push(entry.file_name().to_string_lossy().into_owned()); + } + Ok(entries) +} + +fn push_walk_error( + outcome: &mut WalkOutcome, + response_bytes: &mut usize, + path: PathUri, + message: String, +) -> bool { + let item_bytes = path.to_string().len().saturating_add(message.len()); + if !reserve_walk_response_bytes(outcome, response_bytes, item_bytes) { + return false; + } + outcome.errors.push(WalkError { path, message }); + true +} + +fn reserve_walk_response_bytes( + outcome: &mut WalkOutcome, + response_bytes: &mut usize, + content_bytes: usize, +) -> bool { + let item_bytes = content_bytes.saturating_add(WALK_RESPONSE_ITEM_OVERHEAD_BYTES); + let Some(total_bytes) = response_bytes.checked_add(item_bytes) else { + outcome.truncated = true; + return false; + }; + if total_bytes > MAX_WALK_RESPONSE_BYTES { + outcome.truncated = true; + return false; + } + *response_bytes = total_bytes; + true +} + +fn reject_sandbox_context(sandbox: Option<&FileSystemSandboxContext>) -> io::Result<()> { + if sandbox.is_some() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "direct filesystem operations do not accept sandbox context", + )); + } + Ok(()) +} + +fn file_metadata(metadata: std::fs::Metadata, is_symlink: bool) -> FileMetadata { + FileMetadata { + is_directory: metadata.is_dir(), + is_file: metadata.is_file(), + is_symlink, + size: metadata.len(), + created_at_ms: metadata.created().ok().map_or(0, system_time_to_unix_ms), + modified_at_ms: metadata.modified().ok().map_or(0, system_time_to_unix_ms), + } +} + +fn reject_platform_sandbox_context(sandbox: Option<&FileSystemSandboxContext>) -> io::Result<()> { + if sandbox.is_some_and(FileSystemSandboxContext::should_run_in_sandbox) { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "sandboxed filesystem operations require configured runtime paths", + )); + } + Ok(()) +} + +fn copy_dir_recursive(source: &Path, target: &Path) -> io::Result<()> { + std::fs::create_dir_all(target)?; + for entry in std::fs::read_dir(source)? { + let entry = entry?; + let source_path = entry.path(); + let target_path = target.join(entry.file_name()); + let file_type = entry.file_type()?; + + if file_type.is_dir() { + copy_dir_recursive(&source_path, &target_path)?; + } else if file_type.is_file() { + std::fs::copy(&source_path, &target_path)?; + } else if file_type.is_symlink() { + copy_symlink(&source_path, &target_path)?; + } + } + Ok(()) +} + +fn destination_is_same_or_descendant_of_source( + source: &Path, + destination: &Path, +) -> io::Result { + let source = std::fs::canonicalize(source)?; + let destination = resolve_existing_path(destination)?; + Ok(destination.starts_with(&source)) +} + +pub(crate) fn resolve_existing_path(path: &Path) -> io::Result { + let mut unresolved_suffix = Vec::new(); + let mut existing_path = path; + while !existing_path.exists() { + let Some(file_name) = existing_path.file_name() else { + break; + }; + unresolved_suffix.push(file_name.to_os_string()); + let Some(parent) = existing_path.parent() else { + break; + }; + existing_path = parent; + } + + let mut resolved = std::fs::canonicalize(existing_path)?; + for file_name in unresolved_suffix.iter().rev() { + resolved.push(file_name); + } + Ok(resolved) +} + +pub(crate) fn current_sandbox_cwd() -> io::Result { + let cwd = std::env::current_dir() + .map_err(|err| io::Error::other(format!("failed to read current dir: {err}")))?; + resolve_existing_path(cwd.as_path()) +} + +fn copy_symlink(source: &Path, target: &Path) -> io::Result<()> { + let link_target = std::fs::read_link(source)?; + #[cfg(unix)] + { + std::os::unix::fs::symlink(&link_target, target) + } + #[cfg(windows)] + { + if symlink_points_to_directory(source)? { + std::os::windows::fs::symlink_dir(&link_target, target) + } else { + std::os::windows::fs::symlink_file(&link_target, target) + } + } + #[cfg(not(any(unix, windows)))] + { + let _ = link_target; + let _ = target; + Err(io::Error::new( + io::ErrorKind::Unsupported, + "copying symlinks is unsupported on this platform", + )) + } +} + +#[cfg(windows)] +fn symlink_points_to_directory(source: &Path) -> io::Result { + use std::os::windows::fs::FileTypeExt; + + Ok(std::fs::symlink_metadata(source)? + .file_type() + .is_symlink_dir()) +} + +fn system_time_to_unix_ms(time: SystemTime) -> i64 { + time.duration_since(UNIX_EPOCH) + .ok() + .and_then(|duration| i64::try_from(duration.as_millis()).ok()) + .unwrap_or(0) +} + +#[cfg(all(test, any(unix, windows)))] +#[path = "local_file_system_path_uri_tests.rs"] +mod path_uri_tests; + +#[cfg(all(test, unix))] +mod tests { + use super::*; + use pretty_assertions::assert_eq; + use std::os::unix::fs::symlink; + + #[test] + fn resolve_existing_path_handles_symlink_parent_dotdot_escape() -> io::Result<()> { + let temp_dir = tempfile::TempDir::new()?; + let allowed_dir = temp_dir.path().join("allowed"); + let outside_dir = temp_dir.path().join("outside"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + symlink(&outside_dir, allowed_dir.join("link"))?; + + let resolved = resolve_existing_path( + allowed_dir + .join("link") + .join("..") + .join("secret.txt") + .as_path(), + )?; + + assert_eq!( + resolved, + resolve_existing_path(temp_dir.path())?.join("secret.txt") + ); + Ok(()) + } +} + +#[cfg(all(test, windows))] +mod tests { + use super::*; + use pretty_assertions::assert_eq; + + #[test] + fn symlink_points_to_directory_handles_dangling_directory_symlinks() -> io::Result<()> { + use std::os::windows::fs::symlink_dir; + + let temp_dir = tempfile::TempDir::new()?; + let source_dir = temp_dir.path().join("source"); + let link_path = temp_dir.path().join("source-link"); + std::fs::create_dir(&source_dir)?; + + if symlink_dir(&source_dir, &link_path).is_err() { + return Ok(()); + } + + std::fs::remove_dir(&source_dir)?; + + assert_eq!(symlink_points_to_directory(&link_path)?, true); + Ok(()) + } +} + +#[cfg(test)] +mod walk_tests { + use super::*; + use codex_protocol::models::PermissionProfile; + use codex_protocol::permissions::FileSystemSandboxPolicy; + use codex_protocol::permissions::NetworkSandboxPolicy; + use pretty_assertions::assert_eq; + + #[tokio::test] + async fn sync_walk_rejects_sandbox_context() -> io::Result<()> { + let temp = tempfile::tempdir()?; + let root = PathUri::from_host_native_path(temp.path())?; + let sandbox = FileSystemSandboxContext::from_permission_profile( + PermissionProfile::from_runtime_permissions( + &FileSystemSandboxPolicy::restricted(Vec::new()), + NetworkSandboxPolicy::Restricted, + ), + ); + let options = WalkOptions { + max_depth: 1, + max_directories: 1, + max_entries: 1, + follow_directory_symlinks: false, + prune_hidden_directories: false, + }; + let direct_error = DirectFileSystem + .walk(&root, options, Some(&sandbox)) + .await + .expect_err("direct walk must reject sandbox contexts"); + let wrapper_error = UnsandboxedFileSystem::default() + .walk(&root, options, Some(&sandbox)) + .await + .expect_err("unsandboxed walk must reject restricted contexts"); + assert_eq!(direct_error.kind(), io::ErrorKind::InvalidInput); + assert_eq!(wrapper_error.kind(), io::ErrorKind::InvalidInput); + Ok(()) + } + + #[test] + fn sync_walk_cancellation_stops_before_io() -> io::Result<()> { + let temp = tempfile::tempdir()?; + let missing = PathUri::from_host_native_path(temp.path().join("missing"))?; + let options = WalkOptions { + max_depth: 1, + max_directories: 1, + max_entries: 1, + follow_directory_symlinks: true, + prune_hidden_directories: false, + }; + let cancelled = CancellationToken::new(); + let cancel_on_drop = cancelled.clone().drop_guard(); + drop(cancel_on_drop); + + for result in [ + DirectFileSystem::sync_walk(&missing, options, &cancelled).map(|_| ()), + walk_read_directory(&missing, &cancelled).map(|_| ()), + ] { + assert_eq!( + result + .expect_err("cancelled walks must stop before I/O") + .kind(), + io::ErrorKind::Interrupted, + ); + } + Ok(()) + } + + #[test] + fn sync_walk_response_budget_counts_entries_and_errors() -> io::Result<()> { + let temp = tempfile::tempdir()?; + let root = PathUri::from_host_native_path(temp.path())?; + let mut outcome = WalkOutcome::default(); + let mut response_bytes = + MAX_WALK_RESPONSE_BYTES - WALK_RESPONSE_ITEM_OVERHEAD_BYTES - root.to_string().len(); + assert!(push_walk_error( + &mut outcome, + &mut response_bytes, + root.clone(), + String::new() + )); + assert!(!reserve_walk_response_bytes( + &mut outcome, + &mut response_bytes, + /*content_bytes*/ 0 + )); + assert_eq!( + outcome, + WalkOutcome { + entries: Vec::new(), + errors: vec![WalkError { + path: root, + message: String::new() + }], + truncated: true, + }, + ); + Ok(()) + } +} diff --git a/codex-rs/exec-server/src/local_file_system_path_uri_tests.rs b/codex-rs/exec-server/src/local_file_system_path_uri_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..78a48c5302e65b5598ba5f9a27953367028cfff1 --- /dev/null +++ b/codex-rs/exec-server/src/local_file_system_path_uri_tests.rs @@ -0,0 +1,27 @@ +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +use tokio::io; + +use super::*; + +#[tokio::test] +async fn direct_file_system_rejects_non_native_uri_as_invalid_input() { + let error = DirectFileSystem + .read_file(&non_native_uri(), Default::default(), /*sandbox*/ None) + .await + .expect_err("non-native URI should be rejected"); + + assert_eq!(error.kind(), io::ErrorKind::InvalidInput); +} + +fn non_native_uri() -> PathUri { + #[cfg(unix)] + let uri = "file://server/share/file.txt"; + #[cfg(windows)] + let uri = "file:///usr/local/file.txt"; + + match PathUri::parse(uri) { + Ok(uri) => uri, + Err(err) => panic!("valid non-native URI should parse: {err}"), + } +} diff --git a/codex-rs/exec-server/src/local_process.rs b/codex-rs/exec-server/src/local_process.rs new file mode 100644 index 0000000000000000000000000000000000000000..be4606883c7bb75a955499842335f778577eaab2 --- /dev/null +++ b/codex-rs/exec-server/src/local_process.rs @@ -0,0 +1,2096 @@ +use std::collections::HashMap; +use std::collections::HashSet; +use std::collections::VecDeque; +use std::collections::hash_map::Entry; +use std::sync::Arc; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use crate::process_telemetry::ProcessTelemetry; +use crate::process_telemetry::ProcessTelemetryEvent; +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_network_proxy::NetworkPolicyAuditEvent; +use codex_network_proxy::NetworkPolicyAuditObserver; +use codex_network_proxy::NetworkProtocol; +use codex_network_proxy::NetworkProxyHandle; +use codex_protocol::config_types::EnvironmentVariablePattern; +use codex_protocol::config_types::ShellEnvironmentPolicy; +use codex_protocol::exec_output::ExecToolCallOutput; +use codex_protocol::exec_output::StreamOutput; +use codex_protocol::shell_environment; +use codex_sandboxing::SandboxType; +use codex_sandboxing::is_likely_sandbox_denied; +use codex_utils_pty::ExecCommandSession; +use codex_utils_pty::ProcessSignal as PtyProcessSignal; +use opentelemetry::trace::SpanContext; +use opentelemetry::trace::TraceContextExt; +use tokio::sync::Mutex; +use tokio::sync::Notify; +use tokio::sync::mpsc; +use tokio::sync::watch; +use tokio_util::sync::CancellationToken; +use tracing::Instrument; +use tracing::instrument::WithSubscriber; + +use crate::ExecBackend; +use crate::ExecBackendFuture; +use crate::ExecProcess; +use crate::ExecProcessEvent; +use crate::ExecProcessEventReceiver; +use crate::ExecProcessFuture; +use crate::ExecServerError; +use crate::ExecServerRuntimePaths; +use crate::ProcessId; +use crate::StartedExecProcess; +use crate::network_policy_decisions::network_policy_decider; +use crate::process::ExecProcessEventLog; +use crate::process::sandbox_type_from_protocol; +use crate::process_sandbox::prepare_exec_request_with_telemetry; +use crate::protocol::EXEC_CLOSED_METHOD; +use crate::protocol::ExecClosedNotification; +use crate::protocol::ExecEnvPolicy; +use crate::protocol::ExecExitedNotification; +use crate::protocol::ExecOutputDeltaNotification; +use crate::protocol::ExecOutputStream; +use crate::protocol::ExecParams; +use crate::protocol::ExecResponse; +use crate::protocol::ExecServerNetworkProtocol; +use crate::protocol::MAX_NETWORK_POLICY_PROCESS_ID_BYTES; +use crate::protocol::NETWORK_POLICY_DECISION_METHOD; +use crate::protocol::NetworkPolicyDecisionNotification; +use crate::protocol::ProcessOutputChunk; +use crate::protocol::ProcessSandboxType; +use crate::protocol::ProcessSignal; +use crate::protocol::ReadParams; +use crate::protocol::ReadResponse; +use crate::protocol::SignalParams; +use crate::protocol::SignalResponse; +use crate::protocol::TerminateParams; +use crate::protocol::TerminateResponse; +use crate::protocol::WriteParams; +use crate::protocol::WriteResponse; +use crate::protocol::WriteStatus; +use crate::rpc::RpcNotificationSender; +use crate::rpc::RpcServerOutboundMessage; +use crate::rpc::internal_error; +use crate::rpc::invalid_params; +use crate::rpc::invalid_request; +use crate::rpc_server_requests::RpcServerRequestSender; +#[cfg(unix)] +use crate::shell_snapshot::CapturePurpose; +use crate::telemetry::ExecServerTelemetry; +use crate::telemetry::ProcessMetricGuard; + +const RETAINED_OUTPUT_BYTES_PER_PROCESS: usize = 1024 * 1024; +// Each process/read chunk needs four JSON values. Keep retained replay below the +// shared 256K-value JSON-RPC decoder budget even when output arrives in tiny chunks. +const RETAINED_OUTPUT_CHUNKS_PER_PROCESS: usize = 50_000; +const NOTIFICATION_CHANNEL_CAPACITY: usize = 256; +const PROCESS_EVENT_CHANNEL_CAPACITY: usize = 256; +const RETAINED_STDIN_WRITE_IDS_PER_PROCESS: usize = 4096; +static NEXT_LOCAL_STDIN_WRITE_ID: AtomicU64 = AtomicU64::new(1); +#[cfg(test)] +const EXITED_PROCESS_RETENTION: Duration = Duration::from_millis(25); +#[cfg(not(test))] +const EXITED_PROCESS_RETENTION: Duration = Duration::from_secs(30); + +#[derive(Clone)] +struct RetainedOutputChunk { + seq: u64, + stream: ExecOutputStream, + chunk: Vec, +} + +struct RunningProcess { + session: ExecCommandSession, + tty: bool, + pipe_stdin: bool, + accepted_stdin_write_ids: Arc>, + output: VecDeque, + retained_bytes: usize, + next_seq: u64, + exit_code: Option, + wake_tx: watch::Sender, + events: ExecProcessEventLog, + output_notify: Arc, + open_streams: usize, + closed: bool, + metrics: Option, + termination_requested: bool, + sandbox: SandboxType, + sandbox_denied: bool, + network_proxy_handle: Option, + network_policy_shutdown: Option, +} + +/// Bounded cache of stdin write ids that have already been accepted for one process. +/// +/// A remote client can retry `process/write` after reconnecting. Remembering accepted +/// ids lets the server acknowledge the retried request without writing the same bytes +/// to child stdin twice. +#[derive(Default)] +struct AcceptedStdinWriteIds { + ids: HashSet, + order: VecDeque, +} + +impl AcceptedStdinWriteIds { + fn contains(&self, write_id: &str) -> bool { + self.ids.contains(write_id) + } + + fn remember(&mut self, write_id: String) { + if !self.ids.insert(write_id.clone()) { + return; + } + + self.order.push_back(write_id); + while self.order.len() > RETAINED_STDIN_WRITE_IDS_PER_PROCESS { + let Some(evicted) = self.order.pop_front() else { + break; + }; + self.ids.remove(&evicted); + } + } +} + +struct ProcessStart; + +enum ProcessEntry { + Starting(Arc), + Running(Box), +} + +struct Inner { + notifications: std::sync::RwLock>, + requests: Arc>>, + processes: Mutex>, + #[cfg(unix)] + shell_snapshots: crate::shell_snapshot::ShellSnapshotCache, + telemetry: ExecServerTelemetry, +} + +#[derive(Clone)] +pub(crate) struct LocalProcess { + inner: Arc, + runtime_paths: Option, +} + +struct LocalExecProcess { + process_id: ProcessId, + backend: LocalProcess, + wake_tx: watch::Sender, + events: ExecProcessEventLog, +} + +impl Default for LocalProcess { + fn default() -> Self { + Self::with_discarded_notifications(/*runtime_paths*/ None) + } +} + +impl LocalProcess { + pub(crate) fn with_local_runtime_paths(runtime_paths: ExecServerRuntimePaths) -> Self { + Self::with_discarded_notifications(Some(runtime_paths)) + } + + fn with_discarded_notifications(runtime_paths: Option) -> Self { + let (outgoing_tx, mut outgoing_rx) = + mpsc::channel::(NOTIFICATION_CHANNEL_CAPACITY); + tokio::spawn(async move { while outgoing_rx.recv().await.is_some() {} }); + Self::with_runtime_paths( + RpcNotificationSender::new(outgoing_tx), + ExecServerTelemetry::default(), + runtime_paths, + ) + } + + pub(crate) fn new( + notifications: RpcNotificationSender, + telemetry: ExecServerTelemetry, + runtime_paths: ExecServerRuntimePaths, + ) -> Self { + Self::with_runtime_paths(notifications, telemetry, Some(runtime_paths)) + } + + fn with_runtime_paths( + notifications: RpcNotificationSender, + telemetry: ExecServerTelemetry, + runtime_paths: Option, + ) -> Self { + let requests = notifications.request_sender(); + Self { + inner: Arc::new(Inner { + notifications: std::sync::RwLock::new(Some(notifications)), + requests: Arc::new(std::sync::RwLock::new(Some(requests))), + processes: Mutex::new(HashMap::new()), + #[cfg(unix)] + shell_snapshots: crate::shell_snapshot::ShellSnapshotCache::default(), + telemetry, + }), + runtime_paths, + } + } + + pub(crate) async fn shutdown(&self) { + let remaining = { + let mut processes = self.inner.processes.lock().await; + processes + .drain() + .filter_map(|(_, process)| match process { + ProcessEntry::Starting(_) => None, + ProcessEntry::Running(process) => Some(process), + }) + .collect::>() + }; + for mut process in remaining { + if let Some(network_policy_shutdown) = process.network_policy_shutdown.take() { + network_policy_shutdown.cancel(); + } + if let Some(metrics) = process.metrics.take() { + metrics.finish("terminated"); + } + process.session.terminate(); + } + } + + pub(crate) fn set_notification_sender(&self, notifications: Option) { + let requests = notifications + .as_ref() + .map(RpcNotificationSender::request_sender); + let mut notification_sender = self + .inner + .notifications + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + *notification_sender = notifications; + let previous_requests = std::mem::replace( + &mut *self + .inner + .requests + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner), + requests, + ); + if let Some(previous_requests) = previous_requests { + previous_requests.close(); + } + } + + async fn start_process( + &self, + params: ExecParams, + mut telemetry: ProcessTelemetry, + ) -> Result<(ExecResponse, watch::Sender, ExecProcessEventLog), JSONRPCErrorError> { + telemetry.launch_context = telemetry.launch_context.filter(SpanContext::is_valid); + let metadata = params.metadata.as_ref(); + telemetry.thread_id = metadata + .and_then(|metadata| metadata.thread_id.as_ref()) + .map(ToString::to_string); + // Correlation is controller-supplied, not authorization or arbitrary diagnostic text. + telemetry.tool_call_id = metadata + .and_then(|metadata| metadata.tool_call_id.as_ref()) + .filter(|id| { + !id.is_empty() + && id.len() <= 256 + && id + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || b"_-.:".contains(&byte)) + }) + .cloned(); + let process_id = params.process_id.clone(); + let policy_decision_timeout_ms = params + .network_proxy + .as_ref() + .and_then(|launch| launch.policy_decision_timeout_ms); + if policy_decision_timeout_ms == Some(0) { + return Err(invalid_params( + "network policy decision callback timeout must be nonzero".to_string(), + )); + } + if policy_decision_timeout_ms.is_some() + && (process_id.is_empty() || process_id.len() > MAX_NETWORK_POLICY_PROCESS_ID_BYTES) + { + return Err(invalid_params(format!( + "callback-enabled process ID must be non-empty and at most {MAX_NETWORK_POLICY_PROCESS_ID_BYTES} bytes" + ))); + } + let policy_decision_timeout = policy_decision_timeout_ms.map(Duration::from_millis); + let network_policy_shutdown = policy_decision_timeout.map(|_| CancellationToken::new()); + let network_policy_decider = network_policy_shutdown + .as_ref() + .zip(policy_decision_timeout) + .map(|(process_shutdown, controller_timeout)| { + network_policy_decider( + process_id.clone(), + Arc::clone(&self.inner.requests), + controller_timeout, + process_shutdown.clone(), + ) + }); + let network_policy_audit_observer = params.network_proxy.as_ref().map(|_| { + let process_id = process_id.clone(); + let inner = Arc::downgrade(&self.inner); + Arc::new(move |event: NetworkPolicyAuditEvent| { + let Some(inner) = inner.upgrade() else { + return; + }; + let Some(notifications) = notification_sender(&inner) else { + return; + }; + let notification = NetworkPolicyDecisionNotification { + process_id: process_id.clone(), + timestamp: event.timestamp, + scope: event.scope, + decision: event.decision, + source: event.source, + reason: event.reason, + protocol: match event.protocol { + NetworkProtocol::Http => ExecServerNetworkProtocol::Http, + NetworkProtocol::HttpsConnect => ExecServerNetworkProtocol::HttpsConnect, + NetworkProtocol::Socks5Tcp => ExecServerNetworkProtocol::Socks5Tcp, + NetworkProtocol::Socks5Udp => ExecServerNetworkProtocol::Socks5Udp, + }, + host: event.host, + port: event.port, + method: event.method, + client: event.client, + policy_override: event.policy_override, + }; + let _ = notifications.try_notify(NETWORK_POLICY_DECISION_METHOD, ¬ification); + }) as NetworkPolicyAuditObserver + }); + #[cfg(not(unix))] + if params.shell_snapshot.is_some() { + return Err(invalid_params( + "shell snapshots are unsupported on this platform".to_string(), + )); + } + let prepared = prepare_exec_request_with_telemetry( + ¶ms, + child_env(¶ms), + self.runtime_paths.as_ref(), + network_policy_decider, + network_policy_audit_observer, + &telemetry, + ) + .await?; + #[cfg(unix)] + let mut prepared = prepared; + #[cfg(unix)] + self.inner + .shell_snapshots + .prepare( + ¶ms, + &mut prepared, + &self.inner.telemetry, + CapturePurpose::Execution, + ) + .await?; + if prepared.command.is_empty() { + return Err(invalid_params("argv must not be empty".to_string())); + } + let sandbox_type = match prepared.sandbox { + SandboxType::None => Some(ProcessSandboxType::None), + SandboxType::MacosSeatbelt => Some(ProcessSandboxType::MacosSeatbelt), + SandboxType::LinuxSeccomp => Some(ProcessSandboxType::LinuxSeccomp), + SandboxType::WindowsRestrictedToken => Some(ProcessSandboxType::WindowsRestrictedToken), + SandboxType::WindowsMxc => Some(ProcessSandboxType::WindowsMxc), + }; + + let start = Arc::new(ProcessStart); + { + let mut process_map = self.inner.processes.lock().await; + if process_map.contains_key(&process_id) { + return Err(invalid_request(format!( + "process {process_id} already exists" + ))); + } + process_map.insert( + process_id.clone(), + ProcessEntry::Starting(Arc::clone(&start)), + ); + } + + let spawned_result = codex_sandboxing::spawn_process(codex_sandboxing::SpawnRequest { + command: &prepared.command, + cwd: prepared.cwd.as_path(), + env: &prepared.env, + arg0: &prepared.arg0, + sandbox: prepared.sandbox, + windows_sandbox: prepared.windows_sandbox_spawn_request(), + tty: params.tty, + stdin_open: params.tty || params.pipe_stdin, + inherited_fds: &[], + }) + .await; + let spawned = match spawned_result { + Ok(spawned) => spawned, + Err(err) => { + telemetry.log(ProcessTelemetryEvent::SpawnFailed, prepared.sandbox); + let mut process_map = self.inner.processes.lock().await; + if matches!( + process_map.get(&process_id), + Some(ProcessEntry::Starting(current)) if Arc::ptr_eq(current, &start) + ) { + process_map.remove(&process_id); + } + return Err(internal_error(err.to_string())); + } + }; + let metrics = self.inner.telemetry.process_started(&process_id); + + let output_notify = Arc::new(Notify::new()); + let (wake_tx, _wake_rx) = watch::channel(0); + let events = ExecProcessEventLog::new( + PROCESS_EVENT_CHANNEL_CAPACITY, + RETAINED_OUTPUT_BYTES_PER_PROCESS, + ); + { + let mut process_map = self.inner.processes.lock().await; + if !matches!( + process_map.get(&process_id), + Some(ProcessEntry::Starting(current)) if Arc::ptr_eq(current, &start) + ) { + drop(process_map); + spawned.session.terminate(); + metrics.finish("terminated"); + return Err(invalid_request(format!( + "process {process_id} start was cancelled" + ))); + } + process_map.insert( + process_id.clone(), + ProcessEntry::Running(Box::new(RunningProcess { + session: spawned.session, + tty: params.tty, + pipe_stdin: params.pipe_stdin, + accepted_stdin_write_ids: Arc::new( + Mutex::new(AcceptedStdinWriteIds::default()), + ), + output: VecDeque::new(), + retained_bytes: 0, + next_seq: 1, + exit_code: None, + wake_tx: wake_tx.clone(), + events: events.clone(), + output_notify: Arc::clone(&output_notify), + open_streams: 2, + closed: false, + metrics: Some(metrics), + termination_requested: false, + sandbox: prepared.sandbox, + sandbox_denied: false, + network_proxy_handle: prepared.network_proxy_handle, + network_policy_shutdown, + })), + ); + } + telemetry.log(ProcessTelemetryEvent::Start, prepared.sandbox); + tokio::spawn(stream_output( + process_id.clone(), + if params.tty { + ExecOutputStream::Pty + } else { + ExecOutputStream::Stdout + }, + spawned.stdout_rx, + Arc::clone(&self.inner), + Arc::clone(&output_notify), + )); + tokio::spawn(stream_output( + process_id.clone(), + if params.tty { + ExecOutputStream::Pty + } else { + ExecOutputStream::Stderr + }, + spawned.stderr_rx, + Arc::clone(&self.inner), + Arc::clone(&output_notify), + )); + // Keep the subscriber, but let the request span close independently of process completion. + tokio::spawn( + watch_exit( + process_id.clone(), + spawned.exit_rx, + Arc::clone(&self.inner), + output_notify, + telemetry, + ) + .with_current_subscriber(), + ); + + Ok(( + ExecResponse { + process_id, + sandbox_type, + }, + wake_tx, + events, + )) + } + + pub(crate) async fn exec( + &self, + params: ExecParams, + telemetry: ProcessTelemetry, + ) -> Result { + self.start_process(params, telemetry) + .await + .map(|(response, _, _)| response) + } + + pub(crate) async fn exec_read( + &self, + params: ReadParams, + ) -> Result { + let after_seq = params.after_seq.unwrap_or(0); + let max_bytes = params.max_bytes.unwrap_or(usize::MAX); + let wait = Duration::from_millis(params.wait_ms.unwrap_or(0)); + let deadline = tokio::time::Instant::now() + wait; + + loop { + let (response, output_notify) = { + let process_map = self.inner.processes.lock().await; + let process = process_map.get(¶ms.process_id).ok_or_else(|| { + invalid_request(format!("unknown process id {}", params.process_id)) + })?; + let ProcessEntry::Running(process) = process else { + return Err(invalid_request(format!( + "process id {} is starting", + params.process_id + ))); + }; + + let mut chunks = Vec::new(); + let mut total_bytes = 0; + let mut next_seq = process.next_seq; + for retained in process.output.iter().filter(|chunk| chunk.seq > after_seq) { + let chunk_len = retained.chunk.len(); + if !chunks.is_empty() && total_bytes + chunk_len > max_bytes { + break; + } + total_bytes += chunk_len; + chunks.push(ProcessOutputChunk { + seq: retained.seq, + stream: retained.stream, + chunk: retained.chunk.clone().into(), + }); + next_seq = retained.seq + 1; + if total_bytes >= max_bytes { + break; + } + } + if params.max_bytes.is_none() { + next_seq = process.next_seq; + } + ( + ReadResponse { + chunks, + next_seq, + exited: process.exit_code.is_some(), + exit_code: process.exit_code, + closed: process.closed, + failure: None, + sandbox_denied: process.sandbox_denied, + }, + Arc::clone(&process.output_notify), + ) + }; + + let has_new_terminal_event = + response.exited && after_seq < response.next_seq.saturating_sub(1); + if !response.chunks.is_empty() + || response.closed + || has_new_terminal_event + || tokio::time::Instant::now() >= deadline + { + let _total_bytes: usize = response + .chunks + .iter() + .map(|chunk| chunk.chunk.0.len()) + .sum(); + return Ok(response); + } + + let remaining = deadline.saturating_duration_since(tokio::time::Instant::now()); + if remaining.is_zero() { + return Ok(response); + } + let _ = tokio::time::timeout(remaining, output_notify.notified()).await; + } + } + + pub(crate) async fn exec_write( + &self, + params: WriteParams, + ) -> Result { + let _input_bytes = params.chunk.0.len(); + if params.write_id.is_empty() { + return Err(invalid_params("writeId must not be empty".to_string())); + } + + let (writer_tx, accepted_stdin_write_ids) = { + let process_map = self.inner.processes.lock().await; + let Some(process) = process_map.get(¶ms.process_id) else { + return Ok(WriteResponse { + status: WriteStatus::UnknownProcess, + }); + }; + let ProcessEntry::Running(process) = process else { + return Ok(WriteResponse { + status: WriteStatus::Starting, + }); + }; + if !process.tty && !process.pipe_stdin { + return Ok(WriteResponse { + status: WriteStatus::StdinClosed, + }); + } + ( + process.session.writer_sender(), + Arc::clone(&process.accepted_stdin_write_ids), + ) + }; + + if accepted_stdin_write_ids + .lock() + .await + .contains(¶ms.write_id) + { + return Ok(WriteResponse { + status: WriteStatus::Accepted, + }); + } + + let permit = writer_tx + .reserve() + .await + .map_err(|_| internal_error("failed to write to process stdin".to_string()))?; + let mut accepted_stdin_write_ids = accepted_stdin_write_ids.lock().await; + if accepted_stdin_write_ids.contains(¶ms.write_id) { + return Ok(WriteResponse { + status: WriteStatus::Accepted, + }); + } + + // After this synchronous send, record the write id before any further await. + // Otherwise a cancelled RPC handler could retry and write the same bytes again. + permit.send(params.chunk.into_inner()); + accepted_stdin_write_ids.remember(params.write_id); + + Ok(WriteResponse { + status: WriteStatus::Accepted, + }) + } + + pub(crate) async fn signal_process( + &self, + params: SignalParams, + ) -> Result { + { + let process_map = self.inner.processes.lock().await; + match process_map.get(¶ms.process_id) { + Some(ProcessEntry::Running(process)) => { + if process.exit_code.is_some() { + return Ok(SignalResponse {}); + } + process + .session + .signal(pty_process_signal(params.signal)) + .map_err(|err| internal_error(format!("failed to signal process: {err}")))? + } + Some(ProcessEntry::Starting(_)) | None => {} + } + } + + Ok(SignalResponse {}) + } + + pub(crate) async fn terminate_process( + &self, + params: TerminateParams, + ) -> Result { + let running = { + let mut process_map = self.inner.processes.lock().await; + match process_map.get_mut(¶ms.process_id) { + Some(ProcessEntry::Running(process)) => { + if let Some(network_policy_shutdown) = &process.network_policy_shutdown { + network_policy_shutdown.cancel(); + } + if process.exit_code.is_some() { + return Ok(TerminateResponse { running: false }); + } + process.termination_requested = true; + process.session.terminate(); + true + } + Some(ProcessEntry::Starting(_)) => { + process_map.remove(¶ms.process_id); + true + } + None => false, + } + }; + + Ok(TerminateResponse { running }) + } +} + +fn child_env(params: &ExecParams) -> HashMap { + let mut env = match ¶ms.env_policy { + Some(env_policy) => { + let policy = shell_environment_policy(env_policy); + let mut env = shell_environment::create_env(&policy, /*thread_id*/ None); + env.extend(params.env.clone()); + env + } + None => params.env.clone(), + }; + env.remove(crate::CODEX_EXEC_SERVER_EXIT_ON_STDIN_CLOSE_ENV_VAR); + env.retain(|name, _| !shell_environment::is_non_inheritable_env_var(name)); + env +} + +pub(crate) fn shell_environment_policy(env_policy: &ExecEnvPolicy) -> ShellEnvironmentPolicy { + ShellEnvironmentPolicy { + inherit: env_policy.inherit.clone(), + ignore_default_excludes: env_policy.ignore_default_excludes, + exclude: env_policy + .exclude + .iter() + .map(|pattern| EnvironmentVariablePattern::new_case_insensitive(pattern)) + .collect(), + r#set: env_policy.r#set.clone(), + include_only: env_policy + .include_only + .iter() + .map(|pattern| EnvironmentVariablePattern::new_case_insensitive(pattern)) + .collect(), + use_profile: false, + } +} + +impl LocalProcess { + async fn start(&self, params: ExecParams) -> Result { + let (response, wake_tx, events) = self + .start_process(params, ProcessTelemetry::default()) + .await + .map_err(map_handler_error)?; + let sandbox_type = sandbox_type_from_protocol(response.sandbox_type); + Ok(StartedExecProcess { + process: Arc::new(LocalExecProcess { + process_id: response.process_id, + backend: self.clone(), + wake_tx, + events, + }), + sandbox_type, + }) + } +} + +impl ExecBackend for LocalProcess { + fn start(&self, params: ExecParams) -> ExecBackendFuture<'_> { + Box::pin(LocalProcess::start(self, params)) + } + + #[cfg(unix)] + fn prewarm_shell_snapshot(&self, params: ExecParams) -> ExecProcessFuture<'_, ()> { + Box::pin(async move { + if params.enforce_managed_network + || params.managed_network.is_some() + || params.network_proxy.is_some() + { + return Err(ExecServerError::Protocol( + "shell snapshot prewarming does not support managed networking".to_string(), + )); + } + let mut prepared = prepare_exec_request_with_telemetry( + ¶ms, + child_env(¶ms), + self.runtime_paths.as_ref(), + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + &ProcessTelemetry::default(), + ) + .await + .map_err(map_handler_error)?; + self.inner + .shell_snapshots + .prepare( + ¶ms, + &mut prepared, + &self.inner.telemetry, + CapturePurpose::Prewarm, + ) + .await + .map_err(map_handler_error) + }) + } +} + +impl LocalExecProcess { + async fn read( + &self, + after_seq: Option, + max_bytes: Option, + wait_ms: Option, + ) -> Result { + self.backend + .read(&self.process_id, after_seq, max_bytes, wait_ms) + .await + } + + async fn write(&self, chunk: Vec) -> Result { + self.backend.write(&self.process_id, chunk).await + } + + async fn signal(&self, signal: ProcessSignal) -> Result<(), ExecServerError> { + self.backend.signal(&self.process_id, signal).await + } + + async fn terminate(&self) -> Result<(), ExecServerError> { + self.backend.terminate(&self.process_id).await + } +} + +impl ExecProcess for LocalExecProcess { + fn process_id(&self) -> &ProcessId { + &self.process_id + } + + fn subscribe_wake(&self) -> watch::Receiver { + self.wake_tx.subscribe() + } + + fn subscribe_events(&self) -> ExecProcessEventReceiver { + self.events.subscribe() + } + + fn read( + &self, + after_seq: Option, + max_bytes: Option, + wait_ms: Option, + ) -> ExecProcessFuture<'_, ReadResponse> { + Box::pin(LocalExecProcess::read(self, after_seq, max_bytes, wait_ms)) + } + + fn write(&self, chunk: Vec) -> ExecProcessFuture<'_, WriteResponse> { + Box::pin(LocalExecProcess::write(self, chunk)) + } + + fn signal(&self, signal: ProcessSignal) -> ExecProcessFuture<'_, ()> { + Box::pin(LocalExecProcess::signal(self, signal)) + } + + fn terminate(&self) -> ExecProcessFuture<'_, ()> { + Box::pin(LocalExecProcess::terminate(self)) + } +} + +impl LocalProcess { + async fn read( + &self, + process_id: &ProcessId, + after_seq: Option, + max_bytes: Option, + wait_ms: Option, + ) -> Result { + self.exec_read(ReadParams { + process_id: process_id.clone(), + after_seq, + max_bytes, + wait_ms, + }) + .await + .map_err(map_handler_error) + } + + async fn write( + &self, + process_id: &ProcessId, + chunk: Vec, + ) -> Result { + self.exec_write(WriteParams { + process_id: process_id.clone(), + chunk: chunk.into(), + write_id: format!( + "local-{}", + NEXT_LOCAL_STDIN_WRITE_ID.fetch_add(1, Ordering::Relaxed) + ), + }) + .await + .map_err(map_handler_error) + } + + async fn signal( + &self, + process_id: &ProcessId, + signal: ProcessSignal, + ) -> Result<(), ExecServerError> { + self.signal_process(SignalParams { + process_id: process_id.clone(), + signal, + }) + .await + .map_err(map_handler_error)?; + Ok(()) + } + + async fn terminate(&self, process_id: &ProcessId) -> Result<(), ExecServerError> { + self.terminate_process(TerminateParams { + process_id: process_id.clone(), + }) + .await + .map_err(map_handler_error)?; + Ok(()) + } +} + +fn pty_process_signal(signal: ProcessSignal) -> PtyProcessSignal { + match signal { + ProcessSignal::Interrupt => PtyProcessSignal::Interrupt, + } +} + +fn map_handler_error(error: JSONRPCErrorError) -> ExecServerError { + ExecServerError::Server { + code: error.code, + message: error.message, + } +} + +async fn stream_output( + process_id: ProcessId, + stream: ExecOutputStream, + mut receiver: tokio::sync::mpsc::Receiver>, + inner: Arc, + output_notify: Arc, +) { + while let Some(chunk) = receiver.recv().await { + let _chunk_len = chunk.len(); + let notification = { + let mut processes = inner.processes.lock().await; + let Some(entry) = processes.get_mut(&process_id) else { + break; + }; + let ProcessEntry::Running(process) = entry else { + break; + }; + let seq = process.next_seq; + process.next_seq += 1; + process.retained_bytes += chunk.len(); + process.output.push_back(RetainedOutputChunk { + seq, + stream, + chunk: chunk.clone(), + }); + while process.retained_bytes > RETAINED_OUTPUT_BYTES_PER_PROCESS + || process.output.len() > RETAINED_OUTPUT_CHUNKS_PER_PROCESS + { + let Some(evicted) = process.output.pop_front() else { + break; + }; + process.retained_bytes = process.retained_bytes.saturating_sub(evicted.chunk.len()); + } + let _ = process.wake_tx.send(seq); + let output = ProcessOutputChunk { + seq, + stream, + chunk: chunk.into(), + }; + process + .events + .publish(ExecProcessEvent::Output(output.clone())); + ExecOutputDeltaNotification { + process_id: process_id.clone(), + seq, + stream, + chunk: output.chunk, + } + }; + output_notify.notify_waiters(); + if let Some(notifications) = notification_sender(&inner) { + let _ = notifications + .notify(crate::protocol::EXEC_OUTPUT_DELTA_METHOD, ¬ification) + .await; + } + } + + finish_output_stream(process_id, inner).await; +} + +fn watch_exit( + process_id: ProcessId, + exit_rx: tokio::sync::oneshot::Receiver, + inner: Arc, + output_notify: Arc, + telemetry: ProcessTelemetry, +) -> impl std::future::Future + Send { + // Set the copied OTEL parent before entering; never retain the RPC tracing span. + let process_span = tracing::info_span!(parent: None, "codex.exec_server.process"); + if let Some(launch_context) = &telemetry.launch_context { + codex_otel::set_parent_from_context( + &process_span, + opentelemetry::Context::new().with_remote_span_context(launch_context.clone()), + ); + } + async move { + let exit_code = exit_rx.await.unwrap_or(-1); + let sandboxed = { + let mut processes = inner.processes.lock().await; + match processes.get_mut(&process_id) { + Some(ProcessEntry::Running(process)) => { + let sandboxed = process.sandbox != SandboxType::None; + if let Some(metrics) = process.metrics.take() { + metrics.finish(if process.termination_requested { + "terminated" + } else if exit_code == 0 { + "success" + } else { + "error" + }); + } + sandboxed + } + Some(ProcessEntry::Starting(_)) | None => false, + } + }; + if sandboxed { + let _ = tokio::time::timeout(Duration::from_millis(20), output_notify.notified()).await; + } + let notification = { + let mut processes = inner.processes.lock().await; + if let Some(ProcessEntry::Running(process)) = processes.get_mut(&process_id) { + let seq = process.next_seq; + process.next_seq += 1; + process.exit_code = Some(exit_code); + if process.sandbox != SandboxType::None { + let mut stdout = Vec::new(); + let mut stderr = Vec::new(); + let mut aggregated = Vec::new(); + for chunk in &process.output { + match chunk.stream { + ExecOutputStream::Stdout | ExecOutputStream::Pty => { + stdout.extend_from_slice(&chunk.chunk); + } + ExecOutputStream::Stderr => stderr.extend_from_slice(&chunk.chunk), + } + aggregated.extend_from_slice(&chunk.chunk); + } + let exec_output = ExecToolCallOutput { + exit_code, + stdout: StreamOutput::new(String::from_utf8_lossy(&stdout).into_owned()), + stderr: StreamOutput::new(String::from_utf8_lossy(&stderr).into_owned()), + aggregated_output: StreamOutput::new( + String::from_utf8_lossy(&aggregated).into_owned(), + ), + ..Default::default() + }; + // Keep the classification in the result for caller approval/retry handling. + process.sandbox_denied = + is_likely_sandbox_denied(process.sandbox, &exec_output); + if process.sandbox_denied { + telemetry.log(ProcessTelemetryEvent::SandboxDenied, process.sandbox); + } + } + telemetry.log( + ProcessTelemetryEvent::Exit { + exit_code, + termination_requested: process.termination_requested, + }, + process.sandbox, + ); + let _ = process.wake_tx.send(seq); + process.events.publish(ExecProcessEvent::Exited { + seq, + exit_code, + sandbox_denied: Some(process.sandbox_denied), + }); + Some(ExecExitedNotification { + process_id: process_id.clone(), + seq, + exit_code, + sandbox_denied: Some(process.sandbox_denied), + }) + } else { + None + } + }; + output_notify.notify_waiters(); + if let Some(notification) = notification + && let Some(notifications) = notification_sender(&inner) + { + let _ = notifications + .notify(crate::protocol::EXEC_EXITED_METHOD, ¬ification) + .await; + } + + maybe_emit_closed(process_id, Arc::clone(&inner)).await; + } + .instrument(process_span) +} + +async fn finish_output_stream(process_id: ProcessId, inner: Arc) { + { + let mut processes = inner.processes.lock().await; + let Some(ProcessEntry::Running(process)) = processes.get_mut(&process_id) else { + return; + }; + + if process.open_streams > 0 { + process.open_streams -= 1; + } + } + + maybe_emit_closed(process_id, inner).await; +} + +async fn maybe_emit_closed(process_id: ProcessId, inner: Arc) { + let (notification, output_notify, network_proxy_handle) = { + let mut processes = inner.processes.lock().await; + let Some(ProcessEntry::Running(process)) = processes.get_mut(&process_id) else { + return; + }; + + if process.closed || process.open_streams != 0 || process.exit_code.is_none() { + return; + } + + process.closed = true; + if let Some(network_policy_shutdown) = process.network_policy_shutdown.take() { + network_policy_shutdown.cancel(); + } + let seq = process.next_seq; + process.next_seq += 1; + let _ = process.wake_tx.send(seq); + process.events.publish(ExecProcessEvent::Closed { seq }); + ( + ExecClosedNotification { + process_id: process_id.clone(), + seq, + }, + Arc::clone(&process.output_notify), + process.network_proxy_handle.take(), + ) + }; + + if let Some(network_proxy_handle) = network_proxy_handle + && let Err(err) = network_proxy_handle.shutdown().await + { + tracing::warn!("failed to shut down executor network proxy: {err}"); + } + + output_notify.notify_waiters(); + let cleanup_process_id = process_id.clone(); + let cleanup_inner = Arc::clone(&inner); + tokio::spawn(async move { + tokio::time::sleep(EXITED_PROCESS_RETENTION).await; + let mut processes = cleanup_inner.processes.lock().await; + match processes.entry(cleanup_process_id) { + Entry::Occupied(entry) => { + if matches!(entry.get(), ProcessEntry::Running(process) if process.closed) { + entry.remove(); + } + } + Entry::Vacant(_) => {} + } + }); + + if let Some(notifications) = notification_sender(&inner) { + let _ = notifications + .notify(EXEC_CLOSED_METHOD, ¬ification) + .await; + } +} + +fn notification_sender(inner: &Inner) -> Option { + inner + .notifications + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() +} + +#[cfg(test)] +mod tests { + use super::*; + use codex_exec_server_protocol::JSONRPCMessage; + use codex_exec_server_protocol::JSONRPCResponse; + use codex_exec_server_protocol::RequestId; + use codex_network_proxy::NetworkProxy; + use codex_network_proxy::NetworkProxyConfig; + use codex_network_proxy::NetworkProxyState; + use codex_network_proxy::RemoteNetworkProxyConfig; + use codex_network_proxy::RemoteNetworkProxyLaunchConfig; + use codex_otel::MetricsConfig; + use codex_protocol::config_types::ShellEnvironmentPolicyInherit; + use codex_utils_path_uri::PathUri; + use codex_utils_pty::ProcessDriver; + use opentelemetry_sdk::metrics::InMemoryMetricExporter; + use opentelemetry_sdk::metrics::data::AggregatedMetrics; + use opentelemetry_sdk::metrics::data::MetricData; + use pretty_assertions::assert_eq; + #[cfg(not(target_os = "windows"))] + use tokio::io::AsyncReadExt; + #[cfg(not(target_os = "windows"))] + use tokio::io::AsyncWriteExt; + use tokio::sync::oneshot; + use tokio::time::timeout; + + #[cfg(not(target_os = "windows"))] + use crate::protocol::ExecServerNetworkPolicyDecision; + #[cfg(not(target_os = "windows"))] + use crate::protocol::NETWORK_POLICY_REQUEST_METHOD; + #[cfg(not(target_os = "windows"))] + use crate::protocol::NetworkPolicyRequestParams; + #[cfg(not(target_os = "windows"))] + use crate::protocol::NetworkPolicyRequestResponse; + + fn test_exec_params(env: HashMap) -> ExecParams { + ExecParams { + metadata: None, + process_id: ProcessId::from("env-test"), + argv: vec!["true".to_string()], + cwd: PathUri::from_host_native_path(std::env::current_dir().expect("cwd")) + .expect("cwd URI"), + shell_snapshot: None, + env_policy: None, + env, + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + } + } + + #[cfg(unix)] + #[tokio::test] + async fn executor_proxy_sends_final_network_policy_notification() { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(NOTIFICATION_CHANNEL_CAPACITY); + let backend = LocalProcess::with_runtime_paths( + RpcNotificationSender::new(outgoing_tx), + ExecServerTelemetry::default(), + /*runtime_paths*/ None, + ); + let proxy_config = RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig { + enabled: true, + ..NetworkProxyConfig::default() + }) + .expect("build remote network proxy config"); + let mut params = test_exec_params(HashMap::new()); + params.process_id = ProcessId::from("audit-process"); + params.argv = vec![ + "/bin/sh".to_string(), + "-c".to_string(), + "printf '%s\\n' \"$HTTP_PROXY\"; exec sleep 60".to_string(), + ]; + params.network_proxy = Some( + RemoteNetworkProxyLaunchConfig::new(proxy_config) + .for_execution("environment-1".to_string(), "execution-1".to_string()), + ); + backend + .exec(params, ProcessTelemetry::default()) + .await + .expect("start process with proxy"); + let output = backend + .exec_read(ReadParams { + process_id: ProcessId::from("audit-process"), + after_seq: None, + max_bytes: None, + wait_ms: Some(1_000), + }) + .await + .expect("read executor proxy address"); + let proxy_addr = String::from_utf8( + output + .chunks + .into_iter() + .find(|chunk| matches!(chunk.stream, ExecOutputStream::Stdout)) + .expect("executor proxy address output") + .chunk + .into_inner(), + ) + .expect("UTF-8 proxy address"); + let proxy_addr = proxy_addr + .trim() + .strip_prefix("http://") + .expect("HTTP executor proxy address"); + let mut stream = tokio::net::TcpStream::connect(proxy_addr) + .await + .expect("connect to executor proxy"); + stream + .write_all(b"CONNECT 8.8.8.8:443 HTTP/1.1\r\nHost: 8.8.8.8:443\r\n\r\n") + .await + .expect("write CONNECT request"); + let mut response = [0_u8; 256]; + let response_len = timeout(Duration::from_secs(2), stream.read(&mut response)) + .await + .expect("proxy response timeout") + .expect("read proxy response"); + assert!(String::from_utf8_lossy(&response[..response_len]).starts_with("HTTP/1.1 403")); + + let notification = timeout(Duration::from_secs(2), async { + loop { + match outgoing_rx.recv().await { + Some(RpcServerOutboundMessage::Notification(notification)) + if notification.method == NETWORK_POLICY_DECISION_METHOD => + { + break serde_json::from_value::( + notification + .params + .expect("network policy notification params"), + ) + .expect("deserialize network policy notification"); + } + Some(_) => {} + None => panic!("outbound notifications closed"), + } + } + }) + .await + .expect("network policy notification timeout"); + assert_eq!(notification.process_id, ProcessId::from("audit-process")); + assert_eq!(notification.decision, "deny"); + assert_eq!(notification.host, "8.8.8.8"); + assert_eq!( + notification.protocol, + ExecServerNetworkProtocol::HttpsConnect + ); + backend.shutdown().await; + } + + fn telemetry_backend() -> ( + LocalProcess, + codex_otel::MetricsClient, + InMemoryMetricExporter, + ) { + let exporter = InMemoryMetricExporter::default(); + let metrics = codex_otel::MetricsClient::new(MetricsConfig::in_memory( + "test", + "codex-exec-server", + env!("CARGO_PKG_VERSION"), + exporter.clone(), + )) + .expect("metrics"); + let telemetry = ExecServerTelemetry::new(metrics.clone()); + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(NOTIFICATION_CHANNEL_CAPACITY); + tokio::spawn(async move { while outgoing_rx.recv().await.is_some() {} }); + ( + LocalProcess::with_runtime_paths( + RpcNotificationSender::new(outgoing_tx), + telemetry, + /*runtime_paths*/ None, + ), + metrics, + exporter, + ) + } + + fn assert_finished_process_result( + metrics: codex_otel::MetricsClient, + exporter: &InMemoryMetricExporter, + expected: &str, + ) { + metrics.shutdown().expect("shutdown metrics"); + let resource_metrics = exporter + .get_finished_metrics() + .expect("finished metrics") + .into_iter() + .last() + .expect("metrics export"); + let finished_processes = resource_metrics + .scope_metrics() + .flat_map(opentelemetry_sdk::metrics::data::ScopeMetrics::metrics) + .find(|metric| metric.name() == "exec_server_processes_finished_total") + .expect("finished process metric"); + let AggregatedMetrics::U64(MetricData::Sum(sum)) = finished_processes.data() else { + panic!("finished process metric should be a u64 sum"); + }; + let results = sum + .data_points() + .flat_map(opentelemetry_sdk::metrics::data::SumDataPoint::attributes) + .filter(|attribute| attribute.key.as_str() == "result") + .map(|attribute| attribute.value.as_str().into_owned()) + .collect::>(); + assert_eq!(results, vec![expected.to_string()]); + } + + #[tokio::test] + async fn start_process_rejects_non_native_cwd_before_launch() { + #[cfg(unix)] + let uri = "file://server/share/checkout"; + #[cfg(windows)] + let uri = "file:///usr/local/checkout"; + let cwd = PathUri::parse(uri).expect("non-native cwd URI"); + let source = cwd + .to_abs_path() + .expect_err("cwd should not be native to this host"); + let expected = invalid_params(format!( + "cwd URI `{cwd}` is not valid on this exec-server host: {source}" + )); + let mut params = test_exec_params(HashMap::new()); + params.cwd = cwd; + + let result = LocalProcess::default() + .start_process(params, ProcessTelemetry::default()) + .await; + let Err(error) = result else { + panic!("non-native cwd should be rejected"); + }; + + assert_eq!(error, expected); + } + + #[tokio::test] + async fn callback_enabled_start_bounds_process_id_before_proxy_launch() { + let proxy_config = + RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig::default()) + .expect("remote proxy config"); + let mut proxy = RemoteNetworkProxyLaunchConfig::new(proxy_config); + proxy.policy_decision_timeout_ms = Some(1_000); + let expected = invalid_params(format!( + "callback-enabled process ID must be non-empty and at most {MAX_NETWORK_POLICY_PROCESS_ID_BYTES} bytes" + )); + + for process_id in [ + String::new(), + "p".repeat(MAX_NETWORK_POLICY_PROCESS_ID_BYTES + 1), + ] { + let mut params = test_exec_params(HashMap::new()); + params.process_id = ProcessId::from(process_id); + params.network_proxy = Some(proxy.clone()); + let error = LocalProcess::default() + .start_process(params, ProcessTelemetry::default()) + .await + .err() + .expect("invalid callback process ID should be rejected"); + + assert_eq!(error, expected); + } + + let mut boundary = test_exec_params(HashMap::new()); + boundary.process_id = ProcessId::from("p".repeat(MAX_NETWORK_POLICY_PROCESS_ID_BYTES)); + boundary.network_proxy = Some(proxy); + boundary.argv.clear(); + let error = LocalProcess::default() + .start_process(boundary, ProcessTelemetry::default()) + .await + .err() + .expect("valid boundary process ID should proceed to process preparation"); + #[cfg(not(target_os = "windows"))] + assert!( + error + .message + .contains("executor-local network proxy launch requires an enabled proxy") + ); + #[cfg(target_os = "windows")] + assert_eq!(error, invalid_params("argv must not be empty".to_string())); + + for process_id in [ + String::new(), + "p".repeat(MAX_NETWORK_POLICY_PROCESS_ID_BYTES + 1), + ] { + let mut ordinary = test_exec_params(HashMap::new()); + ordinary.process_id = ProcessId::from(process_id); + ordinary.argv.clear(); + let error = LocalProcess::default() + .start_process(ordinary, ProcessTelemetry::default()) + .await + .err() + .expect("empty argv should be rejected after ID validation"); + + assert_eq!(error, invalid_params("argv must not be empty".to_string())); + } + } + + #[test] + fn child_env_defaults_to_exact_env() { + let params = test_exec_params(HashMap::from([("ONLY_THIS".to_string(), "1".to_string())])); + + assert_eq!( + child_env(¶ms), + HashMap::from([("ONLY_THIS".to_string(), "1".to_string())]) + ); + } + + #[test] + fn child_env_applies_policy_then_overlay() { + let mut params = test_exec_params(HashMap::from([ + ("OVERLAY".to_string(), "overlay".to_string()), + ("POLICY_SET".to_string(), "overlay-wins".to_string()), + ( + "openai_identity_token_file".to_string(), + "/run/identity-token".to_string(), + ), + ])); + params.env_policy = Some(ExecEnvPolicy { + inherit: ShellEnvironmentPolicyInherit::None, + ignore_default_excludes: true, + exclude: Vec::new(), + r#set: HashMap::from([ + ("POLICY_SET".to_string(), "policy".to_string()), + ("OpenAI_Federation_Rule_Id".to_string(), "rule".to_string()), + ]), + include_only: Vec::new(), + }); + + let mut expected = HashMap::from([ + ("OVERLAY".to_string(), "overlay".to_string()), + ("POLICY_SET".to_string(), "overlay-wins".to_string()), + ]); + if cfg!(target_os = "windows") { + expected.insert("PATHEXT".to_string(), ".COM;.EXE;.BAT;.CMD".to_string()); + } + + assert_eq!(child_env(¶ms), expected); + } + + #[tokio::test] + async fn exit_before_shutdown_records_success() { + let (backend, metrics, exporter) = telemetry_backend(); + let mut process = spawn_test_process(&backend, "exit-before-shutdown").await; + + process.exit(/*exit_code*/ 0); + let _ = read_process_until_change(&backend, &process.process_id, /*after_seq*/ None).await; + backend.shutdown().await; + + assert_finished_process_result(metrics, &exporter, "success"); + } + + #[tokio::test] + async fn termination_request_before_exit_records_terminated() { + let (backend, metrics, exporter) = telemetry_backend(); + let mut process = spawn_test_process(&backend, "terminate-before-exit").await; + + assert_eq!( + backend + .terminate_process(TerminateParams { + process_id: process.process_id.clone(), + }) + .await + .expect("terminate process"), + TerminateResponse { running: true }, + ); + process.exit(/*exit_code*/ 0); + let _ = read_process_until_change(&backend, &process.process_id, /*after_seq*/ None).await; + backend.shutdown().await; + + assert_finished_process_result(metrics, &exporter, "terminated"); + } + + #[tokio::test] + async fn termination_request_after_exit_cancels_network_policy_decisions() { + let backend = LocalProcess::default(); + let mut process = spawn_test_process(&backend, "terminate-after-exit").await; + let network_policy_shutdown = CancellationToken::new(); + { + let mut processes = backend.inner.processes.lock().await; + let Some(ProcessEntry::Running(running)) = processes.get_mut(&process.process_id) + else { + panic!("test process should be running"); + }; + running.network_policy_shutdown = Some(network_policy_shutdown.clone()); + } + + process.exit(/*exit_code*/ 0); + let response = + read_process_until_change(&backend, &process.process_id, /*after_seq*/ None).await; + assert!(response.exited); + assert!(!response.closed); + assert!(!network_policy_shutdown.is_cancelled()); + assert_eq!( + backend + .terminate_process(TerminateParams { + process_id: process.process_id.clone(), + }) + .await + .expect("terminate exited process"), + TerminateResponse { running: false }, + ); + assert!(network_policy_shutdown.is_cancelled()); + + drop(process.stdout_tx); + drop(process.stderr_tx); + let _ = read_process_until_closed(&backend, &process.process_id).await; + backend.shutdown().await; + } + + #[tokio::test] + async fn shutdown_before_exit_records_terminated() { + let (backend, metrics, exporter) = telemetry_backend(); + let mut process = spawn_test_process(&backend, "shutdown-before-exit").await; + + backend.shutdown().await; + process.exit(/*exit_code*/ 0); + + assert_finished_process_result(metrics, &exporter, "terminated"); + } + + #[tokio::test] + async fn exited_process_retains_late_output_past_retention() { + let backend = LocalProcess::default(); + let mut process = spawn_test_process(&backend, "proc-late-output").await; + + process.exit(/*exit_code*/ 0); + let exit_response = + read_process_until_change(&backend, &process.process_id, /*after_seq*/ None).await; + assert_eq!( + exit_response, + ReadResponse { + chunks: Vec::new(), + next_seq: 2, + exited: true, + exit_code: Some(0), + closed: false, + failure: None, + sandbox_denied: false, + } + ); + + tokio::time::sleep(EXITED_PROCESS_RETENTION + Duration::from_millis(10)).await; + process + .stdout_tx + .send(b"late output after retention\n".to_vec()) + .await + .expect("send late stdout"); + + let late_response = + read_process_until_change(&backend, &process.process_id, /*after_seq*/ Some(1)).await; + assert_eq!( + late_response.chunks, + vec![ProcessOutputChunk { + seq: 2, + stream: ExecOutputStream::Stdout, + chunk: b"late output after retention\n".to_vec().into(), + }] + ); + assert_eq!(late_response.exit_code, Some(0)); + assert!(!late_response.closed); + + drop(process.stdout_tx); + drop(process.stderr_tx); + let _closed_response = timeout( + Duration::from_secs(1), + read_process_until_closed(&backend, &process.process_id), + ) + .await + .expect("process should close"); + let replay_after_exit = backend + .exec_read(ReadParams { + process_id: process.process_id.clone(), + after_seq: Some(1), + max_bytes: None, + wait_ms: Some(0), + }) + .await + .expect("closed process should remain readable"); + assert_eq!(replay_after_exit.next_seq, 4); + backend.shutdown().await; + } + + #[tokio::test] + async fn process_read_replay_is_bounded_by_chunk_count() { + let backend = LocalProcess::default(); + let process = spawn_test_process(&backend, "proc-chunk-count").await; + let retained_chunk_count = RETAINED_OUTPUT_CHUNKS_PER_PROCESS as u64; + + { + let mut processes = backend.inner.processes.lock().await; + let Some(ProcessEntry::Running(running)) = processes.get_mut(&process.process_id) + else { + panic!("process should be running"); + }; + running.output = (1..=retained_chunk_count) + .map(|seq| RetainedOutputChunk { + seq, + stream: ExecOutputStream::Stdout, + chunk: vec![b'x'], + }) + .collect(); + running.retained_bytes = RETAINED_OUTPUT_CHUNKS_PER_PROCESS; + running.next_seq = retained_chunk_count + 1; + } + + process + .stdout_tx + .send(vec![b'y']) + .await + .expect("send output beyond retained chunk limit"); + timeout(Duration::from_secs(1), async { + loop { + let output_recorded = { + let processes = backend.inner.processes.lock().await; + let Some(ProcessEntry::Running(running)) = processes.get(&process.process_id) + else { + panic!("process should be running"); + }; + running.next_seq == retained_chunk_count + 2 + }; + if output_recorded { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("output should be retained"); + + let response = backend + .exec_read(ReadParams { + process_id: process.process_id.clone(), + after_seq: None, + max_bytes: None, + wait_ms: Some(0), + }) + .await + .expect("read retained output"); + let mut expected_chunks = (2..=retained_chunk_count) + .map(|seq| ProcessOutputChunk { + seq, + stream: ExecOutputStream::Stdout, + chunk: vec![b'x'].into(), + }) + .collect::>(); + expected_chunks.push(ProcessOutputChunk { + seq: retained_chunk_count + 1, + stream: ExecOutputStream::Stdout, + chunk: vec![b'y'].into(), + }); + assert_eq!( + response, + ReadResponse { + chunks: expected_chunks, + next_seq: retained_chunk_count + 2, + exited: false, + exit_code: None, + closed: false, + failure: None, + sandbox_denied: false, + } + ); + + let message = JSONRPCMessage::Response(JSONRPCResponse { + id: RequestId::Integer(1), + result: serde_json::to_value(response).expect("serialize process/read response"), + }); + let encoded = serde_json::to_string(&message).expect("encode JSON-RPC response"); + let decoded = serde_json::from_str::(&encoded) + .expect("retained process/read response should fit the JSON value budget"); + assert_eq!(decoded, message); + + backend.shutdown().await; + } + + #[tokio::test] + async fn exited_process_keeps_network_proxy_until_inherited_streams_close() { + let backend = LocalProcess::default(); + let mut process = spawn_test_process(&backend, "proc-background-child").await; + let (outgoing_tx, outgoing_rx) = mpsc::channel(NOTIFICATION_CHANNEL_CAPACITY); + #[cfg(not(target_os = "windows"))] + let mut outgoing_rx = outgoing_rx; + #[cfg(target_os = "windows")] + let _outgoing_rx = outgoing_rx; + let requests = RpcNotificationSender::new(outgoing_tx).request_sender(); + *backend + .inner + .requests + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(requests.clone()); + let network_policy_shutdown = CancellationToken::new(); + let decider = network_policy_decider( + process.process_id.clone(), + Arc::clone(&backend.inner.requests), + Duration::from_secs(30), + network_policy_shutdown.clone(), + ); + let config = NetworkProxyConfig { + enabled: true, + ..Default::default() + }; + let proxy_config = RemoteNetworkProxyConfig::from_effective_config(&config) + .expect("build remote network proxy config"); + let state = NetworkProxyState::from_remote_launch_config( + RemoteNetworkProxyLaunchConfig::new(proxy_config), + ) + .expect("build network proxy state"); + let proxy = NetworkProxy::builder() + .state(Arc::new(state)) + .policy_decider_arc(decider) + .build() + .await + .expect("build network proxy"); + let handle = proxy.run().await.expect("start network proxy"); + let prepared = proxy + .prepare_for_optional_environment(HashMap::new(), /*environment_id*/ None) + .expect("prepare network proxy environment"); + let proxy_addr: std::net::SocketAddr = prepared + .env + .get("HTTP_PROXY") + .and_then(|value| value.strip_prefix("http://")) + .expect("HTTP proxy address") + .parse() + .expect("parse HTTP proxy address"); + + { + let mut processes = backend.inner.processes.lock().await; + let Some(ProcessEntry::Running(running)) = processes.get_mut(&process.process_id) + else { + panic!("test process should be running"); + }; + running.network_proxy_handle = Some(handle); + running.network_policy_shutdown = Some(network_policy_shutdown.clone()); + } + + process.exit(/*exit_code*/ 0); + let exit_response = + read_process_until_change(&backend, &process.process_id, /*after_seq*/ None).await; + assert!(exit_response.exited); + assert!(!exit_response.closed); + assert!(!network_policy_shutdown.is_cancelled()); + let stream = tokio::net::TcpStream::connect(proxy_addr) + .await + .expect("proxy should remain available to a child holding inherited output streams"); + #[cfg(target_os = "windows")] + { + assert!(proxy.network_proxy_restricting_sid(None).is_some()); + drop(stream); + } + #[cfg(not(target_os = "windows"))] + { + let mut stream = stream; + stream + .write_all(b"CONNECT 8.8.8.8:443 HTTP/1.1\r\nHost: 8.8.8.8:443\r\n\r\n") + .await + .expect("write CONNECT request"); + let outbound = timeout(Duration::from_secs(1), outgoing_rx.recv()) + .await + .expect("policy request should arrive") + .expect("policy request"); + let RpcServerOutboundMessage::Request(request) = outbound else { + panic!("expected policy request"); + }; + assert_eq!(request.method, NETWORK_POLICY_REQUEST_METHOD); + let params: NetworkPolicyRequestParams = + serde_json::from_value(request.params.expect("request params")) + .expect("deserialize policy request"); + assert_eq!(params.process_id, process.process_id); + assert_eq!(params.request.host, "8.8.8.8"); + requests.complete( + request.id, + Ok(serde_json::to_value(NetworkPolicyRequestResponse { + decision: ExecServerNetworkPolicyDecision::Deny { + reason: "not_allowed".to_string(), + }, + }) + .expect("serialize policy response")), + ); + let mut response = [0_u8; 256]; + let response_len = timeout(Duration::from_secs(1), stream.read(&mut response)) + .await + .expect("proxy response timeout") + .expect("read proxy response"); + assert!(String::from_utf8_lossy(&response[..response_len]).starts_with("HTTP/1.1 403")); + } + + drop(process.stdout_tx); + drop(process.stderr_tx); + let closed_response = timeout( + Duration::from_secs(1), + read_process_until_closed(&backend, &process.process_id), + ) + .await + .expect("process should close"); + assert!(closed_response.closed); + assert!(network_policy_shutdown.is_cancelled()); + #[cfg(target_os = "windows")] + assert_eq!(proxy.network_proxy_restricting_sid(None), None); + #[cfg(not(target_os = "windows"))] + assert!(tokio::net::TcpStream::connect(proxy_addr).await.is_err()); + backend.shutdown().await; + } + + #[tokio::test] + async fn closed_process_is_evicted_after_retention() { + let backend = LocalProcess::default(); + let mut process = spawn_test_process(&backend, "proc-closed-eviction").await; + let process_id = process.process_id.clone(); + + process.exit(/*exit_code*/ 0); + drop(process.stdout_tx); + drop(process.stderr_tx); + + let closed_response = timeout( + Duration::from_secs(1), + read_process_until_closed(&backend, &process_id), + ) + .await + .expect("process should close"); + assert!(closed_response.closed); + + timeout(Duration::from_secs(1), async { + loop { + { + let processes = backend.inner.processes.lock().await; + if !processes.contains_key(&process_id) { + break; + } + } + tokio::time::sleep(Duration::from_millis(5)).await; + } + }) + .await + .expect("closed process should be evicted"); + backend.shutdown().await; + } + + struct TestProcess { + process_id: ProcessId, + stdout_tx: mpsc::Sender>, + stderr_tx: mpsc::Sender>, + exit_tx: Option>, + } + + impl TestProcess { + fn exit(&mut self, exit_code: i32) { + self.exit_tx + .take() + .expect("process should not have exited") + .send(exit_code) + .expect("send process exit"); + } + } + + async fn spawn_test_process(backend: &LocalProcess, process_id: &str) -> TestProcess { + let process_id = ProcessId::from(process_id); + let (stdout_tx, stdout_rx) = mpsc::channel(16); + let (stderr_tx, stderr_rx) = mpsc::channel(16); + let (exit_tx, exit_rx) = oneshot::channel(); + let output_notify = Arc::new(Notify::new()); + let (wake_tx, _wake_rx) = watch::channel(0); + let events = ExecProcessEventLog::new( + PROCESS_EVENT_CHANNEL_CAPACITY, + RETAINED_OUTPUT_BYTES_PER_PROCESS, + ); + + let mut processes = backend.inner.processes.lock().await; + let previous = processes.insert( + process_id.clone(), + ProcessEntry::Running(Box::new(RunningProcess { + session: dummy_session(), + tty: false, + pipe_stdin: false, + accepted_stdin_write_ids: Arc::new(Mutex::new(AcceptedStdinWriteIds::default())), + output: VecDeque::new(), + retained_bytes: 0, + next_seq: 1, + exit_code: None, + wake_tx: wake_tx.clone(), + events: events.clone(), + output_notify: Arc::clone(&output_notify), + open_streams: 2, + closed: false, + metrics: Some(backend.inner.telemetry.process_started(&process_id)), + termination_requested: false, + sandbox: SandboxType::None, + sandbox_denied: false, + network_proxy_handle: None, + network_policy_shutdown: None, + })), + ); + assert!(previous.is_none()); + drop(processes); + + tokio::spawn(stream_output( + process_id.clone(), + ExecOutputStream::Stdout, + stdout_rx, + Arc::clone(&backend.inner), + Arc::clone(&output_notify), + )); + tokio::spawn(stream_output( + process_id.clone(), + ExecOutputStream::Stderr, + stderr_rx, + Arc::clone(&backend.inner), + Arc::clone(&output_notify), + )); + tokio::spawn(watch_exit( + process_id.clone(), + exit_rx, + Arc::clone(&backend.inner), + output_notify, + ProcessTelemetry::default(), + )); + + TestProcess { + process_id, + stdout_tx, + stderr_tx, + exit_tx: Some(exit_tx), + } + } + + fn dummy_session() -> ExecCommandSession { + let (writer_tx, _writer_rx) = mpsc::channel(1); + let (_stdout_tx, stdout_rx) = tokio::sync::broadcast::channel(1); + let (_stderr_tx, stderr_rx) = tokio::sync::broadcast::channel(1); + let (_exit_tx, exit_rx) = oneshot::channel(); + + codex_utils_pty::spawn_from_driver(ProcessDriver { + writer_tx, + stdout_rx, + stderr_rx: Some(stderr_rx), + exit_rx, + terminator: None, + writer_handle: None, + resizer: None, + #[cfg(windows)] + tty: false, + }) + .session + } + + async fn read_process_until_change( + backend: &LocalProcess, + process_id: &ProcessId, + after_seq: Option, + ) -> ReadResponse { + timeout( + Duration::from_secs(1), + backend.exec_read(ReadParams { + process_id: process_id.clone(), + after_seq, + max_bytes: None, + wait_ms: Some(1_000), + }), + ) + .await + .expect("process read should finish") + .expect("process read") + } + + async fn read_process_until_closed( + backend: &LocalProcess, + process_id: &ProcessId, + ) -> ReadResponse { + let mut after_seq = None; + loop { + let response = read_process_until_change(backend, process_id, after_seq).await; + if response.closed { + return response; + } + for chunk in &response.chunks { + after_seq = Some(chunk.seq); + } + after_seq = response.next_seq.checked_sub(1).or(after_seq); + } + } +} diff --git a/codex-rs/exec-server/src/network_policy_decisions.rs b/codex-rs/exec-server/src/network_policy_decisions.rs new file mode 100644 index 0000000000000000000000000000000000000000..129ae30c33468fc70c8005d68a86f2db7a0d830d --- /dev/null +++ b/codex-rs/exec-server/src/network_policy_decisions.rs @@ -0,0 +1,97 @@ +use std::sync::Arc; +use std::sync::RwLock; +use std::time::Duration; + +use codex_network_proxy::NetworkDecision; +use codex_network_proxy::NetworkPolicyDecider; +use codex_network_proxy::NetworkPolicyRequest; +use codex_network_proxy::NetworkProtocol; +use tokio_util::sync::CancellationToken; + +use crate::ProcessId; +use crate::protocol::ExecServerNetworkPolicyDecision; +use crate::protocol::ExecServerNetworkPolicyRequest; +use crate::protocol::ExecServerNetworkProtocol; +use crate::protocol::MAX_NETWORK_POLICY_HOST_BYTES; +use crate::protocol::MAX_NETWORK_POLICY_REASON_BYTES; +use crate::protocol::NETWORK_POLICY_REQUEST_METHOD; +use crate::protocol::NetworkPolicyRequestParams; +use crate::protocol::NetworkPolicyRequestResponse; +use crate::rpc_server_requests::RpcServerRequestSender; + +const NETWORK_POLICY_TRANSPORT_TIMEOUT_MARGIN: Duration = Duration::from_secs(5); + +pub(crate) fn network_policy_decider( + process_id: ProcessId, + requests: Arc>>, + controller_timeout: Duration, + process_shutdown: CancellationToken, +) -> Arc { + let request_timeout = + controller_timeout.saturating_add(NETWORK_POLICY_TRANSPORT_TIMEOUT_MARGIN); + Arc::new(move |request: NetworkPolicyRequest| { + let process_id = process_id.clone(); + let requests = Arc::clone(&requests); + let process_shutdown = process_shutdown.clone(); + async move { + let host = request.host.as_str(); + if host.is_empty() + || host.len() > MAX_NETWORK_POLICY_HOST_BYTES + || host.chars().any(char::is_control) + || host.chars().any(char::is_whitespace) + { + return NetworkDecision::deny("not_allowed"); + } + let requests = requests + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone(); + let Some(requests) = requests else { + return NetworkDecision::deny("not_allowed"); + }; + let params = NetworkPolicyRequestParams { + process_id, + request: ExecServerNetworkPolicyRequest { + protocol: match request.protocol { + NetworkProtocol::Http => ExecServerNetworkProtocol::Http, + NetworkProtocol::HttpsConnect => ExecServerNetworkProtocol::HttpsConnect, + NetworkProtocol::Socks5Tcp => ExecServerNetworkProtocol::Socks5Tcp, + NetworkProtocol::Socks5Udp => ExecServerNetworkProtocol::Socks5Udp, + }, + host: request.host, + port: request.port, + }, + }; + tokio::select! { + biased; + _ = process_shutdown.cancelled() => NetworkDecision::deny("not_allowed"), + response = requests.call_with_timeout::<_, NetworkPolicyRequestResponse>( + NETWORK_POLICY_REQUEST_METHOD, + ¶ms, + request_timeout, + ) => response + .map(|response| match response.decision { + ExecServerNetworkPolicyDecision::Allow => NetworkDecision::Allow, + ExecServerNetworkPolicyDecision::Deny { reason } + | ExecServerNetworkPolicyDecision::Ask { reason } + if reason.len() > MAX_NETWORK_POLICY_REASON_BYTES + || reason.chars().any(char::is_control) => + { + NetworkDecision::deny("not_allowed") + } + ExecServerNetworkPolicyDecision::Deny { reason } => { + NetworkDecision::deny(reason) + } + ExecServerNetworkPolicyDecision::Ask { reason } => { + NetworkDecision::ask(reason) + } + }) + .unwrap_or_else(|_| NetworkDecision::deny("not_allowed")), + } + } + }) +} + +#[cfg(test)] +#[path = "network_policy_decisions_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/network_policy_decisions_tests.rs b/codex-rs/exec-server/src/network_policy_decisions_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..ace0c3e1f3b09b747ccc4625645e29c8f8fb9a45 --- /dev/null +++ b/codex-rs/exec-server/src/network_policy_decisions_tests.rs @@ -0,0 +1,225 @@ +use std::sync::Arc; +use std::sync::RwLock; +use std::time::Duration; + +use codex_network_proxy::NetworkDecision; +use codex_network_proxy::NetworkPolicyRequest; +use codex_network_proxy::NetworkPolicyRequestArgs; +use codex_network_proxy::NetworkProtocol; +use pretty_assertions::assert_eq; +use tokio::sync::mpsc; +use tokio::task::JoinHandle; +use tokio::time::timeout; + +use super::*; +use crate::protocol::ExecServerNetworkPolicyDecision; +use crate::protocol::NetworkPolicyRequestParams; +use crate::protocol::NetworkPolicyRequestResponse; +use crate::rpc::RpcServerOutboundMessage; + +struct DeciderHarness { + requests: RpcServerRequestSender, + outgoing: mpsc::Receiver, + controller_timeout: Duration, + process_shutdown: CancellationToken, +} + +impl DeciderHarness { + fn new() -> Self { + let (outgoing_tx, outgoing) = mpsc::channel(/*buffer*/ 8); + Self { + requests: RpcServerRequestSender::new(outgoing_tx), + outgoing, + controller_timeout: Duration::from_secs(60), + process_shutdown: CancellationToken::new(), + } + } + + fn request(&self, host: &str) -> JoinHandle { + let decider = network_policy_decider( + ProcessId::from("process"), + Arc::new(RwLock::new(Some(self.requests.clone()))), + self.controller_timeout, + self.process_shutdown.clone(), + ); + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::HttpsConnect, + host: host.to_string(), + port: 443, + environment_id: None, + client_addr: None, + method: None, + command: None, + exec_policy_hint: None, + }); + tokio::spawn(async move { decider.decide(request).await }) + } + + async fn next_request( + &mut self, + ) -> ( + codex_exec_server_protocol::RequestId, + NetworkPolicyRequestParams, + ) { + let outbound = timeout(Duration::from_secs(1), self.outgoing.recv()) + .await + .expect("policy request should arrive") + .expect("policy request"); + let RpcServerOutboundMessage::Request(request) = outbound else { + panic!("expected policy request"); + }; + assert_eq!(request.method, NETWORK_POLICY_REQUEST_METHOD); + let params = serde_json::from_value(request.params.expect("request params")) + .expect("deserialize policy request"); + (request.id, params) + } +} + +async fn await_decision(decision: JoinHandle) -> NetworkDecision { + timeout(Duration::from_secs(1), decision) + .await + .expect("network policy decision should resolve") + .expect("network policy decision task") +} + +#[tokio::test] +async fn returns_client_policy_decision() { + let mut harness = DeciderHarness::new(); + let decision = harness.request("example.com"); + let (request_id, params) = harness.next_request().await; + assert_eq!(params.process_id, ProcessId::from("process")); + assert_eq!(params.request.host, "example.com"); + harness.requests.complete( + request_id, + Ok(serde_json::to_value(NetworkPolicyRequestResponse { + decision: ExecServerNetworkPolicyDecision::Allow, + }) + .expect("serialize policy response")), + ); + + assert_eq!(await_decision(decision).await, NetworkDecision::Allow); +} + +#[tokio::test] +async fn policy_response_reasons_are_bounded_and_fail_closed() { + let mut harness = DeciderHarness::new(); + let boundary_reason = "d".repeat(MAX_NETWORK_POLICY_REASON_BYTES); + let cases = [ + ( + ExecServerNetworkPolicyDecision::Deny { + reason: boundary_reason.clone(), + }, + NetworkDecision::deny(boundary_reason), + ), + ( + ExecServerNetworkPolicyDecision::Deny { + reason: "d".repeat(MAX_NETWORK_POLICY_REASON_BYTES + 1), + }, + NetworkDecision::deny("not_allowed"), + ), + ( + ExecServerNetworkPolicyDecision::Ask { + reason: "ask permission".to_string(), + }, + NetworkDecision::ask("ask permission"), + ), + ( + ExecServerNetworkPolicyDecision::Ask { + reason: "ask\npermission".to_string(), + }, + NetworkDecision::deny("not_allowed"), + ), + ]; + + for (response, expected) in cases { + let decision = harness.request("example.com"); + let (request_id, _) = harness.next_request().await; + harness.requests.complete( + request_id, + Ok( + serde_json::to_value(NetworkPolicyRequestResponse { decision: response }) + .expect("serialize policy response"), + ), + ); + assert_eq!(await_decision(decision).await, expected); + } +} + +#[tokio::test] +async fn boundary_host_is_relayed() { + let mut harness = DeciderHarness::new(); + let host = "h".repeat(MAX_NETWORK_POLICY_HOST_BYTES); + let decision = harness.request(&host); + let (request_id, params) = harness.next_request().await; + assert_eq!(params.request.host, host); + harness.requests.complete( + request_id, + Ok(serde_json::to_value(NetworkPolicyRequestResponse { + decision: ExecServerNetworkPolicyDecision::Allow, + }) + .expect("serialize policy response")), + ); + + assert_eq!(await_decision(decision).await, NetworkDecision::Allow); +} + +#[tokio::test] +async fn invalid_hosts_fail_closed_before_reverse_rpc() { + let mut harness = DeciderHarness::new(); + let invalid_hosts = [ + String::new(), + "host name".to_string(), + "host\u{0000}name".to_string(), + "h".repeat(MAX_NETWORK_POLICY_HOST_BYTES + 1), + ]; + + for host in invalid_hosts { + assert_eq!( + await_decision(harness.request(&host)).await, + NetworkDecision::deny("not_allowed") + ); + assert!(harness.outgoing.try_recv().is_err()); + assert_eq!(harness.requests.pending_request_count(), 0); + } +} + +#[tokio::test] +async fn process_exit_and_disconnect_fail_closed() { + let mut process_exit = DeciderHarness::new(); + let process_decision = process_exit.request("process-exit.example.com"); + process_exit.next_request().await; + process_exit.process_shutdown.cancel(); + assert_eq!( + await_decision(process_decision).await, + NetworkDecision::deny("not_allowed") + ); + assert_eq!(process_exit.requests.pending_request_count(), 0); + + let mut disconnect = DeciderHarness::new(); + let disconnect_decision = disconnect.request("disconnect.example.com"); + disconnect.next_request().await; + disconnect.requests.close(); + assert_eq!( + await_decision(disconnect_decision).await, + NetworkDecision::deny("not_allowed") + ); + assert_eq!(disconnect.requests.pending_request_count(), 0); +} + +#[tokio::test(start_paused = true)] +async fn configured_decision_timeout_fails_closed() { + let mut harness = DeciderHarness::new(); + harness.controller_timeout = Duration::from_secs(17); + let decision = harness.request("timeout.example.com"); + harness.next_request().await; + + tokio::time::advance(Duration::from_secs(21)).await; + assert!(!decision.is_finished()); + + tokio::time::advance(Duration::from_secs(1)).await; + assert_eq!( + await_decision(decision).await, + NetworkDecision::deny("not_allowed") + ); + assert_eq!(harness.requests.pending_request_count(), 0); +} diff --git a/codex-rs/exec-server/src/no_follow/mod.rs b/codex-rs/exec-server/src/no_follow/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..54461e887c9d7e105a2edb20e2bae4a255a5c93d --- /dev/null +++ b/codex-rs/exec-server/src/no_follow/mod.rs @@ -0,0 +1,60 @@ +use crate::FileMetadata; +use std::io; +use std::path::Path; +#[cfg(windows)] +use std::time::SystemTime; +#[cfg(windows)] +use std::time::UNIX_EPOCH; + +#[cfg(unix)] +mod unix; +#[cfg(windows)] +mod windows; + +#[cfg(unix)] +use unix as imp; +#[cfg(windows)] +use windows as imp; + +pub(crate) async fn open_file(path: &Path) -> io::Result { + imp::open_file(path.to_path_buf()).await +} + +pub(crate) async fn write_file(path: &Path, contents: Vec) -> io::Result<()> { + imp::write_file(path.to_path_buf(), contents).await +} + +#[cfg(unix)] +pub(crate) async fn metadata(path: &Path) -> io::Result { + imp::metadata(path.to_path_buf()).await +} + +#[cfg(windows)] +pub(crate) async fn metadata(path: &Path) -> io::Result { + imp::metadata(path.to_path_buf()) + .await + .map(|metadata| FileMetadata { + is_directory: metadata.is_dir(), + is_file: metadata.is_file(), + is_symlink: false, + size: metadata.len(), + created_at_ms: metadata.created().ok().map_or(0, system_time_to_unix_ms), + modified_at_ms: metadata.modified().ok().map_or(0, system_time_to_unix_ms), + }) +} + +#[cfg(windows)] +fn system_time_to_unix_ms(time: SystemTime) -> i64 { + time.duration_since(UNIX_EPOCH) + .ok() + .and_then(|duration| i64::try_from(duration.as_millis()).ok()) + .unwrap_or(0) +} + +pub(crate) async fn create_directory(path: &Path, recursive: bool) -> io::Result<()> { + imp::create_directory(path.to_path_buf(), recursive).await +} + +pub(crate) async fn remove(path: &Path, recursive: bool, force: bool) -> io::Result<()> { + imp::remove(path.to_path_buf(), recursive, force).await +} diff --git a/codex-rs/exec-server/src/no_follow/unix.rs b/codex-rs/exec-server/src/no_follow/unix.rs new file mode 100644 index 0000000000000000000000000000000000000000..23a7fcc83a1e2e3daf1586b7d11f690c1d7faab6 --- /dev/null +++ b/codex-rs/exec-server/src/no_follow/unix.rs @@ -0,0 +1,329 @@ +use crate::FileMetadata; +use rustix::fs::AtFlags; +use rustix::fs::Mode; +use rustix::fs::OFlags; +use rustix::fs::Stat; +#[cfg(target_os = "linux")] +use rustix::fs::Statx; +#[cfg(target_os = "linux")] +use rustix::fs::StatxFlags; +use rustix::fs::fstat; +use rustix::fs::mkdirat; +use rustix::fs::open; +use rustix::fs::openat; +use rustix::fs::statat; +#[cfg(target_os = "linux")] +use rustix::fs::statx; +use rustix::fs::unlinkat; +use std::ffi::OsString; +use std::io; +use std::io::Write; +use std::os::fd::OwnedFd; +use std::path::Component; +use std::path::Path; +use std::path::PathBuf; + +#[cfg(any(target_os = "linux", target_os = "android"))] +fn directory_access_flags() -> OFlags { + OFlags::PATH +} + +#[cfg(target_vendor = "apple")] +fn directory_access_flags() -> OFlags { + OFlags::from_bits_retain(libc::O_SEARCH as u32) +} + +#[cfg(not(any(target_os = "linux", target_os = "android", target_vendor = "apple")))] +fn directory_access_flags() -> OFlags { + OFlags::RDONLY +} + +fn components(path: &Path) -> io::Result> { + if !path.is_absolute() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "no-follow filesystem operations require an absolute path", + )); + } + path.components() + .filter_map(|component| match component { + Component::RootDir => None, + Component::Normal(component) => Some(Ok(component.to_os_string())), + Component::CurDir | Component::ParentDir | Component::Prefix(_) => { + Some(Err(io::Error::new( + io::ErrorKind::InvalidInput, + "no-follow filesystem operations require a normalized path", + ))) + } + }) + .collect() +} + +fn root() -> io::Result { + open( + "/", + directory_access_flags() | OFlags::DIRECTORY | OFlags::CLOEXEC, + Mode::empty(), + ) + .map_err(io::Error::from) +} + +fn open_directory(parent: &OwnedFd, name: &OsString) -> io::Result { + openat( + parent, + name, + directory_access_flags() | OFlags::DIRECTORY | OFlags::NOFOLLOW | OFlags::CLOEXEC, + Mode::empty(), + ) + .map_err(io::Error::from) +} + +fn parent(path: &Path) -> io::Result<(OwnedFd, OsString)> { + let mut components = components(path)?; + let leaf = components + .pop() + .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "path must name an entry"))?; + let mut directory = root()?; + for component in components { + directory = open_directory(&directory, &component)?; + } + Ok((directory, leaf)) +} + +fn open_entry_sync(path: &Path) -> io::Result { + let (parent, leaf) = parent(path)?; + let file = openat( + &parent, + leaf, + OFlags::RDONLY | OFlags::NOFOLLOW | OFlags::NONBLOCK | OFlags::CLOEXEC, + Mode::empty(), + ) + .map_err(io::Error::from)?; + Ok(std::fs::File::from(file)) +} + +fn open_file_sync(path: &Path) -> io::Result { + let file = open_entry_sync(path)?; + if !file.metadata()?.is_file() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "path is not a regular file", + )); + } + Ok(file) +} + +pub(super) async fn open_file(path: PathBuf) -> io::Result { + tokio::task::spawn_blocking(move || open_file_sync(&path).map(tokio::fs::File::from_std)) + .await + .map_err(|error| io::Error::other(format!("filesystem task failed: {error}")))? +} + +pub(super) async fn write_file(path: PathBuf, contents: Vec) -> io::Result<()> { + tokio::task::spawn_blocking(move || { + let (parent, leaf) = parent(&path)?; + // Prevent FIFOs and devices from blocking during open before the + // descriptor can be validated as a regular file below. + let file = openat( + &parent, + leaf, + OFlags::WRONLY | OFlags::CREATE | OFlags::NOFOLLOW | OFlags::NONBLOCK | OFlags::CLOEXEC, + Mode::from_raw_mode(0o666), + ) + .map_err(io::Error::from)?; + let mut file = std::fs::File::from(file); + if !file.metadata()?.is_file() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "path is not a regular file", + )); + } + file.set_len(0)?; + file.write_all(&contents) + }) + .await + .map_err(|error| io::Error::other(format!("filesystem task failed: {error}")))? +} + +fn metadata_sync(path: &Path) -> io::Result { + let components = components(path)?; + + #[cfg(target_os = "linux")] + { + let (parent, leaf, flags) = if components.is_empty() { + (root()?, OsString::new(), AtFlags::EMPTY_PATH) + } else { + let (parent, leaf) = parent(path)?; + (parent, leaf, AtFlags::SYMLINK_NOFOLLOW) + }; + let metadata = statx( + &parent, + &leaf, + flags, + StatxFlags::BASIC_STATS | StatxFlags::BTIME, + ); + match metadata { + Ok(metadata) => file_metadata_linux(metadata), + Err(error) if error == rustix::io::Errno::NOSYS || error == rustix::io::Errno::PERM => { + let metadata = if components.is_empty() { + fstat(&parent) + } else { + statat(&parent, &leaf, AtFlags::SYMLINK_NOFOLLOW) + } + .map_err(io::Error::from)?; + file_metadata(metadata, /*created_at_ms*/ 0) + } + Err(error) => Err(io::Error::from(error)), + } + } + + #[cfg(not(target_os = "linux"))] + if components.is_empty() { + let root = root()?; + let metadata = fstat(&root).map_err(io::Error::from)?; + let created_at_ms = created_at_ms(&metadata); + file_metadata(metadata, created_at_ms) + } else { + let (parent, leaf) = parent(path)?; + let metadata = + statat(&parent, &leaf, AtFlags::SYMLINK_NOFOLLOW).map_err(io::Error::from)?; + let created_at_ms = created_at_ms(&metadata); + file_metadata(metadata, created_at_ms) + } +} + +#[cfg(target_os = "linux")] +fn file_metadata_linux(metadata: Statx) -> io::Result { + let kind = u32::from(metadata.stx_mode) & libc::S_IFMT; + if kind == libc::S_IFLNK { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "path contains a symbolic link", + )); + } + let created_at_ms = if metadata.stx_mask & StatxFlags::BTIME.bits() == 0 { + 0 + } else { + unix_time_ms(metadata.stx_btime.tv_sec, metadata.stx_btime.tv_nsec) + }; + Ok(FileMetadata { + is_directory: kind == libc::S_IFDIR, + is_file: kind == libc::S_IFREG, + is_symlink: false, + size: metadata.stx_size, + created_at_ms, + modified_at_ms: unix_time_ms(metadata.stx_mtime.tv_sec, metadata.stx_mtime.tv_nsec), + }) +} + +fn file_metadata(metadata: Stat, created_at_ms: i64) -> io::Result { + let kind = metadata.st_mode & libc::S_IFMT; + if kind == libc::S_IFLNK { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "path contains a symbolic link", + )); + } + Ok(FileMetadata { + is_directory: kind == libc::S_IFDIR, + is_file: kind == libc::S_IFREG, + is_symlink: false, + size: u64::try_from(metadata.st_size).unwrap_or(0), + created_at_ms, + modified_at_ms: unix_time_ms(metadata.st_mtime, metadata.st_mtime_nsec), + }) +} + +#[cfg(target_vendor = "apple")] +fn created_at_ms(metadata: &Stat) -> i64 { + unix_time_ms(metadata.st_birthtime, metadata.st_birthtime_nsec) +} + +#[cfg(not(any(target_vendor = "apple", target_os = "linux")))] +fn created_at_ms(_metadata: &Stat) -> i64 { + 0 +} + +fn unix_time_ms(seconds: i64, nanoseconds: impl TryInto) -> i64 { + let nanoseconds = nanoseconds.try_into().unwrap_or_default(); + seconds + .saturating_mul(1_000) + .saturating_add(nanoseconds / 1_000_000) +} + +pub(super) async fn metadata(path: PathBuf) -> io::Result { + tokio::task::spawn_blocking(move || metadata_sync(&path)) + .await + .map_err(|error| io::Error::other(format!("filesystem task failed: {error}")))? +} + +pub(super) async fn create_directory(path: PathBuf, recursive: bool) -> io::Result<()> { + tokio::task::spawn_blocking(move || { + let components = components(&path)?; + if components.is_empty() && !recursive { + return Err(io::Error::new( + io::ErrorKind::AlreadyExists, + "directory already exists", + )); + } + let mut directory = root()?; + for (index, component) in components.iter().enumerate() { + let is_leaf = index + 1 == components.len(); + if !recursive && is_leaf { + mkdirat(&directory, component, Mode::from_raw_mode(0o777)) + .map_err(io::Error::from)?; + return Ok(()); + } + match open_directory(&directory, component) { + Ok(next) => directory = next, + Err(error) if recursive && error.kind() == io::ErrorKind::NotFound => { + match mkdirat(&directory, component, Mode::from_raw_mode(0o777)) { + Ok(()) | Err(rustix::io::Errno::EXIST) => {} + Err(error) => return Err(io::Error::from(error)), + } + directory = open_directory(&directory, component)?; + } + Err(error) => return Err(error), + } + } + Ok(()) + }) + .await + .map_err(|error| io::Error::other(format!("filesystem task failed: {error}")))? +} + +pub(super) async fn remove(path: PathBuf, recursive: bool, force: bool) -> io::Result<()> { + tokio::task::spawn_blocking(move || { + if recursive { + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "recursive no-follow removal is unsupported", + )); + } + let (parent, leaf) = match parent(&path) { + Ok(value) => value, + Err(error) if force && error.kind() == io::ErrorKind::NotFound => return Ok(()), + Err(error) => return Err(error), + }; + let metadata = match statat(&parent, &leaf, AtFlags::SYMLINK_NOFOLLOW) { + Ok(metadata) => metadata, + Err(error) if force && error == rustix::io::Errno::NOENT => return Ok(()), + Err(error) => return Err(io::Error::from(error)), + }; + let kind = metadata.st_mode & libc::S_IFMT; + if kind == libc::S_IFLNK { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "path contains a symbolic link", + )); + } + let flags = if kind == libc::S_IFDIR { + AtFlags::REMOVEDIR + } else { + AtFlags::empty() + }; + unlinkat(&parent, leaf, flags).map_err(io::Error::from) + }) + .await + .map_err(|error| io::Error::other(format!("filesystem task failed: {error}")))? +} diff --git a/codex-rs/exec-server/src/no_follow/windows.rs b/codex-rs/exec-server/src/no_follow/windows.rs new file mode 100644 index 0000000000000000000000000000000000000000..14a3db30a9feb76cb0fe22ea7e07951d999957c9 --- /dev/null +++ b/codex-rs/exec-server/src/no_follow/windows.rs @@ -0,0 +1,333 @@ +use crate::regular_file; +use std::ffi::OsStr; +use std::ffi::c_void; +use std::io; +use std::io::Write; +use std::mem::size_of; +use std::os::windows::ffi::OsStrExt; +use std::os::windows::io::AsRawHandle; +use std::os::windows::io::FromRawHandle; +use std::os::windows::io::OwnedHandle; +use std::os::windows::io::RawHandle; +use std::path::Component; +use std::path::Path; +use std::path::PathBuf; +use std::path::Prefix; +use std::ptr; +use windows_sys::Win32::Foundation::HANDLE; +use windows_sys::Win32::Foundation::INVALID_HANDLE_VALUE; +use windows_sys::Win32::Foundation::NTSTATUS; +use windows_sys::Win32::Foundation::RtlNtStatusToDosError; +use windows_sys::Win32::Foundation::UNICODE_STRING; +use windows_sys::Win32::Security::SECURITY_QUALITY_OF_SERVICE; +use windows_sys::Win32::Security::SecurityIdentification; +use windows_sys::Win32::Storage::FileSystem::DELETE; +use windows_sys::Win32::Storage::FileSystem::FILE_ATTRIBUTE_NORMAL; +use windows_sys::Win32::Storage::FileSystem::FILE_DISPOSITION_INFO; +use windows_sys::Win32::Storage::FileSystem::FILE_GENERIC_READ; +use windows_sys::Win32::Storage::FileSystem::FILE_READ_ATTRIBUTES; +use windows_sys::Win32::Storage::FileSystem::FILE_SHARE_DELETE; +use windows_sys::Win32::Storage::FileSystem::FILE_SHARE_READ; +use windows_sys::Win32::Storage::FileSystem::FILE_SHARE_WRITE; +use windows_sys::Win32::Storage::FileSystem::FILE_WRITE_DATA; +use windows_sys::Win32::Storage::FileSystem::FileDispositionInfo; +use windows_sys::Win32::Storage::FileSystem::SetFileInformationByHandle; +use windows_sys::Win32::System::IO::IO_STATUS_BLOCK; +use windows_sys::Win32::System::IO::IO_STATUS_BLOCK_0; +use windows_sys::Win32::System::Kernel::OBJ_CASE_INSENSITIVE; +use windows_sys::Win32::System::Kernel::OBJ_DONT_REPARSE; + +const FILE_DIRECTORY_FILE: u32 = 0x0000_0001; +const FILE_SYNCHRONOUS_IO_NONALERT: u32 = 0x0000_0020; +const FILE_NON_DIRECTORY_FILE: u32 = 0x0000_0040; +const FILE_OPEN: u32 = 1; +const FILE_CREATE: u32 = 2; +const FILE_OPEN_IF: u32 = 3; +const SYNCHRONIZE_ACCESS: u32 = 0x0010_0000; +const STATUS_REPARSE_POINT_ENCOUNTERED: NTSTATUS = 0xC000_050B_u32 as i32; +const SECURITY_STATIC_TRACKING: u8 = 0; +const BOOLEAN_TRUE: u8 = 1; + +#[repr(C)] +struct ObjectAttributes { + length: u32, + root_directory: HANDLE, + object_name: *const UNICODE_STRING, + attributes: u32, + security_descriptor: *const c_void, + security_quality_of_service: *const c_void, +} + +#[link(name = "ntdll")] +unsafe extern "system" { + fn NtCreateFile( + file_handle: *mut HANDLE, + desired_access: u32, + object_attributes: *const ObjectAttributes, + io_status_block: *mut IO_STATUS_BLOCK, + allocation_size: *const i64, + file_attributes: u32, + share_access: u32, + create_disposition: u32, + create_options: u32, + ea_buffer: *const c_void, + ea_length: u32, + ) -> NTSTATUS; +} + +fn nt_path(path: &Path) -> io::Result> { + if !path.is_absolute() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "no-follow filesystem operations require an absolute path", + )); + } + + let (prefix, source_offset) = match path.components().next() { + Some(Component::Prefix(prefix)) => match prefix.kind() { + Prefix::Disk(_) => ("\\??\\", 0), + Prefix::VerbatimDisk(_) => ("\\??\\", 4), + Prefix::UNC(_, _) => ("\\??\\UNC\\", 2), + Prefix::VerbatimUNC(_, _) => ("\\??\\UNC\\", 8), + _ => { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "no-follow filesystem operations require a local disk or UNC path", + )); + } + }, + _ => { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "no-follow filesystem operations require an absolute Windows path", + )); + } + }; + + let mut result: Vec = OsStr::new(prefix).encode_wide().collect(); + result.extend( + path.as_os_str() + .encode_wide() + .skip(source_offset) + .map(|unit| { + if unit == b'/' as u16 { + b'\\' as u16 + } else { + unit + } + }), + ); + result.push(0); + Ok(result) +} + +fn open_handle( + path: &Path, + desired_access: u32, + create_disposition: u32, + create_options: u32, +) -> io::Result { + let mut path = nt_path(path)?; + let name_length = u16::try_from((path.len() - 1) * size_of::()) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "filesystem path is too long"))?; + let maximum_length = u16::try_from(path.len() * size_of::()) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "filesystem path is too long"))?; + let object_name = UNICODE_STRING { + Length: name_length, + MaximumLength: maximum_length, + Buffer: path.as_mut_ptr(), + }; + let security_quality_of_service = SECURITY_QUALITY_OF_SERVICE { + Length: size_of::() as u32, + ImpersonationLevel: SecurityIdentification, + ContextTrackingMode: SECURITY_STATIC_TRACKING, + EffectiveOnly: BOOLEAN_TRUE, + }; + let object_attributes = ObjectAttributes { + length: size_of::() as u32, + root_directory: 0, + object_name: &object_name, + attributes: OBJ_CASE_INSENSITIVE as u32 | OBJ_DONT_REPARSE as u32, + security_descriptor: ptr::null(), + security_quality_of_service: (&raw const security_quality_of_service).cast(), + }; + let mut io_status_block = IO_STATUS_BLOCK { + Anonymous: IO_STATUS_BLOCK_0 { Status: 0 }, + Information: 0, + }; + let mut handle = 0; + let status = unsafe { + NtCreateFile( + &mut handle, + desired_access | SYNCHRONIZE_ACCESS, + &object_attributes, + &mut io_status_block, + ptr::null(), + FILE_ATTRIBUTE_NORMAL, + FILE_SHARE_READ | FILE_SHARE_WRITE | FILE_SHARE_DELETE, + create_disposition, + create_options | FILE_SYNCHRONOUS_IO_NONALERT, + ptr::null(), + /*ea_length*/ 0, + ) + }; + if status < 0 { + if status == STATUS_REPARSE_POINT_ENCOUNTERED { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "path contains a reparse point", + )); + } + let code = unsafe { RtlNtStatusToDosError(status) }; + return Err(io::Error::from_raw_os_error(code as i32)); + } + if handle == 0 || handle == INVALID_HANDLE_VALUE { + return Err(io::Error::other( + "NtCreateFile returned an invalid filesystem handle", + )); + } + + Ok(unsafe { OwnedHandle::from_raw_handle(handle as RawHandle) }) +} + +fn open_entry(path: &Path) -> io::Result { + let handle = open_handle( + path, + FILE_READ_ATTRIBUTES, + FILE_OPEN, + /*create_options*/ 0, + )?; + Ok(std::fs::File::from(handle)) +} + +fn open_file_sync(path: &Path) -> io::Result { + let handle = open_handle(path, FILE_GENERIC_READ, FILE_OPEN, FILE_NON_DIRECTORY_FILE)?; + let file = std::fs::File::from(handle); + validate_regular_file(&file, path)?; + Ok(file) +} + +fn validate_regular_file(file: &std::fs::File, path: &Path) -> io::Result<()> { + if !regular_file::is_disk_file(file) || !file.metadata()?.is_file() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("path `{}` is not a file", path.display()), + )); + } + Ok(()) +} + +pub(super) async fn open_file(path: PathBuf) -> io::Result { + tokio::task::spawn_blocking(move || open_file_sync(&path).map(tokio::fs::File::from_std)) + .await + .map_err(|error| io::Error::other(format!("filesystem task failed: {error}")))? +} + +pub(super) async fn write_file(path: PathBuf, contents: Vec) -> io::Result<()> { + tokio::task::spawn_blocking(move || { + let handle = open_handle( + &path, + FILE_READ_ATTRIBUTES | FILE_WRITE_DATA, + FILE_OPEN_IF, + FILE_NON_DIRECTORY_FILE, + )?; + let mut file = std::fs::File::from(handle); + validate_regular_file(&file, &path)?; + file.set_len(0)?; + file.write_all(&contents) + }) + .await + .map_err(|error| io::Error::other(format!("filesystem task failed: {error}")))? +} + +pub(super) async fn metadata(path: PathBuf) -> io::Result { + tokio::task::spawn_blocking(move || open_entry(&path)?.metadata()) + .await + .map_err(|error| io::Error::other(format!("filesystem task failed: {error}")))? +} + +fn create_directory_sync(path: &Path, recursive: bool) -> io::Result<()> { + if !recursive { + open_handle(path, FILE_READ_ATTRIBUTES, FILE_CREATE, FILE_DIRECTORY_FILE)?; + return Ok(()); + } + + let mut current = PathBuf::new(); + for component in path.components() { + current.push(component.as_os_str()); + if matches!(component, Component::Normal(_)) { + open_or_create_directory(¤t)?; + } + } + Ok(()) +} + +fn open_or_create_directory(path: &Path) -> io::Result<()> { + match open_handle(path, FILE_READ_ATTRIBUTES, FILE_OPEN, FILE_DIRECTORY_FILE) { + Ok(_) => Ok(()), + Err(error) if error.kind() == io::ErrorKind::NotFound => { + match open_handle(path, FILE_READ_ATTRIBUTES, FILE_CREATE, FILE_DIRECTORY_FILE) { + Ok(_) => Ok(()), + Err(error) if error.kind() == io::ErrorKind::AlreadyExists => { + open_handle(path, FILE_READ_ATTRIBUTES, FILE_OPEN, FILE_DIRECTORY_FILE)?; + Ok(()) + } + Err(error) => Err(error), + } + } + Err(error) => Err(error), + } +} + +pub(super) async fn create_directory(path: PathBuf, recursive: bool) -> io::Result<()> { + tokio::task::spawn_blocking(move || create_directory_sync(&path, recursive)) + .await + .map_err(|error| io::Error::other(format!("filesystem task failed: {error}")))? +} + +fn remove_sync(path: &Path, recursive: bool, force: bool) -> io::Result<()> { + if recursive { + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "recursive no-follow removal is unsupported", + )); + } + + let metadata = match open_entry(path).and_then(|file| file.metadata()) { + Ok(metadata) => metadata, + Err(error) if force && error.kind() == io::ErrorKind::NotFound => return Ok(()), + Err(error) => return Err(error), + }; + let create_options = if metadata.is_dir() { + FILE_DIRECTORY_FILE + } else { + FILE_NON_DIRECTORY_FILE + }; + let handle = open_handle(path, DELETE, FILE_OPEN, create_options)?; + let disposition = FILE_DISPOSITION_INFO { + DeleteFile: BOOLEAN_TRUE, + }; + let result = unsafe { + SetFileInformationByHandle( + handle.as_raw_handle() as HANDLE, + FileDispositionInfo, + (&raw const disposition).cast(), + size_of::() as u32, + ) + }; + if result == 0 { + return Err(io::Error::last_os_error()); + } + drop(handle); + Ok(()) +} + +pub(super) async fn remove(path: PathBuf, recursive: bool, force: bool) -> io::Result<()> { + tokio::task::spawn_blocking(move || remove_sync(&path, recursive, force)) + .await + .map_err(|error| io::Error::other(format!("filesystem task failed: {error}")))? +} + +#[cfg(test)] +#[path = "windows_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/no_follow/windows_tests.rs b/codex-rs/exec-server/src/no_follow/windows_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..d9e3b0bbeb1abce587f5ebdbc4bc0cd484bb9d86 --- /dev/null +++ b/codex-rs/exec-server/src/no_follow/windows_tests.rs @@ -0,0 +1,43 @@ +use super::*; +use windows_sys::Win32::Storage::FileSystem::FILE_GENERIC_WRITE; +use windows_sys::Win32::Storage::FileSystem::PIPE_ACCESS_DUPLEX; +use windows_sys::Win32::System::Pipes::CreateNamedPipeW; +use windows_sys::Win32::System::Pipes::PIPE_READMODE_BYTE; +use windows_sys::Win32::System::Pipes::PIPE_TYPE_BYTE; +use windows_sys::Win32::System::Pipes::PIPE_WAIT; + +#[test] +fn native_open_rejects_named_pipes_before_connecting() { + let pipe_name = format!("codex-no-follow-{}", uuid::Uuid::new_v4()); + let server_path = PathBuf::from(format!(r"\\.\pipe\{pipe_name}")); + let client_path = PathBuf::from(format!(r"\\localhost\pipe\{pipe_name}")); + let wide_path = server_path + .as_os_str() + .encode_wide() + .chain(std::iter::once(0)) + .collect::>(); + let pipe = unsafe { + CreateNamedPipeW( + wide_path.as_ptr(), + PIPE_ACCESS_DUPLEX, + PIPE_TYPE_BYTE | PIPE_READMODE_BYTE | PIPE_WAIT, + /*nmaxinstances*/ 1, + /*noutbuffersize*/ 1, + /*ninbuffersize*/ 1, + /*ndefaulttimeout*/ 1_000, + ptr::null(), + ) + }; + assert_ne!(pipe, INVALID_HANDLE_VALUE); + let _pipe = unsafe { OwnedHandle::from_raw_handle(pipe as RawHandle) }; + + let error = open_handle( + &client_path, + FILE_GENERIC_READ | FILE_GENERIC_WRITE, + FILE_OPEN, + FILE_NON_DIRECTORY_FILE, + ) + .expect_err("strict native open should reject named pipes"); + assert_eq!(error.kind(), io::ErrorKind::InvalidInput); + assert_eq!(error.to_string(), "path contains a reparse point"); +} diff --git a/codex-rs/exec-server/src/noise_channel.rs b/codex-rs/exec-server/src/noise_channel.rs new file mode 100644 index 0000000000000000000000000000000000000000..cde2aae4eeb1ccedf3d3617bd1e7fa2ba2e21796 --- /dev/null +++ b/codex-rs/exec-server/src/noise_channel.rs @@ -0,0 +1,323 @@ +//! Noise channel used by the remote exec-server relay. +//! +//! The harness initiates hybrid IK and pins the exec-server static key returned +//! by the registry. The first handshake message lets the exec-server authenticate +//! the harness static key; the exec-server then asks the registry whether that +//! key is authorized before completing the handshake. +//! +//! "Hybrid" means the session keys include both X25519 and ML-KEM-768 key +//! agreement. Once the two-message handshake finishes, AES-GCM protects the +//! ordered transport records carrying JSON-RPC. + +use base64::Engine; +use base64::engine::general_purpose::STANDARD; +use clatter::HybridHandshake; +use clatter::HybridHandshakeParams; +use clatter::KeyPair; +use clatter::bytearray::ByteArray; +use clatter::constants::MAX_MESSAGE_LEN; +use clatter::crypto::cipher::AesGcm; +use clatter::crypto::dh::X25519; +use clatter::crypto::hash::Sha256; +use clatter::crypto::kem::rust_crypto_ml_kem::MlKem768; +use clatter::handshakepattern::noise_hybrid_ik; +use clatter::traits::Cipher; +use clatter::traits::Dh; +use clatter::traits::Handshaker; +use clatter::traits::Kem; +use clatter::transportstate::TransportState; +use serde::Deserialize; +use serde::Serialize; + +/// Identifies the handshake pattern and algorithms used by this channel. +pub(crate) const NOISE_CHANNEL_SUITE: &str = "Noise_hybridIK_X25519+MLKEM768_AESGCM_SHA256"; + +const PROLOGUE_DOMAIN: &[u8] = b"codex-exec-server-relay-noise/v1"; + +type Handshake = HybridHandshake; +type Transport = TransportState; +type DhKeyPair = KeyPair<::PubKey, ::PrivateKey>; +type MlKem768PublicKey = ::PubKey; +type KemKeyPair = KeyPair<::PubKey, ::SecretKey>; + +/// Public key material for the exec-server Noise suite. +/// The suite tag prevents keys for another protocol from being accepted just +/// because their components have the expected lengths. +#[derive(Clone, Eq, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct NoiseChannelPublicKey { + suite: String, + x25519_public_key: String, + mlkem768_public_key: String, +} + +impl std::fmt::Debug for NoiseChannelPublicKey { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("NoiseChannelPublicKey") + .field("suite", &self.suite) + .field("x25519_public_key", &"") + .field("mlkem768_public_key", &"") + .finish() + } +} + +impl NoiseChannelPublicKey { + /// Decode registry-provided key material before passing it to Clatter. + fn decode(&self) -> Result<(::PubKey, MlKem768PublicKey), NoiseChannelError> { + if self.suite != NOISE_CHANNEL_SUITE { + return Err(NoiseChannelError::InvalidPublicKey( + "unsupported Noise channel suite", + )); + } + let dh = STANDARD + .decode(&self.x25519_public_key) + .map_err(|_| NoiseChannelError::InvalidPublicKey("invalid X25519 public key"))?; + let dh: ::PubKey = dh + .try_into() + .map_err(|_| NoiseChannelError::InvalidPublicKey("invalid X25519 public key length"))?; + let kem = STANDARD + .decode(&self.mlkem768_public_key) + .map_err(|_| NoiseChannelError::InvalidPublicKey("invalid ML-KEM-768 public key"))?; + if kem.len() != MlKem768PublicKey::LENGTH { + return Err(NoiseChannelError::InvalidPublicKey( + "invalid ML-KEM-768 public key length", + )); + } + + Ok((dh, MlKem768PublicKey::from_slice(kem.as_slice()))) + } +} + +/// Static Noise identity kept for the lifetime of an executor or harness process. +#[derive(Clone)] +pub struct NoiseChannelIdentity { + dh: DhKeyPair, + kem: KemKeyPair, +} + +impl NoiseChannelIdentity { + pub fn generate() -> Result { + let dh = X25519::genkey() + .map_err(|error| NoiseChannelError::KeyGeneration(error.to_string()))?; + let kem = MlKem768::genkey() + .map_err(|error| NoiseChannelError::KeyGeneration(error.to_string()))?; + Ok(Self { dh, kem }) + } + + pub fn public_key(&self) -> NoiseChannelPublicKey { + NoiseChannelPublicKey { + suite: NOISE_CHANNEL_SUITE.to_string(), + x25519_public_key: STANDARD.encode(self.dh.public), + mlkem768_public_key: STANDARD.encode(self.kem.public.as_slice()), + } + } +} + +/// Harness-side state between the two hybrid-IK messages. +/// Consuming it in [`Self::finish`] keeps a handshake tied to one relay stream. +pub(crate) struct InitiatorHandshake { + handshake: Handshake, +} + +impl InitiatorHandshake { + /// Start hybrid IK and pin the expected executor key. + /// `payload` carries the short-lived registry authorization inside the first + /// encrypted handshake message. + pub(crate) fn start( + identity: &NoiseChannelIdentity, + responder_public_key: &NoiseChannelPublicKey, + prologue: &[u8], + payload: &[u8], + ) -> Result<(Self, Vec), NoiseChannelError> { + let (responder_dh, responder_kem) = responder_public_key.decode()?; + + // Both executor key components are pinned before any JSON-RPC is sent. + let params = HybridHandshakeParams::new(noise_hybrid_ik(), true) + .with_prologue(prologue) + .with_s(identity.dh.clone()) + .with_s_kem(identity.kem.clone()) + .with_rs(responder_dh) + .with_rs_kem(responder_kem); + let mut handshake = Handshake::new(params)?; + let overhead = handshake.get_next_message_overhead()?; + if payload.len() > MAX_MESSAGE_LEN - overhead { + return Err(NoiseChannelError::InvalidMessage( + "handshake payload is too large", + )); + } + let mut output = vec![0u8; payload.len() + overhead]; + let output_len = handshake.write_message(payload, &mut output)?; + output.truncate(output_len); + Ok((Self { handshake }, output)) + } + + /// Consume the executor response and enter transport mode. + /// The v1 response does not carry an application payload. + pub(crate) fn finish(mut self, response: &[u8]) -> Result { + ensure_noise_frame_len(response.len(), "handshake response is too large")?; + let overhead = self.handshake.get_next_message_overhead()?; + let mut payload = vec![0u8; response.len().saturating_sub(overhead)]; + let payload_len = self.handshake.read_message(response, &mut payload)?; + if payload_len != 0 { + return Err(NoiseChannelError::InvalidMessage( + "handshake response payload must be empty", + )); + } + Ok(NoiseTransport { + transport: self.handshake.finalize()?, + }) + } +} + +/// Executor-side handshake state while harness authorization is pending. +/// This is not a usable transport until the registry accepts the authenticated +/// harness key. +pub(crate) struct PendingResponderHandshake { + handshake: Handshake, + pub(crate) initiator_public_key: NoiseChannelPublicKey, + pub(crate) payload: Vec, +} + +impl PendingResponderHandshake { + /// Parse the first IK message and recover the authenticated harness key. + /// Callers must authorize that key before calling [`Self::complete`]. + pub(crate) fn read_request( + identity: &NoiseChannelIdentity, + prologue: &[u8], + request: &[u8], + ) -> Result { + ensure_noise_frame_len(request.len(), "handshake request is too large")?; + let params = HybridHandshakeParams::new(noise_hybrid_ik(), false) + .with_prologue(prologue) + .with_s(identity.dh.clone()) + .with_s_kem(identity.kem.clone()); + let mut handshake = Handshake::new(params)?; + let overhead = handshake.get_next_message_overhead()?; + let mut payload = vec![0u8; request.len().saturating_sub(overhead)]; + let payload_len = handshake.read_message(request, &mut payload)?; + // Clatter exposes this key only after the first IK message authenticates. + let remote = handshake + .get_remote_static() + .ok_or(NoiseChannelError::InvalidMessage( + "handshake request is missing initiator static key", + ))?; + let initiator_public_key = NoiseChannelPublicKey { + suite: NOISE_CHANNEL_SUITE.to_string(), + x25519_public_key: STANDARD.encode(remote.dh()), + mlkem768_public_key: STANDARD.encode(remote.kem().as_slice()), + }; + payload.truncate(payload_len); + Ok(Self { + handshake, + initiator_public_key, + payload, + }) + } + + /// Finish the handshake after the registry authorizes the harness key. + pub(crate) fn complete(mut self) -> Result<(NoiseTransport, Vec), NoiseChannelError> { + let overhead = self.handshake.get_next_message_overhead()?; + let mut response = vec![0u8; overhead]; + let response_len = self.handshake.write_message(&[], &mut response)?; + response.truncate(response_len); + Ok(( + NoiseTransport { + transport: self.handshake.finalize()?, + }, + response, + )) + } +} + +/// Established channel with independent implicit send and receive nonces. +/// Relay records must be ordered before decryption, and a logical record must +/// not be encrypted again for retry. +pub(crate) struct NoiseTransport { + transport: Transport, +} + +impl NoiseTransport { + /// Encrypt the next transport record. + pub(crate) fn encrypt(&mut self, plaintext: &[u8]) -> Result, NoiseChannelError> { + let frame_len = plaintext.len().checked_add(AesGcm::tag_len()).ok_or( + NoiseChannelError::InvalidMessage("transport plaintext is too large"), + )?; + ensure_noise_frame_len(frame_len, "transport plaintext is too large")?; + Ok(self.transport.send_vec(plaintext)?) + } + + /// Decrypt the next ordered transport record. + pub(crate) fn decrypt(&mut self, ciphertext: &[u8]) -> Result, NoiseChannelError> { + if ciphertext.len() < AesGcm::tag_len() { + return Err(NoiseChannelError::InvalidMessage( + "transport ciphertext is too short", + )); + } + ensure_noise_frame_len(ciphertext.len(), "transport ciphertext is too large")?; + Ok(self.transport.receive_vec(ciphertext)?) + } +} + +/// Bind the handshake to one environment registration and relay stream. +/// Both peers include these values in the Noise transcript before processing +/// the first handshake message. +pub(crate) fn noise_channel_prologue( + environment_id: &str, + executor_registration_id: &str, + stream_id: &str, +) -> Vec { + let mut prologue = Vec::new(); + append_prologue_part(&mut prologue, PROLOGUE_DOMAIN); + append_prologue_part(&mut prologue, environment_id.as_bytes()); + append_prologue_part(&mut prologue, executor_registration_id.as_bytes()); + append_prologue_part(&mut prologue, stream_id.as_bytes()); + prologue +} + +fn append_prologue_part(prologue: &mut Vec, part: &[u8]) { + // Length prefixes make component boundaries unambiguous. Raw concatenation + // would allow different identifier tuples to produce the same prologue. + let len = part.len() as u64; + prologue.extend_from_slice(&len.to_be_bytes()); + prologue.extend_from_slice(part); +} + +fn ensure_noise_frame_len( + frame_len: usize, + message: &'static str, +) -> Result<(), NoiseChannelError> { + if frame_len > MAX_MESSAGE_LEN { + return Err(NoiseChannelError::InvalidMessage(message)); + } + Ok(()) +} + +#[derive(Debug, thiserror::Error)] +pub enum NoiseChannelError { + #[error("Noise channel key generation failed: {0}")] + KeyGeneration(String), + #[error("invalid Noise channel public key: {0}")] + InvalidPublicKey(&'static str), + #[error("invalid Noise channel message: {0}")] + InvalidMessage(&'static str), + #[error("Noise channel handshake failed: {0}")] + Handshake(String), + #[error("Noise channel transport failed: {0}")] + Transport(String), +} + +impl From for NoiseChannelError { + fn from(error: clatter::error::HandshakeError) -> Self { + Self::Handshake(error.to_string()) + } +} + +impl From for NoiseChannelError { + fn from(error: clatter::error::TransportError) -> Self { + Self::Transport(error.to_string()) + } +} + +#[cfg(test)] +#[path = "noise_channel_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/noise_channel_tests.rs b/codex-rs/exec-server/src/noise_channel_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..b298e8e580c74d057a8677581b1b169fc16ac8bc --- /dev/null +++ b/codex-rs/exec-server/src/noise_channel_tests.rs @@ -0,0 +1,226 @@ +use pretty_assertions::assert_eq; + +use super::InitiatorHandshake; +use super::MAX_MESSAGE_LEN; +use super::NOISE_CHANNEL_SUITE; +use super::NoiseChannelError; +use super::NoiseChannelIdentity; +use super::NoiseChannelPublicKey; +use super::PendingResponderHandshake; +use super::noise_channel_prologue; + +#[test] +fn hybrid_ik_roundtrip_authenticates_both_endpoints() { + let initiator = NoiseChannelIdentity::generate().expect("generate initiator identity"); + let responder = NoiseChannelIdentity::generate().expect("generate responder identity"); + let prologue = noise_channel_prologue("env-1", "registration-1", "stream-1"); + let authorization = b"harness-key-authorization"; + + let (initiator_handshake, request) = InitiatorHandshake::start( + &initiator, + &responder.public_key(), + &prologue, + authorization, + ) + .expect("start initiator handshake"); + let responder_handshake = + PendingResponderHandshake::read_request(&responder, &prologue, &request) + .expect("read responder handshake"); + + assert_eq!( + &responder_handshake.initiator_public_key, + &initiator.public_key() + ); + assert_eq!(responder_handshake.payload.as_slice(), authorization); + + let (mut responder_transport, response) = responder_handshake + .complete() + .expect("complete responder handshake"); + let mut initiator_transport = initiator_handshake + .finish(&response) + .expect("complete initiator handshake"); + + let request_ciphertext = initiator_transport + .encrypt(b"request") + .expect("encrypt request"); + assert_ne!(request_ciphertext, b"request"); + assert_eq!( + responder_transport + .decrypt(&request_ciphertext) + .expect("decrypt request"), + b"request" + ); + + let response_ciphertext = responder_transport + .encrypt(b"response") + .expect("encrypt response"); + assert_ne!(response_ciphertext, b"response"); + assert_eq!( + initiator_transport + .decrypt(&response_ciphertext) + .expect("decrypt response"), + b"response" + ); +} + +#[test] +fn initiator_rejects_wrong_responder_key() { + let initiator = NoiseChannelIdentity::generate().expect("generate initiator identity"); + let expected_responder = NoiseChannelIdentity::generate().expect("generate expected identity"); + let actual_responder = NoiseChannelIdentity::generate().expect("generate actual identity"); + let prologue = noise_channel_prologue("env-1", "registration-1", "stream-1"); + + let (_initiator_handshake, request) = InitiatorHandshake::start( + &initiator, + &expected_responder.public_key(), + &prologue, + b"authorization", + ) + .expect("start initiator handshake"); + + assert!( + PendingResponderHandshake::read_request(&actual_responder, &prologue, &request).is_err() + ); +} + +#[test] +fn responder_rejects_mismatched_prologue() { + let initiator = NoiseChannelIdentity::generate().expect("generate initiator identity"); + let responder = NoiseChannelIdentity::generate().expect("generate responder identity"); + let initiator_prologue = noise_channel_prologue("env-1", "registration-1", "stream-1"); + let responder_prologue = noise_channel_prologue("env-1", "registration-1", "stream-2"); + let (_initiator_handshake, request) = InitiatorHandshake::start( + &initiator, + &responder.public_key(), + &initiator_prologue, + b"authorization", + ) + .expect("start initiator handshake"); + + assert!( + PendingResponderHandshake::read_request(&responder, &responder_prologue, &request).is_err() + ); +} + +#[test] +fn prologue_encoding_is_stable_and_unambiguous() { + let prologue = noise_channel_prologue("env-1", "registration-1", "stream-1"); + + assert_eq!( + prologue, + b"\x00\x00\x00\x00\x00\x00\x00\x20codex-exec-server-relay-noise/v1\ + \x00\x00\x00\x00\x00\x00\x00\x05env-1\ + \x00\x00\x00\x00\x00\x00\x00\x0eregistration-1\ + \x00\x00\x00\x00\x00\x00\x00\x08stream-1" + .to_vec() + ); +} + +#[test] +fn transport_rejects_tampered_ciphertext() { + let initiator = NoiseChannelIdentity::generate().expect("generate initiator identity"); + let responder = NoiseChannelIdentity::generate().expect("generate responder identity"); + let prologue = noise_channel_prologue("env-1", "registration-1", "stream-1"); + let (initiator_handshake, request) = InitiatorHandshake::start( + &initiator, + &responder.public_key(), + &prologue, + b"authorization", + ) + .expect("start initiator handshake"); + let responder_handshake = + PendingResponderHandshake::read_request(&responder, &prologue, &request) + .expect("read responder handshake"); + let (mut responder_transport, response) = responder_handshake + .complete() + .expect("complete responder handshake"); + let mut initiator_transport = initiator_handshake + .finish(&response) + .expect("complete initiator handshake"); + let mut ciphertext = initiator_transport + .encrypt(b"request") + .expect("encrypt request"); + ciphertext[0] ^= 1; + + assert!(responder_transport.decrypt(&ciphertext).is_err()); +} + +#[test] +fn transport_rejects_replayed_ciphertext() { + let initiator = NoiseChannelIdentity::generate().expect("generate initiator identity"); + let responder = NoiseChannelIdentity::generate().expect("generate responder identity"); + let prologue = noise_channel_prologue("env-1", "registration-1", "stream-1"); + let (initiator_handshake, request) = InitiatorHandshake::start( + &initiator, + &responder.public_key(), + &prologue, + b"authorization", + ) + .expect("start initiator handshake"); + let responder_handshake = + PendingResponderHandshake::read_request(&responder, &prologue, &request) + .expect("read responder handshake"); + let (mut responder_transport, response) = responder_handshake + .complete() + .expect("complete responder handshake"); + let mut initiator_transport = initiator_handshake + .finish(&response) + .expect("complete initiator handshake"); + let ciphertext = initiator_transport + .encrypt(b"request") + .expect("encrypt request"); + + assert_eq!( + responder_transport + .decrypt(&ciphertext) + .expect("decrypt request"), + b"request" + ); + assert!(matches!( + responder_transport.decrypt(&ciphertext), + Err(NoiseChannelError::Transport(_)) + )); +} + +#[test] +fn public_key_validation_rejects_unknown_suite() { + let key = NoiseChannelIdentity::generate() + .expect("generate identity") + .public_key(); + let json = serde_json::to_value(key).expect("serialize key"); + let mut object = json.as_object().expect("key object").clone(); + object.insert("suite".to_string(), serde_json::json!("unknown")); + let key: NoiseChannelPublicKey = + serde_json::from_value(serde_json::Value::Object(object)).expect("deserialize key"); + + let initiator = NoiseChannelIdentity::generate().expect("generate initiator identity"); + assert!(InitiatorHandshake::start(&initiator, &key, b"prologue", b"").is_err()); +} + +#[test] +fn public_key_serializes_with_expected_suite() { + let key = NoiseChannelIdentity::generate() + .expect("generate identity") + .public_key(); + + let json = serde_json::to_value(key).expect("serialize key"); + + assert_eq!(json["suite"], NOISE_CHANNEL_SUITE); +} + +#[test] +fn initiator_rejects_oversized_handshake_payload() { + let initiator = NoiseChannelIdentity::generate().expect("generate initiator identity"); + let responder = NoiseChannelIdentity::generate().expect("generate responder identity"); + let payload = vec![0; MAX_MESSAGE_LEN]; + + let result = + InitiatorHandshake::start(&initiator, &responder.public_key(), b"prologue", &payload); + + assert!(matches!( + result, + Err(NoiseChannelError::InvalidMessage( + "handshake payload is too large" + )) + )); +} diff --git a/codex-rs/exec-server/src/noise_relay/executor_stream.rs b/codex-rs/exec-server/src/noise_relay/executor_stream.rs new file mode 100644 index 0000000000000000000000000000000000000000..612cc9781f1f607c58d72832073187275713d73c --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/executor_stream.rs @@ -0,0 +1,195 @@ +//! One executor-side virtual stream after the Noise handshake. +//! +//! The environment loop owns reads and a per-stream task owns writes. They share +//! `NoiseTransport` because its send and receive nonces live in the same value; +//! the mutex is never held across `.await`. + +use std::sync::Arc; +use std::sync::Mutex; + +use tokio::sync::mpsc; +use tokio::sync::watch; +use tracing::warn; + +use crate::ExecServerError; +use crate::connection::CHANNEL_CAPACITY; +use crate::noise_channel::NoiseTransport; +use crate::noise_relay::message_framing::MessageDecoder; +use crate::noise_relay::message_framing::NOISE_RECORD_PLAINTEXT_LEN; +use crate::noise_relay::ordered_ciphertext::OrderedCiphertextFrames; +use crate::noise_relay::stream_handler::NoiseStreamConnection; +use crate::noise_relay::stream_handler::NoiseStreamHandler; +use crate::noise_relay::take_next_sequence; +use crate::relay::encode_relay_message_frame; +use crate::relay_proto::RelayData; +use crate::relay_proto::RelayMessageFrame; +use crate::telemetry::ExecutorRegistration; + +/// Identifies one completed virtual-stream instance. +/// +/// Stream IDs are supplied by the untrusted relay peer and may be reused. The +/// instance ID prevents a delayed writer notification from removing a newer +/// stream that happens to use the same routing ID. +pub(crate) struct ClosedNoiseVirtualStream { + pub(crate) stream_id: String, + pub(crate) instance_id: u64, +} + +/// One authenticated application stream carried by the executor's physical relay. +/// +/// Inbound delivery is intentionally nonblocking. An overloaded or abandoned +/// stream fails independently instead of stalling every stream multiplexed over +/// the same physical websocket. +pub(crate) struct NoiseVirtualStream { + incoming_tx: mpsc::Sender, + disconnected_tx: watch::Sender, + transport: Arc>, + inbound_ciphertexts: OrderedCiphertextFrames, + inbound_decoder: MessageDecoder, + pub(crate) instance_id: u64, +} + +impl NoiseVirtualStream { + pub(crate) fn disconnect(self) { + let _ = self.disconnected_tx.send(true); + } + + /// Reorder and decrypt one record, then deliver complete payloads to the handler. + /// This must stay nonblocking because all virtual streams share the read loop. + pub(crate) fn receive_data(&mut self, data: RelayData) -> Result<(), ExecServerError> { + for ciphertext in self.inbound_ciphertexts.push(data.seq, data.payload)? { + let plaintext = { + let mut transport = self + .transport + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + transport.decrypt(&ciphertext).map_err(|error| { + ExecServerError::Protocol(format!("Noise relay decryption failed: {error}")) + })? + }; + for message in self.inbound_decoder.push(&plaintext)? { + self.incoming_tx + .try_send(H::decode(message)?) + .map_err(|_| { + ExecServerError::Protocol( + "Noise virtual stream inbound queue is full or closed".to_string(), + ) + })?; + } + } + Ok(()) + } +} + +/// Hand a completed handshake to its execution or forwarding owner. +/// +/// The returned value is the read half; the spawned task owns outbound framing +/// and reports its instance ID on exit so stream-ID reuse is safe. +pub(crate) fn spawn_noise_virtual_stream( + stream_id: String, + instance_id: u64, + handler: H, + physical_outgoing_tx: mpsc::Sender>, + closed_stream_tx: mpsc::Sender, + transport: NoiseTransport, + executor_registration: Option>, +) -> NoiseVirtualStream { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (incoming_tx, incoming_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (disconnected_tx, disconnected_rx) = watch::channel(false); + let transport = Arc::new(Mutex::new(transport)); + let writer_transport = Arc::clone(&transport); + let owner_stream_id = stream_id.clone(); + let owner_closed_stream_tx = closed_stream_tx.clone(); + let writer_stream_id = stream_id; + let writer_task = tokio::spawn(async move { + let mut next_seq = 0u32; + 'writer: while let Some(message) = outgoing_rx.recv().await { + let message = match H::encode(message) { + Ok(message) => message, + Err(error) => { + warn!("failed to encode Noise virtual stream payload: {error}"); + break; + } + }; + // Each chunk becomes one Noise record and consumes one nonce. + let mut trace = message.trace; + for plaintext_record in message.framed.chunks(NOISE_RECORD_PLAINTEXT_LEN) { + let seq = match take_next_sequence(&mut next_seq) { + Ok(seq) => seq, + Err(error) => { + warn!("Noise virtual stream sequence exhausted: {error}"); + break 'writer; + } + }; + let ciphertext = { + let mut transport = writer_transport + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + transport.encrypt(plaintext_record) + }; + let ciphertext = match ciphertext { + Ok(ciphertext) => ciphertext, + Err(error) => { + warn!("failed to encrypt Noise virtual stream payload: {error}"); + break 'writer; + } + }; + let frame = RelayMessageFrame::data( + writer_stream_id.clone(), + seq, + ciphertext, + trace.take(), + ); + if physical_outgoing_tx + .send(encode_relay_message_frame(&frame)) + .await + .is_err() + { + break 'writer; + } + } + } + + // The physical relay owns reset delivery and rejects stale instance IDs. + let closed_stream = ClosedNoiseVirtualStream { + stream_id: writer_stream_id, + instance_id, + }; + let _ = closed_stream_tx.send(closed_stream).await; + }); + + let connection = NoiseStreamConnection { + outgoing_tx, + incoming_rx, + disconnected_rx, + writer_task, + executor_registration, + }; + tokio::spawn(async move { + handler.run_connection(connection).await; + let _ = owner_closed_stream_tx + .send(ClosedNoiseVirtualStream { + stream_id: owner_stream_id, + instance_id, + }) + .await; + }); + + NoiseVirtualStream { + incoming_tx, + disconnected_tx, + transport, + inbound_ciphertexts: OrderedCiphertextFrames::default(), + inbound_decoder: MessageDecoder::default(), + instance_id, + } +} + +#[cfg(test)] +#[path = "executor_stream_tests.rs"] +mod tests; + +#[cfg(test)] +#[path = "forward_stream_tests.rs"] +mod forward_tests; diff --git a/codex-rs/exec-server/src/noise_relay/executor_stream_tests.rs b/codex-rs/exec-server/src/noise_relay/executor_stream_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4bf8f945dfef8e65e70c5db058683fa2b9bfbd0d --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/executor_stream_tests.rs @@ -0,0 +1,125 @@ +use std::time::Duration; + +use anyhow::Result; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::RequestId; +use codex_protocol::protocol::W3cTraceContext; +use tokio::sync::mpsc; +use tokio::time::timeout; + +use super::ClosedNoiseVirtualStream; +use super::spawn_noise_virtual_stream; +use crate::ExecServerRuntimePaths; +use crate::connection::CHANNEL_CAPACITY; +use crate::noise_channel::InitiatorHandshake; +use crate::noise_channel::NoiseChannelIdentity; +use crate::noise_channel::PendingResponderHandshake; +use crate::noise_relay::message_framing::frame_jsonrpc_message; +use crate::relay_proto::RelayData; +use crate::relay_proto::RelayMessageFrame; +use crate::server::ConnectionProcessor; + +#[test] +fn executor_requests_attach_trace_context_only_to_the_first_noise_record() { + let traceparent = "00-00000000000000000000000000000001-0000000000000002-01"; + let tracestate = "dd=s:1"; + let owned_traceparent = traceparent.to_string(); + let owned_tracestate = tracestate.to_string(); + let traceparent_ptr = owned_traceparent.as_ptr(); + let tracestate_ptr = owned_tracestate.as_ptr(); + let mut request = JSONRPCRequest { + id: RequestId::Integer(1), + method: "approval/request".to_string(), + params: None, + trace: Some(W3cTraceContext { + traceparent: Some(owned_traceparent), + tracestate: Some(owned_tracestate), + }), + }; + + let first = RelayMessageFrame::data( + "stream-1".to_string(), + /*seq*/ 0, + vec![1], + request.trace.take(), + ); + assert_eq!(first.traceparent.as_deref(), Some(traceparent)); + assert_eq!(first.tracestate.as_deref(), Some(tracestate)); + assert_eq!( + first.traceparent.as_ref().unwrap().as_ptr(), + traceparent_ptr + ); + assert_eq!(first.tracestate.as_ref().unwrap().as_ptr(), tracestate_ptr); + + let second = RelayMessageFrame::data( + "stream-1".to_string(), + /*seq*/ 1, + vec![2], + request.trace.take(), + ); + assert!(second.traceparent.is_none()); + assert!(second.tracestate.is_none()); + + let response = RelayMessageFrame::data( + "stream-1".to_string(), + /*seq*/ 2, + vec![3], + /*trace*/ None, + ); + assert!(response.traceparent.is_none()); + assert!(response.tracestate.is_none()); +} + +#[tokio::test] +async fn processor_exit_reports_closed_virtual_stream() -> Result<()> { + let executor_identity = NoiseChannelIdentity::generate()?; + let harness_identity = NoiseChannelIdentity::generate()?; + let prologue = b"test-prologue"; + let (initiator, request) = InitiatorHandshake::start( + &harness_identity, + &executor_identity.public_key(), + prologue, + b"authorization", + )?; + let pending = PendingResponderHandshake::read_request(&executor_identity, prologue, &request)?; + let (executor_transport, response) = pending.complete()?; + let mut harness_transport = initiator.finish(&response)?; + + let (physical_outgoing_tx, _physical_outgoing_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (closed_stream_tx, mut closed_stream_rx) = mpsc::channel(1); + let mut stream = spawn_noise_virtual_stream( + "stream-1".to_string(), + /*instance_id*/ 7, + ConnectionProcessor::new(ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?), + physical_outgoing_tx, + closed_stream_tx, + executor_transport, + /*executor_registration*/ None, + ); + + let message = JSONRPCMessage::Response(JSONRPCResponse { + id: RequestId::Integer(1), + result: serde_json::Value::Null, + }); + let ciphertext = harness_transport.encrypt(&frame_jsonrpc_message(&message)?)?; + stream.receive_data(RelayData { + seq: 0, + segment_index: 0, + segment_count: 1, + payload: ciphertext, + })?; + + assert!(matches!( + timeout(Duration::from_secs(1), closed_stream_rx.recv()).await?, + Some(ClosedNoiseVirtualStream { + stream_id, + instance_id: 7, + }) if stream_id == "stream-1" + )); + Ok(()) +} diff --git a/codex-rs/exec-server/src/noise_relay/forward_stream_tests.rs b/codex-rs/exec-server/src/noise_relay/forward_stream_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..d19c3232a5c3c0ea7912d644fdcbb6c554bac64b --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/forward_stream_tests.rs @@ -0,0 +1,112 @@ +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use bytes::Bytes; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use futures::SinkExt; +use futures::StreamExt; +use pretty_assertions::assert_eq; +use tokio::net::TcpListener; +use tokio::sync::mpsc; +use tokio::time::timeout; +use tokio_tungstenite::accept_async; +use tokio_tungstenite::tungstenite::Message; + +use super::ClosedNoiseVirtualStream; +use super::spawn_noise_virtual_stream; +use crate::ExecServerTelemetry; +use crate::forward::Forwarder; +use crate::noise_channel::InitiatorHandshake; +use crate::noise_channel::NoiseChannelIdentity; +use crate::noise_channel::PendingResponderHandshake; +use crate::noise_relay::message_framing::MessageDecoder; +use crate::noise_relay::message_framing::frame_message; +use crate::relay::decode_relay_message_frame; +use crate::relay_proto::RelayData; + +const TEST_TIMEOUT: Duration = Duration::from_secs(5); + +#[tokio::test] +async fn forwards_opaque_noise_payloads_and_drains_before_closing() -> Result<()> { + let executor_identity = NoiseChannelIdentity::generate()?; + let harness_identity = NoiseChannelIdentity::generate()?; + let prologue = b"forwarding-test"; + let (initiator, request) = InitiatorHandshake::start( + &harness_identity, + &executor_identity.public_key(), + prologue, + b"authorization", + )?; + let pending = PendingResponderHandshake::read_request(&executor_identity, prologue, &request)?; + let (executor_transport, response) = pending.complete()?; + let mut harness_transport = initiator.finish(&response)?; + + let destination = TcpListener::bind("127.0.0.1:0").await?; + let forwarder = Forwarder::new( + format!("ws://{}", destination.local_addr()?), + &HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ExecServerTelemetry::default(), + )?; + // A single queued record makes the final response exercise writer draining. + let (physical_outgoing_tx, mut physical_outgoing_rx) = mpsc::channel(1); + let (closed_stream_tx, mut closed_stream_rx) = mpsc::channel(2); + let mut stream = spawn_noise_virtual_stream( + "stream-1".to_string(), + /*instance_id*/ 7, + forwarder, + physical_outgoing_tx, + closed_stream_tx, + executor_transport, + /*executor_registration*/ None, + ); + let (socket, _) = timeout(TEST_TIMEOUT, destination.accept()).await??; + let mut socket = timeout(TEST_TIMEOUT, accept_async(socket)).await??; + + let request = Bytes::from_static(b"not JSON\x00\xff"); + stream.receive_data(RelayData { + seq: 0, + segment_index: 0, + segment_count: 1, + payload: harness_transport.encrypt(&frame_message(&request)?)?, + })?; + let forwarded = timeout(TEST_TIMEOUT, socket.next()) + .await? + .context("forwarded request")??; + assert_eq!(forwarded.into_data(), request); + + let response = Bytes::from(vec![0xff; 128 * 1024]); + socket.send(Message::Binary(response.clone())).await?; + socket.close(None).await?; + + let mut decoder = MessageDecoder::default(); + let mut messages = Vec::new(); + let mut next_seq = 0; + let closed = timeout(TEST_TIMEOUT, async { + loop { + tokio::select! { + biased; + Some(encoded) = physical_outgoing_rx.recv() => { + let frame = decode_relay_message_frame(&encoded)?; + assert_eq!(frame.stream_id, "stream-1"); + let data = frame.into_data()?; + assert_eq!(data.seq, next_seq); + next_seq += 1; + messages.extend(decoder.push(&harness_transport.decrypt(&data.payload)?)?); + } + closed = closed_stream_rx.recv() => { + break closed.context("stream close notification"); + } + } + } + }) + .await??; + assert!(next_seq > 1); + assert_eq!(messages, vec![response]); + assert!(matches!( + closed, + ClosedNoiseVirtualStream { stream_id, instance_id: 7 } if stream_id == "stream-1" + )); + Ok(()) +} diff --git a/codex-rs/exec-server/src/noise_relay/harness.rs b/codex-rs/exec-server/src/noise_relay/harness.rs new file mode 100644 index 0000000000000000000000000000000000000000..793f7c1e12e9816e58606b3ffeaf9a1b009f5b2a --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/harness.rs @@ -0,0 +1,659 @@ +//! Harness side of the Noise relay. +//! +//! The rendezvous service routes frames by `stream_id`, but does not authenticate +//! the executor or see JSON-RPC plaintext. We claim a stream, complete hybrid IK +//! against the registry-provided executor key, and then expose the result as a +//! normal `JsonRpcConnection`. Outbound JSON-RPC is framed and split into Noise +//! records; inbound records are reordered before decryption and reassembly. + +use futures::FutureExt; +use futures::Sink; +use futures::SinkExt; +use futures::Stream; +use futures::StreamExt; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::sync::watch; +use tokio_tungstenite::tungstenite::Message; +use tracing::Instrument; +use tracing::debug; +use tracing::info; +use tracing::warn; +use uuid::Uuid; + +use crate::ExecServerError; +use crate::connection::CHANNEL_CAPACITY; +use crate::connection::JsonRpcConnection; +use crate::connection::JsonRpcConnectionEvent; +use crate::connection::JsonRpcTransport; +use crate::connection::WEBSOCKET_KEEPALIVE_INTERVAL; +use crate::noise_channel::InitiatorHandshake; +use crate::noise_channel::NoiseChannelIdentity; +use crate::noise_channel::NoiseChannelPublicKey; +use crate::noise_channel::NoiseTransport; +use crate::noise_channel::noise_channel_prologue; +use crate::noise_relay::message_framing::JsonRpcMessageDecoder; +use crate::noise_relay::message_framing::NOISE_RECORD_PLAINTEXT_LEN; +use crate::noise_relay::message_framing::frame_jsonrpc_message; +use crate::noise_relay::ordered_ciphertext::OrderedCiphertextFrames; +use crate::noise_relay::take_next_sequence; +use crate::relay::RelayFrameBodyKind; +use crate::relay::decode_relay_message_frame; +use crate::relay::encode_relay_message_frame; +use crate::relay_proto::RelayData; +use crate::relay_proto::RelayMessageFrame; +use crate::websocket_pong_watchdog::WEBSOCKET_PONG_TIMEOUT; +use crate::websocket_pong_watchdog::WEBSOCKET_PONG_TIMEOUT_REASON; +use crate::websocket_pong_watchdog::WebSocketPongWatchdog; + +/// Values that bind one harness websocket to the intended executor registration. +/// +/// These fields all come from the same registry response. Keeping them together +/// makes that relationship visible at the call site and avoids mixing up the +/// several string and key arguments used to start the handshake. +pub(crate) struct NoiseHarnessConnectionArgs { + pub(crate) connection_label: String, + pub(crate) environment_id: String, + pub(crate) executor_registration_id: String, + pub(crate) identity: NoiseChannelIdentity, + pub(crate) responder_public_key: NoiseChannelPublicKey, + pub(crate) harness_key_authorization: String, +} + +/// One Noise-backed JSON-RPC connection plus a signal that its authenticated +/// transport is ready for application messages. +pub(crate) struct NoiseHarnessConnection { + pub(crate) connection: JsonRpcConnection, + pub(crate) handshake_ready: oneshot::Receiver<()>, +} + +// Reset frames are cleartext relay control and are not authenticated by Noise. +// Preserve the availability signal while replacing attacker-controlled reason +// text before it reaches disconnect diagnostics. +const NOISE_RELAY_RESET_DISCONNECT_REASON: &str = "Noise relay stream reset"; +// Give a Pong already queued behind data a bounded chance to reach the reader. +const MAX_FRAMES_DRAINED_AFTER_PONG_DEADLINE: usize = 32; + +/// Adapt one harness rendezvous websocket and expose when hybrid IK completes. +/// +/// Callers that send application messages immediately after opening the +/// websocket must await handshake_ready first. Dropping that receiver keeps +/// the legacy fire-and-forget behavior for existing connection owners. +pub(crate) fn noise_harness_connection_from_websocket_with_readiness( + stream: T, + args: NoiseHarnessConnectionArgs, +) -> NoiseHarnessConnection +where + T: Sink + Stream> + Unpin + Send + 'static, + E: std::fmt::Display + Send + 'static, +{ + let NoiseHarnessConnectionArgs { + connection_label, + environment_id, + executor_registration_id, + identity, + responder_public_key, + harness_key_authorization, + } = args; + let stream_id = Uuid::new_v4().to_string(); + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (incoming_tx, incoming_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (disconnected_tx, disconnected_rx) = watch::channel(false); + let (handshake_ready_tx, handshake_ready_rx) = oneshot::channel(); + let stream_span = tracing::debug_span!("noise_relay.stream", noise_side = "harness",); + debug!( + environment_id, + executor_registration_id, stream_id, "Noise harness relay details" + ); + + let websocket_task = tokio::spawn(async move { + let mut handshake_ready_tx = Some(handshake_ready_tx); + let mut websocket = stream; + + // Bind the Noise transcript to the exact environment registration and + // virtual relay stream before emitting any handshake bytes. A captured + // handshake cannot be spliced onto a different routed connection. + let prologue = + noise_channel_prologue(&environment_id, &executor_registration_id, &stream_id); + let (initiator_handshake, request) = match InitiatorHandshake::start( + &identity, + &responder_public_key, + &prologue, + harness_key_authorization.as_bytes(), + ) { + Ok(handshake) => handshake, + Err(error) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + format!("failed to start Noise relay handshake: {error}"), + ); + return; + } + }; + + // Resume claims the stream ID at rendezvous; Handshake carries the + // opaque first IK message. No JSON-RPC data is sent before the + // responder proves possession of the pinned static key. + let resume = RelayMessageFrame::resume(stream_id.clone()); + let handshake = RelayMessageFrame::handshake(stream_id.clone(), request); + if websocket + .send(Message::Binary(encode_relay_message_frame(&resume).into())) + .await + .is_err() + || websocket + .send(Message::Binary( + encode_relay_message_frame(&handshake).into(), + )) + .await + .is_err() + { + let _ = disconnected_tx.send(true); + return; + } + + // During the handshake, ignore unrelated routed streams and control + // frames, but reject data on our stream. Accepting early data would + // create a plaintext or unauthenticated application path. + let mut transport = loop { + let Some(incoming_message) = websocket.next().await else { + send_disconnected( + &incoming_tx, + &disconnected_tx, + "Noise relay websocket ended during handshake".to_string(), + ); + return; + }; + let message = match incoming_message { + Ok(Message::Binary(payload)) => payload, + Ok(Message::Close(_)) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + "Noise relay websocket received close frame during handshake".to_string(), + ); + return; + } + Ok(Message::Ping(_) | Message::Pong(_) | Message::Frame(_)) => continue, + Ok(Message::Text(_)) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + "Noise relay transport expects binary protobuf frames".to_string(), + ); + return; + } + Err(error) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + format!( + "failed to read Noise relay websocket from {connection_label}: {error}" + ), + ); + return; + } + }; + let frame = match decode_relay_message_frame(message.as_ref()) { + Ok(frame) => frame, + Err(error) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + format!("failed to parse Noise relay frame: {error}"), + ); + return; + } + }; + if frame.stream_id != stream_id { + debug!("Noise relay ignored frame for unrelated stream during handshake"); + continue; + } + match frame.validate() { + Ok(RelayFrameBodyKind::Handshake) => { + let response = match frame.into_handshake_payload() { + Ok(response) => response, + Err(error) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + format!("invalid Noise relay handshake response: {error}"), + ); + return; + } + }; + match initiator_handshake.finish(&response) { + Ok(transport) => { + info!( + noise_event = "handshake", + noise_outcome = "ok", + "Noise harness handshake completed" + ); + if let Some(handshake_ready_tx) = handshake_ready_tx.take() { + let _ = handshake_ready_tx.send(()); + } + break transport; + } + Err(error) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + format!("Noise relay handshake failed: {error}"), + ); + return; + } + } + } + Ok(RelayFrameBodyKind::Reset) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + NOISE_RELAY_RESET_DISCONNECT_REASON.to_string(), + ); + return; + } + Ok( + RelayFrameBodyKind::Ack + | RelayFrameBodyKind::Resume + | RelayFrameBodyKind::Heartbeat, + ) => {} + Ok(RelayFrameBodyKind::Data) | Err(_) => { + send_disconnected( + &incoming_tx, + &disconnected_tx, + "Noise relay received data before handshake completion".to_string(), + ); + return; + } + } + }; + + // After the handshake, each relay sequence maps to exactly one Noise + // transport record. Outbound records are encrypted once; inbound + // records are reordered and deduplicated before decryption. + let mut websocket = websocket.peekable(); + let mut next_outbound_seq = 0u32; + let mut inbound_ciphertexts = OrderedCiphertextFrames::default(); + let mut inbound_decoder = JsonRpcMessageDecoder::default(); + let mut keepalive = tokio::time::interval_at( + tokio::time::Instant::now() + WEBSOCKET_KEEPALIVE_INTERVAL, + WEBSOCKET_KEEPALIVE_INTERVAL, + ); + keepalive.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + let mut pong_watchdog = WebSocketPongWatchdog::new(WEBSOCKET_PONG_TIMEOUT); + let pong_deadline = tokio::time::sleep(WEBSOCKET_PONG_TIMEOUT); + tokio::pin!(pong_deadline); + // Keep one framed message as a cursor. Sending one Noise record per loop + // creates a scheduling point for keepalive and inbound control frames + // without splitting the WebSocket reader and writer. + let mut pending_outbound = None; + let mut force_incoming = false; + let mut frames_drained_after_pong_deadline = 0usize; + 'relay: loop { + // Consume a due tick before the always-ready record arm below can win + // another select iteration and postpone the keepalive. + if pong_watchdog.deadline().is_none() + && keepalive.tick().now_or_never().is_some() + { + if let Err(error) = send_keepalive_ping( + &mut websocket, + &mut pong_watchdog, + pong_deadline.as_mut(), + ) + .await + { + warn!("failed to write Noise relay keepalive ping: {error}"); + break; + } + frames_drained_after_pong_deadline = 0; + continue; + } + + let pong_deadline_expired = pong_watchdog + .deadline() + .is_some_and(|deadline| tokio::time::Instant::now() >= deadline); + // After expiry, inspect only frames already queued. Forcing the peeked + // item through next() makes the 32-frame grace deterministic. + if pong_deadline_expired && !force_incoming { + if frames_drained_after_pong_deadline + < MAX_FRAMES_DRAINED_AFTER_PONG_DEADLINE + && std::pin::Pin::new(&mut websocket) + .peek() + .now_or_never() + .is_some() + { + force_incoming = true; + } else { + warn!( + noise_reason = WEBSOCKET_PONG_TIMEOUT_REASON, + "Noise harness rendezvous websocket disconnected" + ); + send_disconnected( + &incoming_tx, + &disconnected_tx, + WEBSOCKET_PONG_TIMEOUT_REASON.to_string(), + ); + return; + } + } + + // While a Pong is outstanding, drain already-queued inbound traffic + // before the next fragment so a queued Pong cannot sit behind writes. + if !force_incoming + && pong_watchdog.deadline().is_some() + && pending_outbound.is_some() + && std::pin::Pin::new(&mut websocket) + .peek() + .now_or_never() + .is_some() + { + force_incoming = true; + } + + tokio::select! { + maybe_message = outgoing_rx.recv(), if pending_outbound.is_none() && !force_incoming && !pong_deadline_expired => { + let Some(message) = maybe_message else { + break; + }; + let framed = match frame_jsonrpc_message(&message) { + Ok(framed) => framed, + Err(error) => { + warn!("failed to frame JSON-RPC payload for Noise relay: {error}"); + break; + } + }; + let request_trace = match message { + codex_exec_server_protocol::JSONRPCMessage::Request(request) => request.trace, + codex_exec_server_protocol::JSONRPCMessage::Notification(_) + | codex_exec_server_protocol::JSONRPCMessage::Response(_) + | codex_exec_server_protocol::JSONRPCMessage::Error(_) => None, + }; + pending_outbound = Some((framed, 0, request_trace)); + } + _ = std::future::ready(()), if pending_outbound.is_some() && !force_incoming && !pong_deadline_expired => { + let seq = match take_next_sequence(&mut next_outbound_seq) { + Ok(seq) => seq, + Err(error) => { + warn!("Noise relay sequence exhausted: {error}"); + break 'relay; + } + }; + let (ciphertext, next_offset, message_complete, request_trace) = { + let Some((framed, offset, request_trace)) = pending_outbound.as_mut() else { + continue; + }; + let next_offset = (*offset + NOISE_RECORD_PLAINTEXT_LEN).min(framed.len()); + let ciphertext = match transport.encrypt(&framed[*offset..next_offset]) { + Ok(ciphertext) => ciphertext, + Err(error) => { + warn!("failed to encrypt JSON-RPC payload for Noise relay: {error}"); + break 'relay; + } + }; + ( + ciphertext, + next_offset, + next_offset == framed.len(), + request_trace.take(), + ) + }; + let frame = + RelayMessageFrame::data(stream_id.clone(), seq, ciphertext, request_trace); + // A Pong can arrive after the readiness check while this write owns the + // combined sink and stream. A single bounded record can therefore hit the + // deadline and disconnect with that Pong queued. Treat that as write + // backpressure; this loop yields only between records. + if let Err(error) = send_websocket_message( + &mut websocket, + Message::Binary(encode_relay_message_frame(&frame).into()), + pong_watchdog.write_deadline(tokio::time::Instant::now()), + ) + .await + { + warn!("failed to write Noise relay websocket: {error}"); + break 'relay; + } + if message_complete { + pending_outbound = None; + } else if let Some((_framed, offset, _request_trace)) = pending_outbound.as_mut() { + *offset = next_offset; + } + } + _ = &mut pong_deadline, if pong_watchdog.deadline().is_some() && !force_incoming => { + continue; + } + _ = keepalive.tick(), if pong_watchdog.deadline().is_none() => { + if let Err(error) = send_keepalive_ping( + &mut websocket, + &mut pong_watchdog, + pong_deadline.as_mut(), + ) + .await + { + warn!("failed to write Noise relay keepalive ping: {error}"); + break; + } + frames_drained_after_pong_deadline = 0; + } + incoming_message = websocket.next() => { + force_incoming = false; + let Some(incoming_message) = incoming_message else { + break; + }; + // Count each completed read after expiry. If only the deadline arm + // advanced this counter, reads won by a simultaneously ready incoming + // arm would not count toward the 32-frame cap. + if pong_watchdog + .deadline() + .is_some_and(|deadline| tokio::time::Instant::now() >= deadline) + { + frames_drained_after_pong_deadline += 1; + } + match incoming_message { + Ok(Message::Binary(payload)) => { + let frame = match decode_relay_message_frame(payload.as_ref()) { + Ok(frame) => frame, + Err(error) => { + send_malformed(&incoming_tx, error.to_string()); + break; + } + }; + if frame.stream_id != stream_id { + continue; + } + match frame.validate() { + Ok(RelayFrameBodyKind::Data) => { + let data = match frame.into_data() { + Ok(data) => data, + Err(error) => { + send_malformed(&incoming_tx, error.to_string()); + break; + } + }; + if let Err(error) = receive_data( + &mut inbound_ciphertexts, + &mut transport, + &mut inbound_decoder, + data, + pong_watchdog.write_deadline(tokio::time::Instant::now()), + &incoming_tx, + ) + .await + { + if matches!(error, ExecServerError::Closed) { + break; + } + send_malformed(&incoming_tx, error.to_string()); + break; + } + } + Ok(RelayFrameBodyKind::Reset) => { + let _ = incoming_tx.try_send( + JsonRpcConnectionEvent::Disconnected { + reason: Some( + NOISE_RELAY_RESET_DISCONNECT_REASON.to_string(), + ), + }, + ); + break; + } + Ok( + RelayFrameBodyKind::Ack + | RelayFrameBodyKind::Resume + | RelayFrameBodyKind::Heartbeat, + ) => {} + Ok(RelayFrameBodyKind::Handshake) | Err(_) => { + send_malformed( + &incoming_tx, + "Noise relay received invalid post-handshake frame".to_string(), + ); + break; + } + } + } + Ok(Message::Close(_)) => break, + Ok(Message::Pong(_)) => { + pong_watchdog.received_pong(); + frames_drained_after_pong_deadline = 0; + } + Ok(Message::Ping(_) | Message::Frame(_)) => {} + Ok(Message::Text(_)) => { + send_malformed( + &incoming_tx, + "Noise relay transport expects binary protobuf frames".to_string(), + ); + break; + } + Err(error) => { + debug!("Noise relay websocket read failed: {error}"); + break; + } + } + } + } + } + let _ = disconnected_tx.send(true); + } + .instrument(stream_span)); + + NoiseHarnessConnection { + connection: JsonRpcConnection { + outgoing_tx, + incoming_rx, + disconnected_rx, + task_handles: vec![websocket_task], + transport: JsonRpcTransport::Plain, + }, + handshake_ready: handshake_ready_rx, + } +} + +async fn send_websocket_message( + websocket: &mut T, + message: Message, + deadline: tokio::time::Instant, +) -> Result<(), String> +where + T: Sink + Unpin, + E: std::fmt::Display, +{ + match tokio::time::timeout_at(deadline, websocket.send(message)).await { + Ok(Ok(())) => Ok(()), + Ok(Err(error)) => Err(error.to_string()), + Err(_) => Err("websocket write timed out".to_string()), + } +} + +async fn send_keepalive_ping( + websocket: &mut T, + pong_watchdog: &mut WebSocketPongWatchdog, + pong_deadline: std::pin::Pin<&mut tokio::time::Sleep>, +) -> Result<(), String> +where + T: Sink + Unpin, + E: std::fmt::Display, +{ + send_websocket_message( + websocket, + Message::Ping(Vec::new().into()), + pong_watchdog.write_deadline(tokio::time::Instant::now()), + ) + .await?; + // Start the response clock after the Ping flushes; waiting for sink capacity + // is governed by the write deadline above. + pong_watchdog.ping_sent(tokio::time::Instant::now()); + if let Some(deadline) = pong_watchdog.deadline() { + pong_deadline.reset(deadline); + } + Ok(()) +} + +/// Order and decrypt one relay frame, then emit any complete JSON-RPC messages. +/// Relay records and JSON-RPC messages do not share boundaries, so reassembly +/// happens after decryption. +async fn receive_data( + inbound_ciphertexts: &mut OrderedCiphertextFrames, + transport: &mut NoiseTransport, + decoder: &mut JsonRpcMessageDecoder, + data: RelayData, + delivery_deadline: tokio::time::Instant, + incoming_tx: &mpsc::Sender, +) -> Result<(), ExecServerError> { + // Ordering must happen before decryption because Noise transport nonces are + // implicit. A future or duplicate ciphertext passed directly to Clatter + // would desynchronize the channel. + for ciphertext in inbound_ciphertexts.push(data.seq, data.payload)? { + let plaintext = transport.decrypt(&ciphertext).map_err(|error| { + ExecServerError::Protocol(format!("Noise relay decryption failed: {error}")) + })?; + + // The authenticated byte stream can carry partial or multiple JSON-RPC + // messages; emit only complete, successfully parsed messages. + for message in decoder.push(&plaintext)? { + send_incoming_event( + incoming_tx, + JsonRpcConnectionEvent::message(message), + delivery_deadline, + ) + .await?; + } + } + Ok(()) +} + +async fn send_incoming_event( + incoming_tx: &mpsc::Sender, + event: JsonRpcConnectionEvent, + deadline: tokio::time::Instant, +) -> Result<(), ExecServerError> { + match tokio::time::timeout_at(deadline, incoming_tx.send(event)).await { + Ok(Ok(())) => Ok(()), + Ok(Err(_)) => Err(ExecServerError::Closed), + Err(_) => { + warn!( + noise_reason = "application_backpressure", + "Noise harness application event delivery timed out" + ); + Err(ExecServerError::Closed) + } + } +} + +fn send_malformed(incoming_tx: &mpsc::Sender, reason: String) { + let _ = incoming_tx.try_send(JsonRpcConnectionEvent::MalformedMessage { reason }); +} + +fn send_disconnected( + incoming_tx: &mpsc::Sender, + disconnected_tx: &watch::Sender, + reason: String, +) { + let _ = disconnected_tx.send(true); + let _ = incoming_tx.try_send(JsonRpcConnectionEvent::Disconnected { + reason: Some(reason), + }); +} + +#[cfg(test)] +#[path = "harness_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/noise_relay/harness_tests.rs b/codex-rs/exec-server/src/noise_relay/harness_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..5530941ca0eb7869dfa4e0e29a4535e5cf69cd30 --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/harness_tests.rs @@ -0,0 +1,509 @@ +use std::pin::Pin; +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::task::Context as TaskContext; +use std::task::Poll; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::RequestId; +use codex_protocol::protocol::W3cTraceContext; +use futures::Sink; +use futures::SinkExt; +use futures::StreamExt; +use futures::channel::mpsc as futures_mpsc; +use pretty_assertions::assert_eq; +use tokio::net::TcpListener; +use tokio::sync::mpsc; +use tokio::time::Instant; +use tokio::time::timeout; +use tokio_tungstenite::accept_async; +use tokio_tungstenite::connect_async; +use tokio_tungstenite::tungstenite::Message; + +use super::*; +use crate::connection::JsonRpcConnectionEvent; +use crate::noise_channel::PendingResponderHandshake; + +fn noise_harness_connection_from_websocket( + stream: T, + args: NoiseHarnessConnectionArgs, +) -> JsonRpcConnection +where + T: Sink + Stream> + Unpin + Send + 'static, + E: std::fmt::Display + Send + 'static, +{ + noise_harness_connection_from_websocket_with_readiness(stream, args).connection +} + +const ENVIRONMENT_ID: &str = "environment-1"; +const EXECUTOR_REGISTRATION_ID: &str = "registration-1"; + +#[tokio::test(start_paused = true)] +async fn first_encrypted_request_frame_exposes_only_its_trace_context() -> Result<()> { + let (connection, mut control, mut outbound_rx) = connected_controlled_harness().await?; + let traceparent = "00-00000000000000000000000000000001-0000000000000002-01"; + let tracestate = "dd=s:1"; + + connection + .outgoing_tx + .send(JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(1), + method: "fs/getMetadata".to_string(), + params: Some(serde_json::json!({ + "payload": "x".repeat(NOISE_RECORD_PLAINTEXT_LEN * 2), + })), + trace: Some(W3cTraceContext { + traceparent: Some(traceparent.to_string()), + tracestate: Some(tracestate.to_string()), + }), + })) + .await?; + + control.wait_for_blocked_write(/*expected*/ 1).await?; + control.grant_writes(/*count*/ 1); + let first = read_outbound_frame(&mut outbound_rx).await?; + assert_eq!( + (first.traceparent.as_deref(), first.tracestate.as_deref()), + (Some(traceparent), Some(tracestate)) + ); + let first_payload = first.into_data()?.payload; + assert!( + !first_payload + .windows("fs/getMetadata".len()) + .any(|window| window == b"fs/getMetadata") + ); + + control.wait_for_blocked_write(/*expected*/ 2).await?; + control.grant_writes(/*count*/ 1); + let second = read_outbound_frame(&mut outbound_rx).await?; + assert_eq!((second.traceparent, second.tracestate), (None, None)); + + for task in &connection.task_handles { + task.abort(); + } + Ok(()) +} + +#[tokio::test(start_paused = true)] +async fn fragmented_writes_yield_to_keepalive_and_queued_pong() -> Result<()> { + let (connection, mut control, mut outbound_rx) = connected_controlled_harness().await?; + + connection + .outgoing_tx + .send(JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(1), + method: "large".to_string(), + params: Some(serde_json::json!({ + "payload": "x".repeat(NOISE_RECORD_PLAINTEXT_LEN * 3), + })), + trace: None, + })) + .await?; + + control.wait_for_blocked_write(/*expected*/ 1).await?; + tokio::time::advance(WEBSOCKET_KEEPALIVE_INTERVAL + Duration::from_millis(10)).await; + control.grant_writes(/*count*/ 1); + let first_data = read_outbound_data(&mut outbound_rx).await?; + assert_eq!(first_data.seq, 0); + + control.wait_for_blocked_write(/*expected*/ 2).await?; + control.grant_writes(/*count*/ 1); + let Message::Ping(ping_payload) = timeout(Duration::from_secs(1), outbound_rx.next()) + .await? + .context("harness closed before sending keepalive")? + else { + anyhow::bail!("expected keepalive between fragmented writes"); + }; + + control.wait_for_blocked_write(/*expected*/ 3).await?; + control.send_inbound(Message::Pong(ping_payload))?; + tokio::time::advance(WEBSOCKET_KEEPALIVE_INTERVAL + Duration::from_millis(10)).await; + control.grant_writes(/*count*/ 1); + let second_data = read_outbound_data(&mut outbound_rx).await?; + assert_eq!(second_data.seq, 1); + + control.wait_for_blocked_write(/*expected*/ 4).await?; + control.grant_writes(/*count*/ 1); + let next_message = timeout(Duration::from_secs(1), outbound_rx.next()) + .await? + .context("harness closed after receiving queued Pong")?; + assert!(matches!(next_message, Message::Ping(_))); + + for task in &connection.task_handles { + task.abort(); + } + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn post_deadline_drain_stops_before_frame_33() -> Result<()> { + let (mut connection, mut control, mut outbound_rx) = connected_controlled_harness().await?; + + control.wait_for_blocked_write(/*expected*/ 1).await?; + control.grant_writes(/*count*/ 1); + let Message::Ping(ping_payload) = timeout(Duration::from_secs(1), outbound_rx.next()) + .await? + .context("harness closed before sending keepalive")? + else { + anyhow::bail!("expected keepalive ping"); + }; + let reads_before_deadline = control.inbound_reads(); + + let unrelated_frame = + encode_relay_message_frame(&RelayMessageFrame::resume("unrelated-stream".to_string())); + for _ in 0..MAX_FRAMES_DRAINED_AFTER_PONG_DEADLINE { + control.send_inbound(Message::Binary(unrelated_frame.clone().into()))?; + } + control.send_inbound(Message::Pong(ping_payload))?; + + // Keep the current-thread runtime from consuming the queued frames until the + // Pong deadline and every frame are ready together. + std::thread::sleep(WEBSOCKET_PONG_TIMEOUT + Duration::from_millis(10)); + + let event = timeout(Duration::from_secs(1), connection.incoming_rx.recv()).await?; + let Some(JsonRpcConnectionEvent::Disconnected { reason }) = event else { + anyhow::bail!("expected Pong timeout, got {event:?}"); + }; + assert_eq!(reason.as_deref(), Some(WEBSOCKET_PONG_TIMEOUT_REASON)); + assert_eq!( + control.inbound_reads() - reads_before_deadline, + MAX_FRAMES_DRAINED_AFTER_PONG_DEADLINE + ); + Ok(()) +} + +#[tokio::test] +async fn pong_keeps_harness_alive_until_peer_stops_responding() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let harness_connection = tokio::spawn(connect_async(websocket_url)); + let (socket, _peer_addr) = listener.accept().await?; + let mut executor_websocket = accept_async(socket).await?; + let (harness_websocket, _response) = harness_connection.await??; + + let executor_identity = NoiseChannelIdentity::generate()?; + let mut connection = noise_harness_connection_from_websocket( + harness_websocket, + NoiseHarnessConnectionArgs { + connection_label: "test rendezvous".to_string(), + environment_id: ENVIRONMENT_ID.to_string(), + executor_registration_id: EXECUTOR_REGISTRATION_ID.to_string(), + identity: NoiseChannelIdentity::generate()?, + responder_public_key: executor_identity.public_key(), + harness_key_authorization: "authorization".to_string(), + }, + ); + + let resume_message = timeout(Duration::from_secs(1), executor_websocket.next()) + .await? + .context("harness closed before sending resume")??; + let Message::Binary(resume_payload) = resume_message else { + anyhow::bail!("expected resume frame, got {resume_message:?}"); + }; + let resume = decode_relay_message_frame(resume_payload.as_ref())?; + assert_eq!(resume.validate()?, RelayFrameBodyKind::Resume); + + let handshake_message = timeout(Duration::from_secs(1), executor_websocket.next()) + .await? + .context("harness closed before sending handshake")??; + let Message::Binary(handshake_payload) = handshake_message else { + anyhow::bail!("expected handshake frame, got {handshake_message:?}"); + }; + let handshake = decode_relay_message_frame(handshake_payload.as_ref())?; + assert_eq!(handshake.stream_id, resume.stream_id); + let stream_id = handshake.stream_id.clone(); + let prologue = + noise_channel_prologue(ENVIRONMENT_ID, EXECUTOR_REGISTRATION_ID, stream_id.as_str()); + let pending = PendingResponderHandshake::read_request( + &executor_identity, + &prologue, + &handshake.into_handshake_payload()?, + )?; + let (_transport, response) = pending.complete()?; + let response = RelayMessageFrame::handshake(stream_id, response); + executor_websocket + .send(Message::Binary( + encode_relay_message_frame(&response).into(), + )) + .await?; + + let mut pings = 0; + while pings < 6 { + let message = timeout(Duration::from_secs(1), executor_websocket.next()) + .await? + .context("harness disconnected before six keepalive pings")??; + match message { + Message::Ping(payload) => { + executor_websocket.send(Message::Pong(payload)).await?; + pings += 1; + } + Message::Pong(_) | Message::Frame(_) => {} + message => anyhow::bail!("expected keepalive ping, got {message:?}"), + } + } + + // Keep non-Pong traffic flowing after responses stop. It must not defeat + // the bounded grace for a Pong already queued behind data. + let unrelated_frame = + encode_relay_message_frame(&RelayMessageFrame::resume("unrelated-stream".to_string())); + let traffic_task = tokio::spawn(async move { + loop { + if executor_websocket + .send(Message::Binary(unrelated_frame.clone().into())) + .await + .is_err() + { + break; + } + tokio::time::sleep(Duration::from_millis(5)).await; + } + }); + let event = timeout(Duration::from_secs(1), connection.incoming_rx.recv()).await?; + traffic_task.abort(); + let _ = traffic_task.await; + let Some(JsonRpcConnectionEvent::Disconnected { reason }) = event else { + anyhow::bail!("expected pong timeout, got {event:?}"); + }; + assert_eq!(reason.as_deref(), Some(WEBSOCKET_PONG_TIMEOUT_REASON)); + Ok(()) +} + +#[tokio::test] +async fn application_event_delivery_is_bounded() -> Result<()> { + let (incoming_tx, _incoming_rx) = mpsc::channel(1); + incoming_tx + .send(JsonRpcConnectionEvent::MalformedMessage { + reason: "fill queue".to_string(), + }) + .await?; + + let result = send_incoming_event( + &incoming_tx, + JsonRpcConnectionEvent::MalformedMessage { + reason: "blocked event".to_string(), + }, + Instant::now() + Duration::from_millis(10), + ) + .await; + + assert!(matches!(result, Err(ExecServerError::Closed))); + Ok(()) +} + +async fn read_outbound_data( + outbound_rx: &mut futures_mpsc::UnboundedReceiver, +) -> Result { + read_outbound_frame(outbound_rx) + .await? + .into_data() + .map_err(anyhow::Error::from) +} + +async fn read_outbound_frame( + outbound_rx: &mut futures_mpsc::UnboundedReceiver, +) -> Result { + let Message::Binary(payload) = timeout(Duration::from_secs(1), outbound_rx.next()) + .await? + .context("harness closed before sending data")? + else { + anyhow::bail!("expected relay data frame"); + }; + let frame = decode_relay_message_frame(payload.as_ref())?; + assert_eq!(frame.validate()?, RelayFrameBodyKind::Data); + Ok(frame) +} + +async fn connected_controlled_harness() -> Result<( + JsonRpcConnection, + ControlledWebSocketHandle, + futures_mpsc::UnboundedReceiver, +)> { + let (websocket, control, mut outbound_rx) = ControlledWebSocket::new(/*write_permits*/ 2); + let executor_identity = NoiseChannelIdentity::generate()?; + let connection = noise_harness_connection_from_websocket( + websocket, + NoiseHarnessConnectionArgs { + connection_label: "test rendezvous".to_string(), + environment_id: ENVIRONMENT_ID.to_string(), + executor_registration_id: EXECUTOR_REGISTRATION_ID.to_string(), + identity: NoiseChannelIdentity::generate()?, + responder_public_key: executor_identity.public_key(), + harness_key_authorization: "authorization".to_string(), + }, + ); + + let Message::Binary(resume_payload) = timeout(Duration::from_secs(1), outbound_rx.next()) + .await? + .context("harness closed before sending resume")? + else { + anyhow::bail!("expected resume frame"); + }; + let resume = decode_relay_message_frame(resume_payload.as_ref())?; + let Message::Binary(handshake_payload) = timeout(Duration::from_secs(1), outbound_rx.next()) + .await? + .context("harness closed before sending handshake")? + else { + anyhow::bail!("expected handshake frame"); + }; + let handshake = decode_relay_message_frame(handshake_payload.as_ref())?; + let stream_id = handshake.stream_id.clone(); + assert_eq!(stream_id, resume.stream_id); + let prologue = + noise_channel_prologue(ENVIRONMENT_ID, EXECUTOR_REGISTRATION_ID, stream_id.as_str()); + let pending = PendingResponderHandshake::read_request( + &executor_identity, + &prologue, + &handshake.into_handshake_payload()?, + )?; + let (_transport, response) = pending.complete()?; + control.send_inbound(Message::Binary( + encode_relay_message_frame(&RelayMessageFrame::handshake(stream_id, response)).into(), + ))?; + Ok((connection, control, outbound_rx)) +} + +struct ControlledWebSocket { + inbound_rx: futures_mpsc::UnboundedReceiver>, + outbound_tx: futures_mpsc::UnboundedSender, + write_permit_rx: futures_mpsc::UnboundedReceiver<()>, + blocked_write_tx: futures_mpsc::UnboundedSender, + write_waiting: bool, + blocked_writes: usize, + inbound_reads: Arc, +} + +struct ControlledWebSocketHandle { + inbound_tx: futures_mpsc::UnboundedSender>, + write_permit_tx: futures_mpsc::UnboundedSender<()>, + blocked_write_rx: futures_mpsc::UnboundedReceiver, + inbound_reads: Arc, +} + +impl ControlledWebSocket { + fn new( + write_permits: usize, + ) -> ( + Self, + ControlledWebSocketHandle, + futures_mpsc::UnboundedReceiver, + ) { + let (inbound_tx, inbound_rx) = futures_mpsc::unbounded(); + let (outbound_tx, outbound_rx) = futures_mpsc::unbounded(); + let (write_permit_tx, write_permit_rx) = futures_mpsc::unbounded(); + let (blocked_write_tx, blocked_write_rx) = futures_mpsc::unbounded(); + for _ in 0..write_permits { + write_permit_tx + .unbounded_send(()) + .expect("test write permit receiver should stay open"); + } + let inbound_reads = Arc::new(AtomicUsize::new(0)); + ( + Self { + inbound_rx, + outbound_tx, + write_permit_rx, + blocked_write_tx, + write_waiting: false, + blocked_writes: 0, + inbound_reads: Arc::clone(&inbound_reads), + }, + ControlledWebSocketHandle { + inbound_tx, + write_permit_tx, + blocked_write_rx, + inbound_reads, + }, + outbound_rx, + ) + } +} + +impl ControlledWebSocketHandle { + fn send_inbound(&self, message: Message) -> Result<()> { + self.inbound_tx + .unbounded_send(Ok(message)) + .map_err(anyhow::Error::from) + } + + fn grant_writes(&self, count: usize) { + for _ in 0..count { + self.write_permit_tx + .unbounded_send(()) + .expect("test write permit receiver should stay open"); + } + } + + fn inbound_reads(&self) -> usize { + self.inbound_reads.load(Ordering::Acquire) + } + + async fn wait_for_blocked_write(&mut self, expected: usize) -> Result<()> { + let actual = timeout(Duration::from_secs(1), self.blocked_write_rx.next()) + .await? + .context("websocket closed before blocking the expected write")?; + assert_eq!(actual, expected); + Ok(()) + } +} + +impl Sink for ControlledWebSocket { + type Error = std::convert::Infallible; + + fn poll_ready(self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll> { + let this = self.get_mut(); + match Pin::new(&mut this.write_permit_rx).poll_next(cx) { + Poll::Ready(Some(())) => { + this.write_waiting = false; + Poll::Ready(Ok(())) + } + Poll::Ready(None) | Poll::Pending => { + if !this.write_waiting { + this.write_waiting = true; + this.blocked_writes += 1; + this.blocked_write_tx + .unbounded_send(this.blocked_writes) + .expect("test blocked-write receiver should stay open"); + } + Poll::Pending + } + } + } + + fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> { + self.outbound_tx + .unbounded_send(item) + .expect("test outbound receiver should stay open"); + Ok(()) + } + + fn poll_flush( + self: Pin<&mut Self>, + _cx: &mut TaskContext<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close( + self: Pin<&mut Self>, + _cx: &mut TaskContext<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } +} + +impl futures::Stream for ControlledWebSocket { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll> { + let result = Pin::new(&mut self.inbound_rx).poll_next(cx); + if matches!(result, Poll::Ready(Some(_))) { + self.inbound_reads.fetch_add(1, Ordering::Release); + } + result + } +} diff --git a/codex-rs/exec-server/src/noise_relay/message_framing.rs b/codex-rs/exec-server/src/noise_relay/message_framing.rs new file mode 100644 index 0000000000000000000000000000000000000000..af9fbd49986f0e23da32acf44ecee4ac377f536c --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/message_framing.rs @@ -0,0 +1,114 @@ +use bytes::Buf; +use bytes::Bytes; +use bytes::BytesMut; +use codex_exec_server_protocol::JSONRPCMessage; + +use crate::ExecServerError; + +const LENGTH_PREFIX_BYTES: usize = size_of::(); +pub(crate) const MAX_NOISE_JSONRPC_MESSAGE_LEN: usize = 64 * 1024 * 1024; +pub(crate) const NOISE_RECORD_PLAINTEXT_LEN: usize = 60 * 1024; + +/// Serialize one JSON-RPC message into the encrypted record byte stream. +/// +/// Clatter limits an individual Noise message to 65,535 bytes, while valid +/// exec-server responses can be much larger. A four-byte authenticated length +/// prefix lets the caller split this byte stream into bounded Noise records and +/// lets the receiver reconstruct exact JSON-RPC message boundaries. +pub(crate) fn frame_jsonrpc_message(message: &JSONRPCMessage) -> Result, ExecServerError> { + let mut framed = vec![0; LENGTH_PREFIX_BYTES]; + serde_json::to_writer(&mut framed, message)?; + let prefix = message_length_prefix(framed.len() - LENGTH_PREFIX_BYTES)?; + framed[..LENGTH_PREFIX_BYTES].copy_from_slice(&prefix); + Ok(framed) +} + +pub(crate) fn frame_message(message: &[u8]) -> Result, ExecServerError> { + let prefix = message_length_prefix(message.len())?; + let mut framed = Vec::with_capacity(LENGTH_PREFIX_BYTES + message.len()); + framed.extend_from_slice(&prefix); + framed.extend_from_slice(message); + Ok(framed) +} + +fn message_length_prefix(message_len: usize) -> Result<[u8; LENGTH_PREFIX_BYTES], ExecServerError> { + if message_len == 0 || message_len > MAX_NOISE_JSONRPC_MESSAGE_LEN { + return Err(ExecServerError::Protocol( + "Noise relay JSON-RPC message exceeds maximum length".to_string(), + )); + } + Ok((message_len as u32).to_be_bytes()) +} + +/// Incrementally reconstructs authenticated JSON-RPC messages from Noise records. +/// +/// The length prefix is encrypted along with the message. It is still bounded +/// here so a bad authenticated peer cannot grow the reassembly buffer forever. +#[derive(Default)] +pub(crate) struct JsonRpcMessageDecoder { + decoder: MessageDecoder, +} + +impl JsonRpcMessageDecoder { + pub(crate) fn push( + &mut self, + plaintext_record: &[u8], + ) -> Result, ExecServerError> { + self.decoder + .push(plaintext_record)? + .into_iter() + .map(|message| serde_json::from_slice(&message).map_err(Into::into)) + .collect() + } +} + +/// Reassembles opaque application payloads without interpreting their schema. +#[derive(Default)] +pub(crate) struct MessageDecoder { + buffered: BytesMut, +} + +impl MessageDecoder { + /// Append one decrypted record and return all complete framed messages. + pub(crate) fn push(&mut self, plaintext_record: &[u8]) -> Result, ExecServerError> { + if plaintext_record.len() > NOISE_RECORD_PLAINTEXT_LEN { + return Err(ExecServerError::Protocol( + "Noise relay plaintext record exceeds maximum length".to_string(), + )); + } + self.buffered.extend_from_slice(plaintext_record); + + // One record can finish multiple messages, and one message can span + // multiple records. Parse only after the authenticated length prefix + // and the full declared payload are present. + let mut messages = Vec::new(); + while let Some(prefix) = self.buffered.get(..LENGTH_PREFIX_BYTES) { + let message_len = + u32::from_be_bytes([prefix[0], prefix[1], prefix[2], prefix[3]]) as usize; + // Reject the authenticated length before waiting for its payload. + if message_len == 0 || message_len > MAX_NOISE_JSONRPC_MESSAGE_LEN { + return Err(ExecServerError::Protocol( + "Noise relay JSON-RPC message has invalid length".to_string(), + )); + } + let framed_len = LENGTH_PREFIX_BYTES + message_len; + if self.buffered.len() < framed_len { + break; + } + self.buffered.advance(LENGTH_PREFIX_BYTES); + messages.push(self.buffered.split_to(message_len).freeze()); + } + + // Even before a message is complete, keep reassembly memory bounded. + if self.buffered.len() > LENGTH_PREFIX_BYTES + MAX_NOISE_JSONRPC_MESSAGE_LEN { + return Err(ExecServerError::Protocol( + "Noise relay JSON-RPC reassembly buffer exceeds maximum length".to_string(), + )); + } + Ok(messages) + } +} + +#[cfg(test)] +#[path = "message_framing_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/noise_relay/message_framing_tests.rs b/codex-rs/exec-server/src/noise_relay/message_framing_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..563e9e3bade2a86d193c9b5e71991e653c51db2a --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/message_framing_tests.rs @@ -0,0 +1,68 @@ +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCNotification; +use pretty_assertions::assert_eq; + +use super::JsonRpcMessageDecoder; +use super::MAX_NOISE_JSONRPC_MESSAGE_LEN; +use super::NOISE_RECORD_PLAINTEXT_LEN; +use super::frame_jsonrpc_message; +use crate::ExecServerError; + +#[test] +fn fragments_and_reassembles_large_jsonrpc_message() { + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "large/test".to_string(), + params: Some(serde_json::json!({ + "data": "x".repeat(128 * 1024), + })), + }); + let framed = frame_jsonrpc_message(&message).unwrap(); + assert!(framed.len() > 128 * 1024); + + let mut decoder = JsonRpcMessageDecoder::default(); + let mut decoded = Vec::new(); + for record in framed.chunks(NOISE_RECORD_PLAINTEXT_LEN) { + decoded.extend(decoder.push(record).unwrap()); + } + + assert_eq!(decoded, vec![message]); +} + +#[test] +fn rejects_declared_message_length_above_limit_without_payload() { + let mut decoder = JsonRpcMessageDecoder::default(); + let declared_len = (MAX_NOISE_JSONRPC_MESSAGE_LEN as u32 + 1).to_be_bytes(); + + assert!(matches!( + decoder.push(&declared_len), + Err(ExecServerError::Protocol(message)) + if message == "Noise relay JSON-RPC message has invalid length" + )); +} + +#[test] +fn rejects_oversized_plaintext_record() { + let mut decoder = JsonRpcMessageDecoder::default(); + + assert!(matches!( + decoder.push(&vec![0; NOISE_RECORD_PLAINTEXT_LEN + 1]), + Err(ExecServerError::Protocol(message)) + if message == "Noise relay plaintext record exceeds maximum length" + )); +} + +#[test] +fn reassembles_many_messages_from_one_record() { + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "small/test".to_string(), + params: None, + }); + let framed = frame_jsonrpc_message(&message).expect("frame message"); + let message_count = NOISE_RECORD_PLAINTEXT_LEN / framed.len(); + let record = framed.repeat(message_count); + + let mut decoder = JsonRpcMessageDecoder::default(); + let decoded = decoder.push(&record).expect("decode record"); + + assert_eq!(decoded, vec![message; message_count]); +} diff --git a/codex-rs/exec-server/src/noise_relay/mod.rs b/codex-rs/exec-server/src/noise_relay/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..dc6030e13b71f8ddfcde6241609eee59768f6759 --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/mod.rs @@ -0,0 +1,35 @@ +pub(crate) mod executor_stream; +mod harness; +pub(crate) mod message_framing; +mod ordered_ciphertext; +pub(crate) mod stream_handler; + +use tokio_tungstenite::tungstenite::protocol::WebSocketConfig; + +use crate::ExecServerError; + +pub(crate) use harness::NoiseHarnessConnectionArgs; +pub(crate) use harness::noise_harness_connection_from_websocket_with_readiness; + +pub(crate) const NOISE_RELAY_RESET_REASON: &str = "noise_relay_protocol_error"; + +// This bounds allocation in tungstenite before protobuf and Noise record +// validation run. It comfortably fits one maximum Noise record plus metadata. +const MAX_NOISE_RELAY_WEBSOCKET_MESSAGE_SIZE: usize = 256 * 1024; + +/// Return the websocket limits required by every Noise relay endpoint. +pub(crate) fn noise_relay_websocket_config() -> WebSocketConfig { + WebSocketConfig::default() + .max_frame_size(Some(MAX_NOISE_RELAY_WEBSOCKET_MESSAGE_SIZE)) + .max_message_size(Some(MAX_NOISE_RELAY_WEBSOCKET_MESSAGE_SIZE)) +} + +fn take_next_sequence(next_seq: &mut u32) -> Result { + // Never wrap: relay sequence is the explicit ordering key for an implicit + // Noise nonce. Reusing zero after u32::MAX would be ambiguous and unsafe. + let seq = *next_seq; + *next_seq = next_seq.checked_add(1).ok_or_else(|| { + ExecServerError::Protocol("Noise relay sequence number exhausted".to_string()) + })?; + Ok(seq) +} diff --git a/codex-rs/exec-server/src/noise_relay/ordered_ciphertext.rs b/codex-rs/exec-server/src/noise_relay/ordered_ciphertext.rs new file mode 100644 index 0000000000000000000000000000000000000000..92bbd291d72f8130881a9364e43323a5566571ad --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/ordered_ciphertext.rs @@ -0,0 +1,70 @@ +use std::collections::BTreeMap; + +use crate::ExecServerError; + +const MAX_REORDER_DISTANCE: u32 = 64; +const MAX_PENDING_BYTES: usize = 1024 * 1024; + +/// Reorders relay records before they reach Noise's implicit receive nonce. +/// The window is bounded, and each sequence number is released at most once. +#[derive(Default)] +pub(crate) struct OrderedCiphertextFrames { + next_seq: u32, + pending: BTreeMap>, + pending_bytes: usize, +} + +impl OrderedCiphertextFrames { + /// Accept one relay record and return the newly contiguous ciphertext run. + /// + /// Returns nothing for duplicates or while a gap remains. Closing a gap also + /// releases any buffered records that now follow it contiguously. + pub(crate) fn push( + &mut self, + seq: u32, + payload: Vec, + ) -> Result>, ExecServerError> { + // Keep the first ciphertext for a sequence. Later copies are duplicates. + if seq < self.next_seq || self.pending.contains_key(&seq) { + return Ok(Vec::new()); + } + if seq > self.next_seq { + // Bound both the sequence gap and buffered bytes. + if seq - self.next_seq > MAX_REORDER_DISTANCE { + return Err(ExecServerError::Protocol( + "Noise relay ciphertext exceeds reorder window".to_string(), + )); + } + let pending_bytes = self.pending_bytes + payload.len(); + if pending_bytes > MAX_PENDING_BYTES { + return Err(ExecServerError::Protocol( + "Noise relay pending ciphertext buffer is full".to_string(), + )); + } + self.pending.insert(seq, payload); + self.pending_bytes = pending_bytes; + return Ok(Vec::new()); + } + + // Release the expected record and anything now contiguous behind it. + let mut ready = vec![payload]; + self.advance()?; + while let Some(payload) = self.pending.remove(&self.next_seq) { + self.pending_bytes -= payload.len(); + ready.push(payload); + self.advance()?; + } + Ok(ready) + } + + fn advance(&mut self) -> Result<(), ExecServerError> { + self.next_seq = self.next_seq.checked_add(1).ok_or_else(|| { + ExecServerError::Protocol("Noise relay sequence number exhausted".to_string()) + })?; + Ok(()) + } +} + +#[cfg(test)] +#[path = "ordered_ciphertext_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/noise_relay/ordered_ciphertext_tests.rs b/codex-rs/exec-server/src/noise_relay/ordered_ciphertext_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..6aa86fcdf0559f68d6e05b55d0f6db60094b2f09 --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/ordered_ciphertext_tests.rs @@ -0,0 +1,52 @@ +use pretty_assertions::assert_eq; + +use super::MAX_PENDING_BYTES; +use super::OrderedCiphertextFrames; + +#[test] +fn releases_ciphertexts_only_in_nonce_order() { + let mut frames = OrderedCiphertextFrames::default(); + + assert_eq!( + frames.push(/*seq*/ 1, b"second".to_vec()).unwrap(), + Vec::>::new() + ); + assert_eq!( + frames.push(/*seq*/ 0, b"first".to_vec()).unwrap(), + vec![b"first".to_vec(), b"second".to_vec()] + ); +} + +#[test] +fn ignores_duplicate_ciphertexts_without_replacing_buffered_record() { + let mut frames = OrderedCiphertextFrames::default(); + + assert_eq!( + frames.push(/*seq*/ 1, b"first copy".to_vec()).unwrap(), + Vec::>::new() + ); + assert_eq!( + frames.push(/*seq*/ 1, b"replacement".to_vec()).unwrap(), + Vec::>::new() + ); + assert_eq!( + frames.push(/*seq*/ 0, b"zero".to_vec()).unwrap(), + vec![b"zero".to_vec(), b"first copy".to_vec()] + ); + assert_eq!( + frames.push(/*seq*/ 0, b"duplicate".to_vec()).unwrap(), + Vec::>::new() + ); +} + +#[test] +fn rejects_unbounded_reordering() { + let mut frames = OrderedCiphertextFrames::default(); + + assert!(frames.push(/*seq*/ 65, Vec::new()).is_err()); + assert!( + frames + .push(/*seq*/ 1, vec![0; MAX_PENDING_BYTES + 1]) + .is_err() + ); +} diff --git a/codex-rs/exec-server/src/noise_relay/stream_handler.rs b/codex-rs/exec-server/src/noise_relay/stream_handler.rs new file mode 100644 index 0000000000000000000000000000000000000000..b0c610933c09d3b98bfe10fa6f7e5e2e744b7131 --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/stream_handler.rs @@ -0,0 +1,89 @@ +use std::future::Future; +use std::sync::Arc; + +use bytes::Bytes; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_protocol::protocol::W3cTraceContext; +use tokio::sync::mpsc; +use tokio::sync::watch; +use tokio::task::JoinHandle; + +use crate::ExecServerError; +use crate::connection::JsonRpcConnection; +use crate::connection::JsonRpcConnectionEvent; +use crate::connection::JsonRpcTransport; +use crate::noise_relay::message_framing::frame_jsonrpc_message; +use crate::server::ConnectionProcessor; +use crate::telemetry::ExecutorRegistration; + +pub(crate) struct NoiseStreamConnection { + pub(crate) outgoing_tx: mpsc::Sender, + pub(crate) incoming_rx: mpsc::Receiver, + pub(crate) disconnected_rx: watch::Receiver, + pub(crate) writer_task: JoinHandle<()>, + pub(crate) executor_registration: Option>, +} + +pub(crate) struct NoiseOutboundMessage { + /// Payload with the authenticated length prefix supplied by message_framing. + pub(crate) framed: Vec, + pub(crate) trace: Option, +} + +/// Adapts complete Noise payloads to one owner without adding an inbound queue. +/// Decoding is synchronous so execution spans begin before queue admission. +pub(crate) trait NoiseStreamHandler: Clone + Send + 'static { + type Incoming: Send + 'static; + type Outgoing: Send + 'static; + + fn decode(payload: Bytes) -> Result; + fn encode(message: Self::Outgoing) -> Result; + fn run_connection( + self, + connection: NoiseStreamConnection, + ) -> impl Future + Send; +} + +impl NoiseStreamHandler for ConnectionProcessor { + type Incoming = JsonRpcConnectionEvent; + type Outgoing = JSONRPCMessage; + + fn decode(payload: Bytes) -> Result { + Ok(JsonRpcConnectionEvent::message(serde_json::from_slice( + &payload, + )?)) + } + + fn encode(message: Self::Outgoing) -> Result { + let framed = frame_jsonrpc_message(&message)?; + let trace = match message { + JSONRPCMessage::Request(request) => request.trace, + JSONRPCMessage::Notification(_) + | JSONRPCMessage::Response(_) + | JSONRPCMessage::Error(_) => None, + }; + Ok(NoiseOutboundMessage { framed, trace }) + } + + async fn run_connection( + self, + connection: NoiseStreamConnection, + ) { + ConnectionProcessor::run_registered_connection( + &self, + JsonRpcConnection { + outgoing_tx: connection.outgoing_tx, + incoming_rx: connection.incoming_rx, + disconnected_rx: connection.disconnected_rx, + task_handles: vec![connection.writer_task], + transport: JsonRpcTransport::Plain, + }, + connection.executor_registration, + ) + .await; + } +} + +#[cfg(test)] +#[path = "stream_handler_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/noise_relay/stream_handler_tests.rs b/codex-rs/exec-server/src/noise_relay/stream_handler_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..0015176684495f656c72b9b72cb746157c9215a3 --- /dev/null +++ b/codex-rs/exec-server/src/noise_relay/stream_handler_tests.rs @@ -0,0 +1,41 @@ +use std::time::Instant; + +use anyhow::Result; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::RequestId; +use pretty_assertions::assert_eq; + +use super::NoiseStreamHandler; +use crate::connection::JsonRpcConnectionEvent; +use crate::server::ConnectionProcessor; + +#[test] +fn local_decode_starts_request_span_before_queueing() -> Result<()> { + let _subscriber = tracing::subscriber::set_default(tracing_subscriber::registry()); + tracing::callsite::rebuild_interest_cache(); + + let request = JSONRPCRequest { + id: RequestId::Integer(1), + method: "test/queued".to_string(), + params: None, + trace: None, + }; + let before_decode = Instant::now(); + let event = ::decode( + serde_json::to_vec(&JSONRPCMessage::Request(request.clone()))?.into(), + )?; + let after_decode = Instant::now(); + let JsonRpcConnectionEvent::QueuedRequest { + request: actual_request, + request_span, + queued_at, + } = event + else { + panic!("local ingress must create the queued request synchronously"); + }; + assert_eq!(actual_request, request); + assert!(!request_span.is_disabled()); + assert!((before_decode..=after_decode).contains(&queued_at)); + Ok(()) +} diff --git a/codex-rs/exec-server/src/process.rs b/codex-rs/exec-server/src/process.rs new file mode 100644 index 0000000000000000000000000000000000000000..5e795faa1af2634ffd5eee36608f03e7459e07f7 --- /dev/null +++ b/codex-rs/exec-server/src/process.rs @@ -0,0 +1,317 @@ +use std::collections::VecDeque; +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; + +use codex_network_proxy::NetworkPolicyDecider; +use codex_sandboxing::SandboxType; +use tokio::sync::broadcast; +use tokio::sync::watch; + +use crate::ExecServerError; +use crate::ProcessId; +use crate::protocol::ExecParams; +use crate::protocol::ProcessOutputChunk; +use crate::protocol::ProcessSandboxType; +use crate::protocol::ProcessSignal; +use crate::protocol::ReadResponse; +use crate::protocol::WriteResponse; + +pub struct StartedExecProcess { + pub process: Arc, + /// `None` means the exec-server peer did not report its sandbox type. + pub sandbox_type: Option, +} + +pub(crate) fn sandbox_type_from_protocol( + sandbox_type: Option, +) -> Option { + match sandbox_type { + None => None, + Some(ProcessSandboxType::None) => Some(SandboxType::None), + Some(ProcessSandboxType::MacosSeatbelt) => Some(SandboxType::MacosSeatbelt), + Some(ProcessSandboxType::LinuxSeccomp) => Some(SandboxType::LinuxSeccomp), + Some(ProcessSandboxType::WindowsRestrictedToken) => { + Some(SandboxType::WindowsRestrictedToken) + } + Some(ProcessSandboxType::WindowsMxc) => Some(SandboxType::WindowsMxc), + } +} + +/// Pushed process events for consumers that want to follow process output as it +/// arrives instead of polling retained output with [`ExecProcess::read`]. +/// +/// The stream is scoped to one [`ExecProcess`] handle. `Output` events carry +/// stdout, stderr, or pty bytes. `Exited` reports the process exit status, while +/// `Closed` means all output streams have ended and no more output events will +/// arrive. `Failed` is used when the process session cannot continue, for +/// example because the remote environment connection disconnected. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ExecProcessEvent { + Output(ProcessOutputChunk), + Exited { + seq: u64, + exit_code: i32, + sandbox_denied: Option, + }, + Closed { + seq: u64, + }, + Failed(String), +} + +/// Replay buffer plus live fan-out for pushed process events. +/// +/// New subscribers first drain a bounded replay history, then continue on the +/// live broadcast channel. The history is bounded by event count and retained +/// output bytes: count protects against many tiny events, while bytes protects +/// against a few very large output chunks. +#[derive(Clone)] +pub(crate) struct ExecProcessEventLog { + inner: Arc, +} + +struct ExecProcessEventLogInner { + history: StdMutex, + live_tx: broadcast::Sender, + event_capacity: usize, + byte_capacity: usize, +} + +#[derive(Default)] +struct ExecProcessEventHistory { + events: VecDeque, + retained_bytes: usize, +} + +impl ExecProcessEvent { + /// Sequence number used to order process-owned events. + /// + /// `Failed` is intentionally unsequenced because it is synthesized by the + /// client when the session or transport fails, not emitted by the process. + pub(crate) fn seq(&self) -> Option { + match self { + ExecProcessEvent::Output(chunk) => Some(chunk.seq), + ExecProcessEvent::Exited { seq, .. } | ExecProcessEvent::Closed { seq } => Some(*seq), + ExecProcessEvent::Failed(_) => None, + } + } + + fn retained_len(&self) -> usize { + match self { + ExecProcessEvent::Output(chunk) => chunk.chunk.0.len(), + ExecProcessEvent::Failed(message) => message.len(), + ExecProcessEvent::Exited { .. } | ExecProcessEvent::Closed { .. } => 0, + } + } +} + +impl ExecProcessEventLog { + pub(crate) fn new(event_capacity: usize, byte_capacity: usize) -> Self { + let (live_tx, _live_rx) = broadcast::channel(event_capacity); + Self { + inner: Arc::new(ExecProcessEventLogInner { + history: StdMutex::new(ExecProcessEventHistory::default()), + live_tx, + event_capacity, + byte_capacity, + }), + } + } + + pub(crate) fn publish(&self, event: ExecProcessEvent) { + let mut history = self + .inner + .history + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + history.retained_bytes += event.retained_len(); + history.events.push_back(event.clone()); + while history.events.len() > self.inner.event_capacity + || history.retained_bytes > self.inner.byte_capacity + { + let Some(evicted) = history.events.pop_front() else { + break; + }; + history.retained_bytes = history + .retained_bytes + .saturating_sub(evicted.retained_len()); + } + + let _ = self.inner.live_tx.send(event); + } + + pub(crate) fn subscribe(&self) -> ExecProcessEventReceiver { + let history = self + .inner + .history + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let live_rx = self.inner.live_tx.subscribe(); + let replay = history.events.iter().cloned().collect(); + + ExecProcessEventReceiver { + replay, + live_rx, + _keepalive: None, + } + } +} + +pub struct ExecProcessEventReceiver { + replay: VecDeque, + live_rx: broadcast::Receiver, + _keepalive: Option>, +} + +impl ExecProcessEventReceiver { + /// Returns a receiver that remains open without yielding events. + pub fn empty() -> Self { + let (live_tx, live_rx) = broadcast::channel(1); + Self { + replay: VecDeque::new(), + live_rx, + _keepalive: Some(live_tx), + } + } + + /// Returns the next replayed or live event. + /// + /// `Lagged` means this receiver fell behind the bounded live channel. The + /// caller should recover through [`ExecProcess::read`] using the last + /// delivered sequence number, then continue receiving pushed events. + pub async fn recv(&mut self) -> Result { + if let Some(event) = self.replay.pop_front() { + return Ok(event); + } + + self.live_rx.recv().await + } +} + +/// Handle for an executor-managed process. +/// +/// Implementations must support both retained-output reads and pushed events: +/// `read` is the request/response API for callers that want to page through +/// buffered output, while `subscribe_events` is the streaming API for callers +/// that want output and lifecycle changes delivered as they happen. +pub trait ExecProcess: Send + Sync { + fn process_id(&self) -> &ProcessId; + + fn subscribe_wake(&self) -> watch::Receiver; + + fn subscribe_events(&self) -> ExecProcessEventReceiver; + + fn read( + &self, + after_seq: Option, + max_bytes: Option, + wait_ms: Option, + ) -> ExecProcessFuture<'_, ReadResponse>; + + fn write(&self, chunk: Vec) -> ExecProcessFuture<'_, WriteResponse>; + + fn signal(&self, signal: ProcessSignal) -> ExecProcessFuture<'_, ()>; + + fn terminate(&self) -> ExecProcessFuture<'_, ()>; +} + +pub type ExecProcessFuture<'a, T> = + Pin> + Send + 'a>>; + +pub trait ExecBackend: Send + Sync { + fn start(&self, params: ExecParams) -> ExecBackendFuture<'_>; + + /// Captures a local shell snapshot without starting the requested command. + /// Failures must remain retryable by real commands. Remote backends do not + /// support this operation; callers should leave them on the lazy path. + fn prewarm_shell_snapshot(&self, _params: ExecParams) -> ExecProcessFuture<'_, ()> { + Box::pin(async { + Err(ExecServerError::Protocol( + "exec backend does not support shell snapshot prewarming".to_string(), + )) + }) + } + + /// Starts a process with an authoritative controller-side policy decider. + fn start_with_network_policy_decider( + &self, + _params: ExecParams, + _decider: Arc, + ) -> ExecBackendFuture<'_> { + Box::pin(async { + Err(ExecServerError::Protocol( + "exec backend does not support remote network policy decisions".to_string(), + )) + }) + } +} + +pub type ExecBackendFuture<'a> = + Pin> + Send + 'a>>; + +#[cfg(test)] +mod tests { + use pretty_assertions::assert_eq; + use tokio::time::Duration; + use tokio::time::timeout; + + use super::ExecProcessEvent; + use super::ExecProcessEventLog; + use super::ExecProcessEventReceiver; + use crate::protocol::ExecOutputStream; + use crate::protocol::ProcessOutputChunk; + + #[tokio::test] + async fn empty_event_receiver_stays_open() { + let mut events = ExecProcessEventReceiver::empty(); + + assert!( + timeout(Duration::from_millis(10), events.recv()) + .await + .is_err() + ); + } + + #[tokio::test] + async fn event_history_replay_is_bounded_by_retained_bytes() { + let log = ExecProcessEventLog::new(/*event_capacity*/ 8, /*byte_capacity*/ 3); + + log.publish(ExecProcessEvent::Output(ProcessOutputChunk { + seq: 1, + stream: ExecOutputStream::Stdout, + chunk: b"large".to_vec().into(), + })); + log.publish(ExecProcessEvent::Exited { + seq: 2, + exit_code: 0, + sandbox_denied: Some(false), + }); + log.publish(ExecProcessEvent::Closed { seq: 3 }); + + let mut events = log.subscribe(); + let replay = vec![ + timeout(Duration::from_secs(1), events.recv()) + .await + .expect("exit event replay should not time out") + .expect("exit event replay should be available"), + timeout(Duration::from_secs(1), events.recv()) + .await + .expect("closed event replay should not time out") + .expect("closed event replay should be available"), + ]; + + assert_eq!( + replay, + vec![ + ExecProcessEvent::Exited { + seq: 2, + exit_code: 0, + sandbox_denied: Some(false), + }, + ExecProcessEvent::Closed { seq: 3 }, + ] + ); + } +} diff --git a/codex-rs/exec-server/src/process_sandbox.rs b/codex-rs/exec-server/src/process_sandbox.rs new file mode 100644 index 0000000000000000000000000000000000000000..035bfa8831b16ad2dd7447c27d178311a27372c4 --- /dev/null +++ b/codex-rs/exec-server/src/process_sandbox.rs @@ -0,0 +1,424 @@ +use std::collections::HashMap; +use std::sync::Arc; + +use crate::process_telemetry::ProcessTelemetry; +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_file_system::WindowsSandboxSelection; +use codex_network_proxy::CUSTOM_CA_ENV_KEYS; +use codex_network_proxy::ManagedNetworkSandboxContext; +use codex_network_proxy::ManagedProxyRouting; +use codex_network_proxy::NetworkPolicyAuditObserver; +use codex_network_proxy::NetworkPolicyDecider; +use codex_network_proxy::NetworkProxy; +use codex_network_proxy::NetworkProxyHandle; +use codex_network_proxy::NetworkProxyState; +use codex_network_proxy::RemoteNetworkProxyLaunchConfig; +use codex_network_proxy::is_managed_mitm_ca_trust_bundle_path; +#[cfg(target_os = "windows")] +use codex_network_proxy::strip_managed_proxy_env; +use codex_protocol::config_types::WindowsSandboxLevel; +use codex_protocol::models::PermissionProfile; +use codex_sandboxing::SandboxCommand; +use codex_sandboxing::SandboxDirectSpawnTransformRequest; +use codex_sandboxing::SandboxManager; +use codex_sandboxing::SandboxTransformRequest; +use codex_sandboxing::SandboxType; +use codex_sandboxing::WindowsSandboxFilesystemOverrides; +use codex_sandboxing::WindowsSandboxProxySettingsMode; +use codex_sandboxing::WindowsSandboxSpawnRequest; +use codex_sandboxing::resolve_windows_elevated_filesystem_overrides; +use codex_sandboxing::resolve_windows_restricted_token_filesystem_overrides; +use codex_sandboxing::windows_sandbox_uses_elevated_backend; +use codex_sandboxing::with_managed_mitm_ca_readable_root; +use codex_utils_absolute_path::AbsolutePathBuf; +use codex_utils_path_uri::PathUri; + +#[cfg(unix)] +use crate::CODEX_ARG0_EXEC_HELPER_ARG1; +use crate::ExecServerRuntimePaths; +use crate::protocol::ExecParams; +use crate::rpc::internal_error; +use crate::rpc::invalid_params; +use crate::sandbox_selection::select_sandbox; + +pub(crate) struct PreparedExecRequest { + pub(crate) command: Vec, + pub(crate) cwd: AbsolutePathBuf, + pub(crate) env: HashMap, + pub(crate) arg0: Option, + pub(crate) sandbox: SandboxType, + pub(crate) network_proxy_handle: Option, + windows_sandbox: Option, +} + +struct PreparedWindowsSandboxRequest { + permission_profile: PermissionProfile, + workspace_roots: Vec, + windows_sandbox_level: WindowsSandboxLevel, + proxy_enforced: bool, + network_proxy_restricting_sid: Option, + proxy_settings_mode: WindowsSandboxProxySettingsMode, + filesystem_overrides: Option, + use_private_desktop: bool, +} + +impl PreparedExecRequest { + pub(crate) fn windows_sandbox_spawn_request(&self) -> Option> { + self.windows_sandbox + .as_ref() + .map(|request| WindowsSandboxSpawnRequest { + permission_profile: &request.permission_profile, + workspace_roots: &request.workspace_roots, + windows_sandbox_level: request.windows_sandbox_level, + proxy_enforced: request.proxy_enforced, + network_proxy_restricting_sid: request.network_proxy_restricting_sid.as_deref(), + proxy_settings_mode: request.proxy_settings_mode, + filesystem_overrides: request.filesystem_overrides.as_ref(), + use_private_desktop: request.use_private_desktop, + }) + } +} + +pub(crate) async fn prepare_exec_request_with_telemetry( + params: &ExecParams, + env: HashMap, + runtime_paths: Option<&ExecServerRuntimePaths>, + network_policy_decider: Option>, + network_policy_audit_observer: Option, + telemetry: &ProcessTelemetry, +) -> Result { + if let Some(sandbox) = params.sandbox.as_ref() + && sandbox.windows_sandbox_selection == WindowsSandboxSelection::Mxc + { + if params.arg0.is_some() || sandbox.windows_sandbox_private_desktop { + return Err(invalid_params( + "MXC custom argv0 and private-desktop launches are not supported".to_owned(), + )); + } + if !codex_sandboxing::windows_mxc_available() { + return Err(invalid_params( + "native MXC is unavailable on this executor".to_owned(), + )); + } + } + #[cfg(target_os = "windows")] + let mut env = env; + #[cfg(target_os = "windows")] + let network_proxy = if params.sandbox.is_none() { + // Shared Windows ingress selects a route from the sandbox token's SID. Native launches + // have no route SID, so leave them direct. + if params.network_proxy.is_some() { + strip_managed_proxy_env(&mut env); + } + None + } else { + params.network_proxy.as_ref() + }; + #[cfg(not(target_os = "windows"))] + let network_proxy = params.network_proxy.as_ref(); + + let (env, managed_network, network_proxy_handle, network_proxy_restricting_sid) = + prepare_managed_network( + params.managed_network.as_ref(), + network_proxy, + if params.sandbox.as_ref().is_some_and(|sandbox| { + sandbox.windows_sandbox_selection == WindowsSandboxSelection::Mxc + }) { + ManagedProxyRouting::DedicatedListeners + } else { + ManagedProxyRouting::SharedIngress + }, + env, + network_policy_decider, + network_policy_audit_observer, + telemetry, + ) + .await?; + let Some(sandbox_context) = params.sandbox.as_ref() else { + return Ok(PreparedExecRequest { + command: params.argv.clone(), + cwd: native_path(¶ms.cwd, "cwd")?, + env, + arg0: params.arg0.clone(), + sandbox: SandboxType::None, + network_proxy_handle, + windows_sandbox: None, + }); + }; + let windows_sandbox_proxy_settings_mode = sandbox_context + .windows_sandbox_proxy_settings_mode + .unwrap_or_default(); + let runtime_paths = runtime_paths + .ok_or_else(|| invalid_params("sandbox runtime paths are not configured".to_string()))?; + // TODO(jif): Transport permissions before orchestrator-local paths are materialized, + // then resolve executor-local helper and workspace paths here. + let permissions: PermissionProfile = sandbox_context + .permissions + .clone() + .try_into() + .map_err(|err| invalid_params(format!("invalid sandbox permission path URI: {err}")))?; + let sandbox_policy_cwd = sandbox_context.cwd.as_ref().unwrap_or(¶ms.cwd); + let native_sandbox_policy_cwd = native_path(sandbox_policy_cwd, "sandbox cwd")?; + let native_workspace_roots = sandbox_context + .workspace_roots + .iter() + .map(|root| native_path(root, "sandbox workspace root")) + .collect::, _>>()?; + let workspace_roots = native_workspace_roots.as_slice(); + let permissions = permissions.materialize_project_roots_with_workspace_roots(workspace_roots); + let managed_mitm_ca_trust_bundle_path = managed_network.as_ref().and_then(|_| { + CUSTOM_CA_ENV_KEYS.iter().find_map(|key| { + let path = env.get(*key)?; + if !is_managed_mitm_ca_trust_bundle_path(path) { + return None; + } + AbsolutePathBuf::from_absolute_path(path).ok() + }) + }); + let permissions = with_managed_mitm_ca_readable_root( + permissions, + managed_mitm_ca_trust_bundle_path.as_ref(), + native_sandbox_policy_cwd.as_path(), + ); + #[cfg(unix)] + let (file_system_policy, network_policy) = permissions.to_runtime_permissions(); + #[cfg(unix)] + let sandbox_helper_paths = params + .arg0 + .iter() + .map(|_| runtime_paths.codex_self_exe.clone()) + .collect::>(); + // Bubblewrap launches the configured helper, which may re-enter this executable to apply + // seccomp, so the outer filesystem sandbox must expose both paths. + #[cfg(target_os = "linux")] + let sandbox_helper_paths = { + let mut sandbox_helper_paths = sandbox_helper_paths; + if !sandbox_helper_paths.contains(&runtime_paths.codex_self_exe) { + sandbox_helper_paths.push(runtime_paths.codex_self_exe.clone()); + } + sandbox_helper_paths.extend(runtime_paths.codex_linux_sandbox_exe.iter().cloned()); + sandbox_helper_paths + }; + #[cfg(unix)] + let file_system_policy = file_system_policy + .with_additional_readable_roots(native_sandbox_policy_cwd.as_path(), &sandbox_helper_paths); + #[cfg(unix)] + let permissions = PermissionProfile::from_runtime_permissions_with_enforcement( + permissions.enforcement(), + &file_system_policy, + network_policy, + ); + let sandbox_manager = SandboxManager::new(); + #[cfg(target_os = "macos")] + let sandbox_manager = sandbox_manager + .with_allowed_symlinked_codex_home(runtime_paths.allowed_symlinked_codex_home.clone()); + let (sandbox, windows_sandbox_level) = select_sandbox( + &sandbox_manager, + &permissions, + sandbox_context, + params.enforce_managed_network, + ); + if sandbox == SandboxType::None { + return Err(invalid_params( + "sandbox intent cannot be enforced on this executor".to_string(), + )); + } + let (program, args) = params + .argv + .split_first() + .ok_or_else(|| invalid_params("argv must not be empty".to_string()))?; + #[cfg(unix)] + let (program, args) = params.arg0.as_ref().map_or_else( + || (program.into(), args.to_vec()), + |arg0| { + let mut helper_args = Vec::with_capacity(params.argv.len() + 2); + helper_args.push(CODEX_ARG0_EXEC_HELPER_ARG1.to_string()); + helper_args.push(arg0.clone()); + helper_args.extend(params.argv.iter().cloned()); + ( + runtime_paths + .codex_self_exe + .as_path() + .as_os_str() + .to_owned(), + helper_args, + ) + }, + ); + #[cfg(not(unix))] + let (program, args) = (program.into(), args.to_vec()); + let transform_request = SandboxDirectSpawnTransformRequest { + workspace_roots, + windows_sandbox_proxy_settings_mode, + transform: SandboxTransformRequest { + command: SandboxCommand { + program, + args, + cwd: params.cwd.clone(), + env, + managed_network, + additional_permissions: None, + }, + permissions: &permissions, + sandbox, + enforce_managed_network: params.enforce_managed_network, + environment_id: None, + network: None, + sandbox_policy_cwd, + sandbox_exe: if cfg!(windows) { + Some(runtime_paths.codex_self_exe.as_path()) + } else { + runtime_paths.codex_linux_sandbox_exe.as_deref() + }, + use_legacy_landlock: sandbox_context.use_legacy_landlock, + windows_sandbox_level: windows_sandbox_level.unwrap_or(WindowsSandboxLevel::Disabled), + windows_sandbox_private_desktop: sandbox_context.windows_sandbox_private_desktop, + }, + }; + let mut request = if sandbox == SandboxType::WindowsRestrictedToken { + // The shared launcher invokes the native Windows session spawner directly. + sandbox_manager.transform(transform_request.transform) + } else { + sandbox_manager.transform_for_direct_spawn(transform_request) + } + .map_err(|err| invalid_params(format!("failed to prepare process sandbox: {err}")))?; + let windows_sandbox = if sandbox == SandboxType::WindowsRestrictedToken { + let windows_sandbox_level = windows_sandbox_level.ok_or_else(|| { + invalid_params("restricted token sandbox requires a sandbox level".to_string()) + })?; + request.arg0 = params.arg0.clone(); + let proxy_enforced = params.enforce_managed_network; + let use_elevated = windows_sandbox_uses_elevated_backend(windows_sandbox_level); + let filesystem_overrides = if use_elevated { + resolve_windows_elevated_filesystem_overrides( + sandbox, + &permissions, + &native_sandbox_policy_cwd, + use_elevated, + ) + } else { + resolve_windows_restricted_token_filesystem_overrides( + sandbox, + &permissions, + &native_sandbox_policy_cwd, + windows_sandbox_level, + ) + } + .map_err(|err| invalid_params(format!("failed to prepare process sandbox: {err}")))?; + Some(PreparedWindowsSandboxRequest { + permission_profile: permissions, + workspace_roots: native_workspace_roots, + windows_sandbox_level, + proxy_enforced, + network_proxy_restricting_sid, + proxy_settings_mode: windows_sandbox_proxy_settings_mode, + filesystem_overrides, + use_private_desktop: sandbox_context.windows_sandbox_private_desktop, + }) + } else { + None + }; + Ok(PreparedExecRequest { + command: request.command, + cwd: native_path(&request.cwd, "cwd")?, + env: request.env, + arg0: request.arg0, + sandbox: request.sandbox, + network_proxy_handle, + windows_sandbox, + }) +} + +async fn prepare_managed_network( + managed_network: Option<&ManagedNetworkSandboxContext>, + network_proxy: Option<&RemoteNetworkProxyLaunchConfig>, + routing: ManagedProxyRouting, + env: HashMap, + network_policy_decider: Option>, + network_policy_audit_observer: Option, + telemetry: &ProcessTelemetry, +) -> Result< + ( + HashMap, + Option, + Option, + Option, + ), + JSONRPCErrorError, +> { + let Some(network_proxy) = network_proxy.cloned() else { + return Ok((env, managed_network.cloned(), None, None)); + }; + let mut state = NetworkProxyState::from_remote_launch_config(network_proxy) + .map_err(|err| invalid_params(format!("invalid network proxy config: {err}")))?; + if let Some(observer) = network_policy_audit_observer { + state.set_policy_audit_observer(observer); + } + if let Some(launch_context) = &telemetry.launch_context { + state.set_launch_span_context(launch_context.clone()); + } + state.set_process_log_metadata(codex_network_proxy::NetworkProxyProcessLogMetadata { + thread_id: telemetry.thread_id.clone(), + tool_call_id: telemetry.tool_call_id.clone(), + executor_identity: telemetry + .executor_registration + .as_ref() + .map(|registration| codex_network_proxy::ExecutorLogIdentity { + environment_id: registration.environment_id.clone(), + registration_id: registration.executor_registration_id.clone(), + }), + }); + let mut builder = NetworkProxy::builder() + .state(Arc::new(state)) + .managed_proxy_routing(routing); + if let Some(network_policy_decider) = network_policy_decider { + builder = builder.policy_decider_arc(network_policy_decider); + } + let proxy = builder + .build() + .await + .map_err(|err| internal_error(format!("failed to build executor network proxy: {err}")))?; + let handle = proxy + .run() + .await + .map_err(|err| internal_error(format!("failed to start executor network proxy: {err}")))?; + #[cfg(target_os = "windows")] + let network_proxy_restricting_sid = if routing == ManagedProxyRouting::SharedIngress { + Some( + proxy + .network_proxy_restricting_sid(/*environment_id*/ None) + .ok_or_else(|| { + internal_error( + "managed Windows proxy route is missing its restricting SID".to_string(), + ) + })?, + ) + } else { + None + }; + #[cfg(not(target_os = "windows"))] + let network_proxy_restricting_sid = None; + let prepared = proxy + .prepare_for_optional_environment(env, /*environment_id*/ None) + .map_err(|err| { + internal_error(format!("failed to prepare executor network proxy: {err}")) + })?; + Ok(( + prepared.env, + Some(prepared.sandbox_context), + Some(handle), + network_proxy_restricting_sid, + )) +} + +fn native_path(path: &PathUri, label: &str) -> Result { + path.to_abs_path().map_err(|err| { + invalid_params(format!( + "{label} URI `{path}` is not valid on this exec-server host: {err}" + )) + }) +} + +#[cfg(test)] +#[path = "process_sandbox_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/process_sandbox_tests.rs b/codex-rs/exec-server/src/process_sandbox_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..7e66b2e419c55599cbc8bfabda53a67aeaccb76d --- /dev/null +++ b/codex-rs/exec-server/src/process_sandbox_tests.rs @@ -0,0 +1,686 @@ +use std::collections::HashMap; +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCErrorError; +#[cfg(unix)] +use codex_file_system::WindowsSandboxSelection; +#[cfg(target_os = "macos")] +use codex_network_proxy::ManagedNetworkSandboxContext; +use codex_network_proxy::NetworkPolicyAuditObserver; +use codex_network_proxy::NetworkPolicyDecider; +use codex_network_proxy::NetworkProxyConfig; +#[cfg(target_os = "macos")] +use codex_network_proxy::NetworkUnixSocketPermission; +#[cfg(target_os = "macos")] +use codex_network_proxy::NetworkUnixSocketPermissions; +use codex_network_proxy::PROXY_ATTRIBUTION_TOKEN_ENV_KEY; +use codex_network_proxy::RemoteNetworkProxyConfig; +use codex_network_proxy::RemoteNetworkProxyLaunchConfig; +#[cfg(windows)] +use codex_protocol::config_types::WindowsSandboxLevel; +#[cfg(any(unix, windows))] +use codex_protocol::models::PermissionProfile; +#[cfg(target_os = "linux")] +use codex_sandboxing::landlock::CODEX_LINUX_SANDBOX_ARG0; +use codex_utils_absolute_path::AbsolutePathBuf; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +#[cfg(windows)] +use test_case::test_case; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::time::timeout; + +use super::PreparedExecRequest; +use super::prepare_exec_request_with_telemetry; +#[cfg(unix)] +use crate::CODEX_ARG0_EXEC_HELPER_ARG1; +use crate::ExecParams; +use crate::ExecServerRuntimePaths; +#[cfg(any(unix, windows))] +use crate::FileSystemSandboxContext; +use crate::ProcessId; +use crate::process_telemetry::ProcessTelemetry; + +async fn prepare_exec_request( + params: &ExecParams, + env: HashMap, + runtime_paths: Option<&ExecServerRuntimePaths>, + network_policy_decider: Option>, + network_policy_audit_observer: Option, +) -> Result { + prepare_exec_request_with_telemetry( + params, + env, + runtime_paths, + network_policy_decider, + network_policy_audit_observer, + &ProcessTelemetry::default(), + ) + .await +} + +#[cfg(unix)] +#[tokio::test] +async fn sandbox_request_wraps_native_argv_on_executor() { + let cwd: AbsolutePathBuf = std::env::current_dir() + .expect("current directory") + .try_into() + .expect("absolute cwd"); + let cwd_uri = PathUri::from_abs_path(&cwd); + let self_exe = std::env::current_exe().expect("current executable"); + let runtime_paths = + ExecServerRuntimePaths::new(self_exe.clone(), Some(self_exe)).expect("runtime paths"); + let sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::workspace_write(), + cwd_uri.clone(), + ); + let params = ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("process-1"), + argv: vec![ + "/bin/bash".to_string(), + "-lc".to_string(), + "pwd".to_string(), + ], + cwd: cwd_uri, + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: Some(sandbox), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }; + + let prepared = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + ) + .await + .expect("prepare sandboxed request"); + + assert_ne!(prepared.command, params.argv); + assert_eq!(prepared.cwd, cwd); + #[cfg(target_os = "linux")] + { + assert_eq!( + prepared.command.first(), + Some(&runtime_paths.codex_self_exe.to_string_lossy().into_owned()) + ); + let permission_profile_json = prepared + .command + .iter() + .position(|arg| arg == "--permission-profile") + .and_then(|index| prepared.command.get(index + 1)) + .expect("sandbox wrapper permission profile"); + let permission_profile: PermissionProfile = + serde_json::from_str(permission_profile_json).expect("permission profile JSON"); + assert_eq!( + permission_profile, + PermissionProfile::workspace_write() + .materialize_project_roots_with_workspace_roots(std::slice::from_ref(&cwd)) + ); + } + #[cfg(target_os = "macos")] + assert_eq!( + prepared.command.first().map(String::as_str), + Some("/usr/bin/sandbox-exec") + ); + + let mut params = params; + params.sandbox.as_mut().unwrap().windows_sandbox_selection = WindowsSandboxSelection::Mxc; + let error = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + ) + .await + .err() + .expect("unsupported MXC must fail closed"); + assert_eq!(error.message, "native MXC is unavailable on this executor"); +} + +#[cfg(unix)] +#[tokio::test] +async fn sandbox_request_routes_custom_arg0_to_inner_helper() { + let cwd: AbsolutePathBuf = std::env::current_dir() + .expect("current directory") + .try_into() + .expect("absolute cwd"); + let cwd_uri = PathUri::from_abs_path(&cwd); + let self_exe = std::env::current_exe().expect("current executable"); + let runtime_paths = + ExecServerRuntimePaths::new(self_exe.clone(), Some(self_exe)).expect("runtime paths"); + let sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::workspace_write(), + cwd_uri.clone(), + ); + let params = ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("process-custom-arg0"), + argv: vec!["/bin/sh".to_string(), "-c".to_string(), "true".to_string()], + cwd: cwd_uri, + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: Some("custom-arg0".to_string()), + sandbox: Some(sandbox), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }; + + let prepared = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + ) + .await + .expect("prepare sandboxed request"); + let helper_mode = prepared + .command + .iter() + .position(|arg| arg == CODEX_ARG0_EXEC_HELPER_ARG1) + .expect("sandboxed command should invoke arg0 helper"); + + assert_eq!( + prepared.command[helper_mode..], + [ + CODEX_ARG0_EXEC_HELPER_ARG1, + "custom-arg0", + "/bin/sh", + "-c", + "true", + ] + ); + #[cfg(target_os = "linux")] + assert_eq!(prepared.arg0, Some(CODEX_LINUX_SANDBOX_ARG0.to_string())); + #[cfg(target_os = "macos")] + assert_eq!(prepared.arg0, None); +} + +#[cfg(target_os = "macos")] +fn managed_network_sandbox_request() -> (ExecParams, ExecServerRuntimePaths) { + let cwd: AbsolutePathBuf = std::env::current_dir() + .expect("current directory") + .try_into() + .expect("absolute cwd"); + let cwd_uri = PathUri::from_abs_path(&cwd); + let self_exe = std::env::current_exe().expect("current executable"); + let runtime_paths = + ExecServerRuntimePaths::new(self_exe.clone(), Some(self_exe)).expect("runtime paths"); + let sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::workspace_write(), + cwd_uri.clone(), + ); + let params = ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("process-managed-network"), + argv: vec!["/usr/bin/true".to_string()], + cwd: cwd_uri, + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: Some(sandbox), + enforce_managed_network: true, + managed_network: None, + network_proxy: None, + }; + (params, runtime_paths) +} + +#[cfg(target_os = "macos")] +fn seatbelt_policy_arg(command: &[String]) -> &str { + command + .windows(2) + .find_map(|args| (args[0] == "-p").then_some(args[1].as_str())) + .expect("Seatbelt policy argument") +} + +#[cfg(target_os = "macos")] +#[tokio::test] +async fn sandbox_request_preserves_prepared_managed_network_policy() { + let (mut params, runtime_paths) = managed_network_sandbox_request(); + let socket_dir = tempfile::tempdir().expect("temporary socket directory"); + let allowed_socket = socket_dir + .path() + .canonicalize() + .expect("canonical socket directory") + .join("allowed.sock") + .to_string_lossy() + .into_owned(); + params.managed_network = Some(ManagedNetworkSandboxContext { + loopback_ports: vec![43123], + allow_local_binding: false, + allow_unix_sockets: vec![allowed_socket.clone()], + dangerously_allow_all_unix_sockets: false, + }); + + let prepared = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + ) + .await + .expect("prepare managed-network sandbox request"); + let policy = seatbelt_policy_arg(&prepared.command); + let unix_socket_definitions = prepared + .command + .iter() + .filter(|arg| arg.starts_with("-DUNIX_SOCKET_PATH_")) + .cloned() + .collect::>(); + + assert!(policy.contains("(allow network-outbound (remote ip \"localhost:43123\"))")); + assert!(policy.contains("(allow system-socket (socket-domain AF_UNIX))")); + assert!(policy.contains( + "(allow network-outbound (remote unix-socket (subpath (param \"UNIX_SOCKET_PATH_0\"))))" + )); + assert_eq!( + unix_socket_definitions, + vec![format!("-DUNIX_SOCKET_PATH_0={allowed_socket}")] + ); + assert!(!policy.contains("(allow network-outbound (remote unix-socket))")); + assert!(!policy.contains("(allow network-outbound)\n")); +} + +#[cfg(target_os = "macos")] +#[tokio::test] +async fn sandbox_request_only_allows_all_unix_sockets_when_configured() { + let (mut params, runtime_paths) = managed_network_sandbox_request(); + params.managed_network = Some(ManagedNetworkSandboxContext { + loopback_ports: vec![43123], + ..ManagedNetworkSandboxContext::default() + }); + + for allow_all in [false, true] { + params + .managed_network + .as_mut() + .expect("managed network context") + .dangerously_allow_all_unix_sockets = allow_all; + let prepared = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + ) + .await + .expect("prepare managed-network sandbox request"); + let policy = seatbelt_policy_arg(&prepared.command); + + assert_eq!( + ( + policy.contains("(allow system-socket (socket-domain AF_UNIX))"), + policy.contains("(allow network-bind (local unix-socket))"), + policy.contains("(allow network-outbound (remote unix-socket))"), + ), + (allow_all, allow_all, allow_all) + ); + assert!(!policy.contains("(allow network-outbound)\n")); + assert!( + !prepared + .command + .iter() + .any(|arg| arg.starts_with("-DUNIX_SOCKET_PATH_")) + ); + } +} + +#[cfg(target_os = "macos")] +#[tokio::test] +async fn sandbox_request_preserves_executor_local_proxy_unix_socket_policy() { + let (mut params, runtime_paths) = managed_network_sandbox_request(); + let socket_dir = tempfile::tempdir().expect("temporary socket directory"); + let socket_root = socket_dir + .path() + .canonicalize() + .expect("canonical socket directory"); + let allowed_socket = socket_root + .join("allowed.sock") + .to_string_lossy() + .into_owned(); + let denied_socket = socket_root + .join("denied.sock") + .to_string_lossy() + .into_owned(); + let config = NetworkProxyConfig { + enabled: true, + enable_socks5: false, + unix_sockets: Some(NetworkUnixSocketPermissions { + entries: [ + (allowed_socket.clone(), NetworkUnixSocketPermission::Allow), + (denied_socket, NetworkUnixSocketPermission::Deny), + ] + .into_iter() + .collect(), + }), + ..NetworkProxyConfig::default() + }; + params.network_proxy = Some(RemoteNetworkProxyLaunchConfig::new( + RemoteNetworkProxyConfig::from_effective_config(&config) + .expect("supported remote proxy config"), + )); + + let prepared = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + ) + .await + .expect("prepare sandbox request with executor-local proxy"); + let policy = seatbelt_policy_arg(&prepared.command); + let proxy_addr: SocketAddr = prepared + .env + .get("HTTP_PROXY") + .expect("HTTP proxy env") + .strip_prefix("http://") + .expect("HTTP proxy scheme") + .parse() + .expect("HTTP proxy address"); + let proxy_port = proxy_addr.port(); + let unix_socket_definitions = prepared + .command + .iter() + .filter(|arg| arg.starts_with("-DUNIX_SOCKET_PATH_")) + .cloned() + .collect::>(); + + assert!(policy.contains(&format!( + "(allow network-outbound (remote ip \"localhost:{proxy_port}\"))" + ))); + assert!(policy.contains( + "(allow network-outbound (remote unix-socket (subpath (param \"UNIX_SOCKET_PATH_0\"))))" + )); + assert_eq!( + unix_socket_definitions, + vec![format!("-DUNIX_SOCKET_PATH_0={allowed_socket}")] + ); + assert!(!policy.contains("(allow network-outbound (remote unix-socket))")); + assert!(!policy.contains("(allow network-outbound)\n")); + + prepared + .network_proxy_handle + .expect("running executor proxy") + .shutdown() + .await + .expect("shut down executor proxy"); +} + +#[tokio::test] +async fn native_request_preserves_native_launch_fields() { + let cwd: AbsolutePathBuf = std::env::current_dir() + .expect("current directory") + .try_into() + .expect("absolute cwd"); + let cwd_uri = PathUri::from_abs_path(&cwd); + let env = HashMap::from([("TEST_ENV".to_string(), "value".to_string())]); + let params = ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("process-1"), + argv: vec!["echo".to_string(), "hello".to_string()], + cwd: cwd_uri, + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: Some("custom-arg0".to_string()), + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }; + + let prepared = prepare_exec_request( + ¶ms, + env.clone(), + /*runtime_paths*/ None, + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + ) + .await + .expect("prepare native request"); + + assert_eq!(prepared.command, params.argv); + assert_eq!(prepared.cwd, cwd); + assert_eq!(prepared.env, env); + assert_eq!(prepared.arg0, params.arg0); +} + +#[tokio::test] +async fn native_request_handles_remote_proxy_config_for_platform() { + let cwd: AbsolutePathBuf = std::env::current_dir() + .expect("current directory") + .try_into() + .expect("absolute cwd"); + let mut config = NetworkProxyConfig { + enabled: true, + ..NetworkProxyConfig::default() + }; + config.set_allowed_domains(vec!["allowed.example".to_string()]); + let proxy_config = RemoteNetworkProxyConfig::from_effective_config(&config) + .expect("supported remote proxy config"); + let params = ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("process-remote-proxy"), + argv: vec!["echo".to_string(), "hello".to_string()], + cwd: PathUri::from_abs_path(&cwd), + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: Some( + RemoteNetworkProxyLaunchConfig::new(proxy_config) + .for_execution("remote".to_string(), "execution-1".to_string()), + ), + }; + let stale_proxy = "http://127.0.0.1:9".to_string(); + let env = HashMap::from([ + ("HTTP_PROXY".to_string(), stale_proxy.clone()), + ("TEST_ENV".to_string(), "value".to_string()), + ( + PROXY_ATTRIBUTION_TOKEN_ENV_KEY.to_string(), + "foreign-token".to_string(), + ), + ]); + + let prepared = prepare_exec_request( + ¶ms, env, /*runtime_paths*/ None, /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + ) + .await + .expect("prepare request with executor-local proxy"); + + if cfg!(target_os = "windows") { + assert_eq!(prepared.env.get("HTTP_PROXY"), None); + assert_eq!(prepared.env.get("TEST_ENV"), Some(&"value".to_string())); + assert!(prepared.network_proxy_handle.is_none()); + return; + } + + let http_proxy = prepared.env.get("HTTP_PROXY").expect("HTTP proxy env"); + assert_ne!(http_proxy, &stale_proxy); + assert!(http_proxy.starts_with("http://127.0.0.1:")); + assert!(!prepared.env.contains_key(PROXY_ATTRIBUTION_TOKEN_ENV_KEY)); + let proxy_addr: SocketAddr = http_proxy + .strip_prefix("http://") + .expect("HTTP proxy scheme") + .parse() + .expect("HTTP proxy address"); + let mut stream = tokio::net::TcpStream::connect(proxy_addr) + .await + .expect("connect to executor proxy"); + stream + .write_all(b"CONNECT blocked.example:443 HTTP/1.1\r\nHost: blocked.example:443\r\n\r\n") + .await + .expect("write CONNECT request"); + let mut response = [0_u8; 256]; + let response_len = timeout(Duration::from_secs(2), stream.read(&mut response)) + .await + .expect("proxy response timeout") + .expect("read proxy response"); + assert!(String::from_utf8_lossy(&response[..response_len]).starts_with("HTTP/1.1 403")); + + prepared + .network_proxy_handle + .expect("running executor proxy") + .shutdown() + .await + .expect("shut down executor proxy"); +} + +#[cfg(not(target_os = "windows"))] +#[tokio::test] +async fn disabled_remote_proxy_config_is_rejected_before_exporting_ports() { + let cwd: AbsolutePathBuf = std::env::current_dir() + .expect("current directory") + .try_into() + .expect("absolute cwd"); + let proxy_config = + RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig::default()) + .expect("serializable disabled proxy config"); + let params = ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("process-disabled-remote-proxy"), + argv: vec!["echo".to_string(), "hello".to_string()], + cwd: PathUri::from_abs_path(&cwd), + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: Some(RemoteNetworkProxyLaunchConfig::new(proxy_config)), + }; + + let error = prepare_exec_request( + ¶ms, + HashMap::new(), + /*runtime_paths*/ None, + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + ) + .await + .err() + .expect("disabled executor proxy launch must fail closed"); + + assert_eq!(error.code, -32602); + assert!( + error + .message + .contains("executor-local network proxy launch requires an enabled proxy") + ); +} + +#[cfg(windows)] +#[test_case(WindowsSandboxLevel::RestrictedToken ; "unelevated is rejected")] +#[test_case(WindowsSandboxLevel::Elevated ; "elevated is accepted")] +#[tokio::test] +async fn managed_network_honors_windows_sandbox_level(windows_sandbox_level: WindowsSandboxLevel) { + let cwd: AbsolutePathBuf = std::env::current_dir() + .expect("current directory") + .try_into() + .expect("absolute cwd"); + let cwd_uri = PathUri::from_abs_path(&cwd); + let self_exe = std::env::current_exe().expect("current executable"); + let runtime_paths = ExecServerRuntimePaths::new(self_exe, None).expect("runtime paths"); + let permissions = PermissionProfile::read_only(); + let mut sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + permissions.clone(), + cwd_uri.clone(), + ); + sandbox.windows_sandbox_selection = windows_sandbox_level.into(); + sandbox.windows_sandbox_proxy_settings_mode = + Some(codex_sandboxing::WindowsSandboxProxySettingsMode::Preserve); + let proxy_config = RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig { + enabled: true, + enable_socks5: false, + ..NetworkProxyConfig::default() + }) + .expect("supported remote proxy config"); + let params = ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("process-managed-network"), + argv: vec!["cmd.exe".to_string(), "/c".to_string(), "exit".to_string()], + cwd: cwd_uri, + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: Some(sandbox), + enforce_managed_network: true, + managed_network: None, + network_proxy: Some(RemoteNetworkProxyLaunchConfig::new(proxy_config)), + }; + + let prepared = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + ) + .await; + + if windows_sandbox_level == WindowsSandboxLevel::RestrictedToken { + let error = prepared + .err() + .expect("managed networking must reject an unelevated Windows sandbox"); + assert_eq!(error.code, -32602); + assert!( + error + .message + .contains("managed networking requires the elevated Windows sandbox backend") + ); + return; + } + + let mut prepared = prepared.expect("managed networking accepts an elevated Windows sandbox"); + let spawn = prepared + .windows_sandbox_spawn_request() + .expect("Windows sandbox spawn request"); + assert_eq!(spawn.windows_sandbox_level, WindowsSandboxLevel::Elevated); + assert!(spawn.proxy_enforced); + assert!(spawn.network_proxy_restricting_sid.is_some()); + prepared + .network_proxy_handle + .take() + .expect("running executor proxy") + .shutdown() + .await + .expect("shut down executor proxy"); +} diff --git a/codex-rs/exec-server/src/process_telemetry.rs b/codex-rs/exec-server/src/process_telemetry.rs new file mode 100644 index 0000000000000000000000000000000000000000..f926fa9889c9600dea5316a0e9e541b22a2781f5 --- /dev/null +++ b/codex-rs/exec-server/src/process_telemetry.rs @@ -0,0 +1,78 @@ +//! Emits bounded process lifecycle telemetry using identity captured at launch. +//! Reconnects and later process operations must not replace that identity. + +use std::sync::Arc; + +use codex_sandboxing::SandboxType; +use opentelemetry::trace::SpanContext; + +use crate::telemetry::ExecutorRegistration; + +/// Log fields captured at launch, never refreshed from a resumed session. +#[derive(Clone, Default)] +pub(crate) struct ProcessTelemetry { + pub(crate) launch_context: Option, + pub(crate) thread_id: Option, + pub(crate) tool_call_id: Option, + pub(crate) executor_registration: Option>, +} + +/// Lifecycle events with outcome fields only when a process has exited. +pub(crate) enum ProcessTelemetryEvent { + Start, + SpawnFailed, + SandboxDenied, + Exit { + exit_code: i32, + termination_requested: bool, + }, +} + +impl ProcessTelemetry { + pub(crate) fn log(&self, event: ProcessTelemetryEvent, sandbox: SandboxType) { + let (event_name, exit_code, termination_requested, reason) = match event { + ProcessTelemetryEvent::Start => ("codex.exec_server.process_start", None, None, None), + ProcessTelemetryEvent::SpawnFailed => { + ("codex.exec_server.process_spawn_failed", None, None, None) + } + ProcessTelemetryEvent::SandboxDenied => ( + "codex.exec_server.sandbox_denied", + None, + None, + Some("inferred_denial"), + ), + ProcessTelemetryEvent::Exit { + exit_code, + termination_requested, + } => ( + "codex.exec_server.process_exit", + Some(exit_code), + Some(termination_requested), + None, + ), + }; + let trace_id = self + .launch_context + .as_ref() + .map(|span| span.trace_id().to_string()); + let span_id = self + .launch_context + .as_ref() + .map(|span| span.span_id().to_string()); + tracing::event!( + target: "codex_otel.log_only", + tracing::Level::INFO, + event.name = event_name, + launch.trace_id = trace_id.as_deref(), + launch.span_id = span_id.as_deref(), + conversation.id = self.thread_id.as_deref(), + tool.call_id = self.tool_call_id.as_deref(), + executor.environment_id = self.executor_registration.as_ref().map(|registration| registration.environment_id.as_str()), + executor.registration_id = self.executor_registration.as_ref().map(|registration| registration.executor_registration_id.as_str()), + sandbox.type = ?sandbox, + process.exit_code = exit_code, + process.termination_requested = termination_requested, + reason, + ); + } +} diff --git a/codex-rs/exec-server/src/proto/codex.exec_server.relay.v1.proto b/codex-rs/exec-server/src/proto/codex.exec_server.relay.v1.proto new file mode 100644 index 0000000000000000000000000000000000000000..a9d2100137230b80e014f7e99cb0d8a8093d2473 --- /dev/null +++ b/codex-rs/exec-server/src/proto/codex.exec_server.relay.v1.proto @@ -0,0 +1,45 @@ +syntax = "proto3"; + +package codex.exec_server.relay.v1; + +message RelayMessageFrame { + uint32 version = 1; + string stream_id = 2; + uint32 ack = 3; + uint32 ack_bits = 4; + + oneof body { + RelayData data = 5; + RelayAck ack_frame = 6; + RelayResume resume = 7; + RelayReset reset = 8; + RelayHeartbeat heartbeat = 9; + RelayHandshake handshake = 10; + } + + optional string traceparent = 11; + optional string tracestate = 12; +} + +message RelayData { + uint32 seq = 1; + uint32 segment_index = 2; + uint32 segment_count = 3; + bytes payload = 4; +} + +message RelayAck {} + +message RelayResume { + uint32 next_seq = 1; +} + +message RelayReset { + string reason = 1; +} + +message RelayHeartbeat {} + +message RelayHandshake { + bytes payload = 1; +} diff --git a/codex-rs/exec-server/src/proto/codex.exec_server.relay.v1.rs b/codex-rs/exec-server/src/proto/codex.exec_server.relay.v1.rs new file mode 100644 index 0000000000000000000000000000000000000000..4694ed4024da99232dbae8d626fd551e9b6ba25f --- /dev/null +++ b/codex-rs/exec-server/src/proto/codex.exec_server.relay.v1.rs @@ -0,0 +1,65 @@ +// This file is @generated by prost-build. +#[derive(Clone, PartialEq, ::prost::Message)] +pub struct RelayMessageFrame { + #[prost(uint32, tag = "1")] + pub version: u32, + #[prost(string, tag = "2")] + pub stream_id: ::prost::alloc::string::String, + #[prost(uint32, tag = "3")] + pub ack: u32, + #[prost(uint32, tag = "4")] + pub ack_bits: u32, + #[prost(oneof = "relay_message_frame::Body", tags = "5, 6, 7, 8, 9, 10")] + pub body: ::core::option::Option, + #[prost(string, optional, tag = "11")] + pub traceparent: ::core::option::Option<::prost::alloc::string::String>, + #[prost(string, optional, tag = "12")] + pub tracestate: ::core::option::Option<::prost::alloc::string::String>, +} +pub mod relay_message_frame { + #[derive(Clone, PartialEq, ::prost::Oneof)] + pub enum Body { + #[prost(message, tag = "5")] + Data(super::RelayData), + #[prost(message, tag = "6")] + AckFrame(super::RelayAck), + #[prost(message, tag = "7")] + Resume(super::RelayResume), + #[prost(message, tag = "8")] + Reset(super::RelayReset), + #[prost(message, tag = "9")] + Heartbeat(super::RelayHeartbeat), + #[prost(message, tag = "10")] + Handshake(super::RelayHandshake), + } +} +#[derive(Clone, PartialEq, ::prost::Message)] +pub struct RelayData { + #[prost(uint32, tag = "1")] + pub seq: u32, + #[prost(uint32, tag = "2")] + pub segment_index: u32, + #[prost(uint32, tag = "3")] + pub segment_count: u32, + #[prost(bytes = "vec", tag = "4")] + pub payload: ::prost::alloc::vec::Vec, +} +#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] +pub struct RelayAck {} +#[derive(Clone, Copy, PartialEq, Eq, Hash, ::prost::Message)] +pub struct RelayResume { + #[prost(uint32, tag = "1")] + pub next_seq: u32, +} +#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] +pub struct RelayReset { + #[prost(string, tag = "1")] + pub reason: ::prost::alloc::string::String, +} +#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] +pub struct RelayHeartbeat {} +#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] +pub struct RelayHandshake { + #[prost(bytes = "vec", tag = "1")] + pub payload: ::prost::alloc::vec::Vec, +} diff --git a/codex-rs/exec-server/src/regular_file.rs b/codex-rs/exec-server/src/regular_file.rs new file mode 100644 index 0000000000000000000000000000000000000000..4182046a623bec4e3e50c6108255a6e20e3fedb6 --- /dev/null +++ b/codex-rs/exec-server/src/regular_file.rs @@ -0,0 +1,82 @@ +use std::io; +use std::path::Path; +use tokio::io::AsyncReadExt; + +pub(crate) async fn open(path: &Path) -> io::Result { + let mut options = tokio::fs::OpenOptions::new(); + options.read(true); + configure_open(&mut options); + + let file = options.open(path).await?; + if !is_disk_file(&file) || !file.metadata().await?.is_file() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("path `{}` is not a file", path.display()), + )); + } + Ok(file) +} + +/// Reads a regular UTF-8 file without following a symlink at its final path component. +pub async fn read_sensitive_file_to_string(path: &Path) -> io::Result { + let mut options = tokio::fs::OpenOptions::new(); + options.read(true); + configure_open(&mut options); + + #[cfg(unix)] + options.custom_flags(libc::O_NONBLOCK | libc::O_NOFOLLOW); + + #[cfg(windows)] + { + use windows_sys::Win32::Storage::FileSystem::FILE_FLAG_OPEN_REPARSE_POINT; + + options.custom_flags(FILE_FLAG_OPEN_REPARSE_POINT); + } + + let mut file = options.open(path).await?; + let metadata = file.metadata().await?; + if !is_disk_file(&file) || !metadata.is_file() || metadata.file_type().is_symlink() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("path `{}` is not a regular file", path.display()), + )); + } + + let mut contents = String::new(); + file.read_to_string(&mut contents).await?; + Ok(contents) +} + +#[cfg(unix)] +fn configure_open(options: &mut tokio::fs::OpenOptions) { + options.custom_flags(libc::O_NONBLOCK); +} + +#[cfg(windows)] +fn configure_open(options: &mut tokio::fs::OpenOptions) { + use windows_sys::Win32::Storage::FileSystem::SECURITY_IDENTIFICATION; + + options.security_qos_flags(SECURITY_IDENTIFICATION); +} + +#[cfg(not(any(unix, windows)))] +fn configure_open(_options: &mut tokio::fs::OpenOptions) {} + +#[cfg(windows)] +pub(crate) fn is_disk_file(file: &impl std::os::windows::io::AsRawHandle) -> bool { + use windows_sys::Win32::Foundation::HANDLE; + use windows_sys::Win32::Storage::FileSystem::FILE_TYPE_DISK; + use windows_sys::Win32::Storage::FileSystem::GetFileType; + + // SAFETY: `file` owns this handle for the duration of the call. + unsafe { GetFileType(file.as_raw_handle() as HANDLE) == FILE_TYPE_DISK } +} + +#[cfg(not(windows))] +fn is_disk_file(_file: &tokio::fs::File) -> bool { + true +} + +#[cfg(test)] +#[path = "regular_file_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/regular_file_tests.rs b/codex-rs/exec-server/src/regular_file_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4cda98114c246e2ecc1d4d6a964a0752f9e72b14 --- /dev/null +++ b/codex-rs/exec-server/src/regular_file_tests.rs @@ -0,0 +1,47 @@ +use super::read_sensitive_file_to_string; +use tempfile::TempDir; + +#[tokio::test] +async fn read_sensitive_file_reads_regular_file() { + let directory = TempDir::new().expect("temporary directory"); + let path = directory.path().join("role.toml"); + tokio::fs::write(&path, "developer_instructions = 'stay focused'") + .await + .expect("write regular file"); + + assert_eq!( + read_sensitive_file_to_string(&path) + .await + .expect("read regular file"), + "developer_instructions = 'stay focused'", + ); +} + +#[tokio::test] +async fn read_sensitive_file_rejects_directory() { + let directory = TempDir::new().expect("temporary directory"); + + assert!( + read_sensitive_file_to_string(directory.path()) + .await + .is_err() + ); +} + +#[cfg(any(unix, windows))] +#[tokio::test] +async fn read_sensitive_file_rejects_symlink() { + let directory = TempDir::new().expect("temporary directory"); + let target = directory.path().join("target.toml"); + let link = directory.path().join("role.toml"); + tokio::fs::write(&target, "model_provider = 'attacker'") + .await + .expect("write symlink target"); + + #[cfg(unix)] + std::os::unix::fs::symlink(&target, &link).expect("create symlink"); + #[cfg(windows)] + std::os::windows::fs::symlink_file(&target, &link).expect("create symlink"); + + assert!(read_sensitive_file_to_string(&link).await.is_err()); +} diff --git a/codex-rs/exec-server/src/relay.rs b/codex-rs/exec-server/src/relay.rs new file mode 100644 index 0000000000000000000000000000000000000000..40afbf0e288895451c07f04de89d33c651fb50f8 --- /dev/null +++ b/codex-rs/exec-server/src/relay.rs @@ -0,0 +1,1342 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCMessage; +use codex_protocol::protocol::W3cTraceContext; +use futures::Sink; +use futures::SinkExt; +use futures::Stream; +use futures::StreamExt; +use prost::Message as ProstMessage; +use tokio::sync::mpsc; +use tokio::sync::watch; +use tokio::task::JoinSet; +use tokio::time::timeout; +use tokio_tungstenite::tungstenite::Message; +use tracing::debug; +use tracing::info; +use tracing::warn; +use uuid::Uuid; + +use crate::ExecServerError; +use crate::connection::CHANNEL_CAPACITY; +use crate::connection::JsonRpcConnection; +use crate::connection::JsonRpcConnectionEvent; +use crate::connection::JsonRpcTransport; +use crate::connection::WEBSOCKET_KEEPALIVE_INTERVAL; +use crate::noise_channel::NoiseChannelIdentity; +use crate::noise_channel::NoiseChannelPublicKey; +use crate::noise_channel::PendingResponderHandshake; +use crate::noise_channel::noise_channel_prologue; +use crate::noise_relay::NOISE_RELAY_RESET_REASON; +use crate::noise_relay::executor_stream::ClosedNoiseVirtualStream; +use crate::noise_relay::executor_stream::NoiseVirtualStream; +use crate::noise_relay::executor_stream::spawn_noise_virtual_stream; +use crate::noise_relay::stream_handler::NoiseStreamHandler; +use crate::relay_proto::RelayData; +use crate::relay_proto::RelayHandshake; +use crate::relay_proto::RelayMessageFrame; +use crate::relay_proto::RelayReset; +use crate::relay_proto::RelayResume; +use crate::relay_proto::relay_message_frame; +#[cfg(test)] +use crate::server::ConnectionProcessor; +use crate::telemetry::ExecutorRegistration; +use crate::websocket_pong_watchdog::WEBSOCKET_PONG_TIMEOUT; +use crate::websocket_pong_watchdog::WEBSOCKET_PONG_TIMEOUT_REASON; +use crate::websocket_pong_watchdog::WebSocketPongWatchdog; + +const RELAY_MESSAGE_FRAME_VERSION: u32 = 1; +const MAX_ACTIVE_NOISE_RELAY_STREAMS: usize = 128; +const MAX_FAILED_NOISE_HANDSHAKES: usize = 8; +const MAX_HARNESS_KEY_AUTHORIZATION_BYTES: usize = 4096; +const MAX_PENDING_HANDSHAKE_VALIDATIONS: usize = 32; +const HARNESS_KEY_VALIDATION_TIMEOUT: Duration = Duration::from_secs(10); + +#[derive(Debug, Clone, Copy, Eq, PartialEq)] +pub(crate) enum RendezvousDisconnectReason { + PeerClose, + ReadError, + WriteError, + PongTimeout, + LocalShutdown, +} + +impl RendezvousDisconnectReason { + pub(crate) fn as_str(self) -> &'static str { + match self { + Self::PeerClose => "peer_close", + Self::ReadError => "read_error", + Self::WriteError => "write_error", + Self::PongTimeout => WEBSOCKET_PONG_TIMEOUT_REASON, + Self::LocalShutdown => "local_shutdown", + } + } +} + +#[derive(Debug, Clone, Copy, Eq, PartialEq)] +pub(crate) enum RelayFrameBodyKind { + Data, + Ack, + Resume, + Reset, + Heartbeat, + Handshake, +} + +impl RelayMessageFrame { + pub(crate) fn data( + stream_id: String, + seq: u32, + payload: Vec, + trace: Option, + ) -> Self { + let (traceparent, tracestate) = trace + .map(|trace| (trace.traceparent, trace.tracestate)) + .unwrap_or_default(); + Self { + version: RELAY_MESSAGE_FRAME_VERSION, + stream_id, + traceparent, + tracestate, + body: Some(relay_message_frame::Body::Data(RelayData { + seq, + segment_index: 0, + segment_count: 1, + payload, + })), + ..Self::default() + } + } + + pub(crate) fn resume(stream_id: String) -> Self { + Self { + version: RELAY_MESSAGE_FRAME_VERSION, + stream_id, + body: Some(relay_message_frame::Body::Resume(RelayResume { + next_seq: 0, + })), + ..Self::default() + } + } + + pub(crate) fn handshake(stream_id: String, payload: Vec) -> Self { + Self { + version: RELAY_MESSAGE_FRAME_VERSION, + stream_id, + body: Some(relay_message_frame::Body::Handshake(RelayHandshake { + payload, + })), + ..Self::default() + } + } + + pub(crate) fn reset(stream_id: String, reason: String) -> Self { + Self { + version: RELAY_MESSAGE_FRAME_VERSION, + stream_id, + body: Some(relay_message_frame::Body::Reset(RelayReset { reason })), + ..Self::default() + } + } + + pub(crate) fn validate(&self) -> Result { + if self.version != RELAY_MESSAGE_FRAME_VERSION { + return Err(ExecServerError::Protocol(format!( + "unsupported relay message frame version {}", + self.version + ))); + } + if self.stream_id.trim().is_empty() { + return Err(ExecServerError::Protocol( + "relay message frame is missing stream_id".to_string(), + )); + } + match self.body.as_ref() { + Some(relay_message_frame::Body::Data(data)) => { + if data.segment_index != 0 || data.segment_count != 1 || data.payload.is_empty() { + return Err(ExecServerError::Protocol( + "relay data message frame is missing required fields".to_string(), + )); + } + Ok(RelayFrameBodyKind::Data) + } + Some(relay_message_frame::Body::AckFrame(_)) => Ok(RelayFrameBodyKind::Ack), + Some(relay_message_frame::Body::Resume(_)) => Ok(RelayFrameBodyKind::Resume), + Some(relay_message_frame::Body::Reset(reset)) => { + if reset.reason.is_empty() { + return Err(ExecServerError::Protocol( + "relay reset message frame is missing reason".to_string(), + )); + } + Ok(RelayFrameBodyKind::Reset) + } + Some(relay_message_frame::Body::Heartbeat(_)) => Ok(RelayFrameBodyKind::Heartbeat), + Some(relay_message_frame::Body::Handshake(handshake)) => { + if handshake.payload.is_empty() { + return Err(ExecServerError::Protocol( + "relay handshake message frame is missing payload".to_string(), + )); + } + Ok(RelayFrameBodyKind::Handshake) + } + None => Err(ExecServerError::Protocol( + "relay message frame is missing body".to_string(), + )), + } + } + + pub(crate) fn into_data(self) -> Result { + let kind = self.validate()?; + if kind != RelayFrameBodyKind::Data { + return Err(ExecServerError::Protocol( + "expected relay data message frame".to_string(), + )); + } + match self.body { + Some(relay_message_frame::Body::Data(data)) => Ok(data), + _ => Err(ExecServerError::Protocol( + "expected relay data message frame".to_string(), + )), + } + } + + fn into_jsonrpc_message(self) -> Result { + let payload = self.into_data()?.payload; + serde_json::from_slice(&payload).map_err(ExecServerError::Json) + } + + pub(crate) fn into_handshake_payload(self) -> Result, ExecServerError> { + let kind = self.validate()?; + if kind != RelayFrameBodyKind::Handshake { + return Err(ExecServerError::Protocol( + "expected relay handshake message frame".to_string(), + )); + } + match self.body { + Some(relay_message_frame::Body::Handshake(handshake)) => Ok(handshake.payload), + _ => Err(ExecServerError::Protocol( + "expected relay handshake message frame".to_string(), + )), + } + } + + pub(crate) fn into_reset_reason(self) -> Option { + match self.body { + Some(relay_message_frame::Body::Reset(reset)) if !reset.reason.is_empty() => { + Some(reset.reason) + } + _ => None, + } + } +} + +pub(crate) fn encode_relay_message_frame(frame: &RelayMessageFrame) -> Vec { + frame.encode_to_vec() +} + +pub(crate) fn decode_relay_message_frame( + payload: &[u8], +) -> Result { + RelayMessageFrame::decode(payload) + .map_err(|err| ExecServerError::Protocol(format!("invalid relay message frame: {err}"))) +} + +pub(crate) fn jsonrpc_payload(message: &JSONRPCMessage) -> Result, ExecServerError> { + serde_json::to_vec(message).map_err(ExecServerError::Json) +} + +enum RelayEventSendError { + IncomingClosed, + WebSocketClosed, +} + +async fn send_event_with_keepalive( + websocket: &mut T, + keepalive: &mut tokio::time::Interval, + incoming_tx: &mpsc::Sender, + event: JsonRpcConnectionEvent, +) -> Result<(), RelayEventSendError> +where + T: Sink + Unpin, +{ + let send = incoming_tx.send(event); + tokio::pin!(send); + loop { + tokio::select! { + result = &mut send => { + return result.map_err(|_| RelayEventSendError::IncomingClosed); + } + _ = keepalive.tick() => { + websocket + .send(Message::Ping(Vec::new().into())) + .await + .map_err(|_| RelayEventSendError::WebSocketClosed)?; + } + } + } +} + +pub(crate) fn harness_connection_from_websocket( + stream: T, + connection_label: String, +) -> JsonRpcConnection +where + T: Sink + Stream> + Unpin + Send + 'static, + E: std::fmt::Display + Send + 'static, +{ + let stream_id = Uuid::new_v4().to_string(); + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (incoming_tx, incoming_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (disconnected_tx, disconnected_rx) = watch::channel(false); + + let websocket_task = tokio::spawn(async move { + let mut websocket = stream; + let reader_label = connection_label; + let reader_stream_id = stream_id.clone(); + let resume = RelayMessageFrame::resume(stream_id.clone()); + if websocket + .send(Message::Binary(encode_relay_message_frame(&resume).into())) + .await + .is_err() + { + let _ = disconnected_tx.send(true); + return; + } + + let mut keepalive = tokio::time::interval_at( + tokio::time::Instant::now() + WEBSOCKET_KEEPALIVE_INTERVAL, + WEBSOCKET_KEEPALIVE_INTERVAL, + ); + keepalive.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + let mut next_seq = 0u32; + loop { + tokio::select! { + maybe_message = outgoing_rx.recv() => { + let Some(message) = maybe_message else { + break; + }; + let payload = match jsonrpc_payload(&message) { + Ok(payload) => payload, + Err(err) => { + warn!("failed to serialize JSON-RPC payload for relay transport: {err}"); + break; + } + }; + let trace = match message { + JSONRPCMessage::Request(request) => request.trace, + JSONRPCMessage::Notification(_) + | JSONRPCMessage::Response(_) + | JSONRPCMessage::Error(_) => None, + }; + let frame = RelayMessageFrame::data(stream_id.clone(), next_seq, payload, trace); + next_seq = next_seq.wrapping_add(1); + if websocket + .send(Message::Binary(encode_relay_message_frame(&frame).into())) + .await + .is_err() + { + let _ = disconnected_tx.send(true); + break; + } + } + _ = keepalive.tick() => { + if websocket.send(Message::Ping(Vec::new().into())).await.is_err() { + let _ = disconnected_tx.send(true); + break; + } + } + incoming_message = websocket.next() => { + match incoming_message { + Some(Ok(Message::Binary(payload))) => { + let frame = match decode_relay_message_frame(payload.as_ref()) { + Ok(frame) => frame, + Err(err) => { + let _ = incoming_tx + .send(JsonRpcConnectionEvent::MalformedMessage { + reason: format!( + "failed to parse relay message frame from {reader_label}: {err}" + ), + }) + .await; + continue; + } + }; + if frame.stream_id != reader_stream_id { + continue; + } + let kind = match frame.validate() { + Ok(kind) => kind, + Err(err) => { + let _ = incoming_tx + .send(JsonRpcConnectionEvent::MalformedMessage { + reason: err.to_string(), + }) + .await; + continue; + } + }; + match kind { + RelayFrameBodyKind::Data => match frame.into_jsonrpc_message() { + Ok(message) => { + match send_event_with_keepalive( + &mut websocket, + &mut keepalive, + &incoming_tx, + JsonRpcConnectionEvent::message(message), + ) + .await + { + Ok(()) => {} + Err(RelayEventSendError::IncomingClosed) => break, + Err(RelayEventSendError::WebSocketClosed) => { + let _ = disconnected_tx.send(true); + break; + } + } + } + Err(err) => { + let _ = incoming_tx + .send(JsonRpcConnectionEvent::MalformedMessage { + reason: err.to_string(), + }) + .await; + } + }, + RelayFrameBodyKind::Reset => { + let _ = disconnected_tx.send(true); + let _ = incoming_tx + .send(JsonRpcConnectionEvent::Disconnected { + reason: frame.into_reset_reason(), + }) + .await; + break; + } + RelayFrameBodyKind::Ack + | RelayFrameBodyKind::Resume + | RelayFrameBodyKind::Heartbeat + | RelayFrameBodyKind::Handshake => {} + } + } + Some(Ok(Message::Close(_))) | None => { + let _ = disconnected_tx.send(true); + let _ = incoming_tx + .send(JsonRpcConnectionEvent::Disconnected { reason: None }) + .await; + break; + } + Some(Ok(Message::Ping(_) | Message::Pong(_) | Message::Frame(_))) => {} + Some(Ok(Message::Text(_))) => { + let _ = incoming_tx + .send(JsonRpcConnectionEvent::MalformedMessage { + reason: "relay exec-server transport expects binary protobuf frames" + .to_string(), + }) + .await; + } + Some(Err(err)) => { + let _ = disconnected_tx.send(true); + let _ = incoming_tx + .send(JsonRpcConnectionEvent::Disconnected { + reason: Some(format!( + "failed to read relay websocket frame from {reader_label}: {err}" + )), + }) + .await; + break; + } + } + } + } + } + }); + + JsonRpcConnection { + outgoing_tx, + incoming_rx, + disconnected_rx, + task_handles: vec![websocket_task], + transport: JsonRpcTransport::Plain, + } +} + +/// Validates that a Noise-authenticated harness public key is authorized. +/// +/// Implementations must consult an authority independent of rendezvous. The +/// exec-server invokes this after parsing the first IK message and before +/// completing the responder handshake. +pub(crate) trait HarnessKeyValidator: Send + Sync { + fn validate_harness_key( + &self, + harness_public_key: &NoiseChannelPublicKey, + authorization: &str, + ) -> impl std::future::Future> + Send; +} + +/// Serve authenticated virtual JSON-RPC streams over one executor websocket. +/// +/// Parsing the first Noise message authenticates the harness key. Only a +/// successful registry check turns that pending handshake into a virtual stream. +#[tracing::instrument(level = "debug", skip_all, fields(noise_side = "executor"))] +pub(crate) async fn run_multiplexed_environment( + stream: T, + handler: H, + environment_id: String, + executor_registration_id: String, + identity: NoiseChannelIdentity, + validator: V, +) -> RendezvousDisconnectReason +where + T: Sink + Stream> + Unpin + Send + 'static, + E: std::fmt::Display + Send + 'static, + V: HarnessKeyValidator + Clone + 'static, + H: NoiseStreamHandler, +{ + debug!( + environment_id, + executor_registration_id, "Noise executor relay details" + ); + let executor_registration = + ExecutorRegistration::new(environment_id.clone(), executor_registration_id.clone()) + .map(Arc::new); + let (mut websocket_sink, mut websocket_stream) = stream.split(); + let (physical_outgoing_tx, mut physical_outgoing_rx) = + mpsc::channel::>(CHANNEL_CAPACITY); + let (closed_stream_tx, mut closed_stream_rx) = + mpsc::channel::(MAX_ACTIVE_NOISE_RELAY_STREAMS); + let (pong_tx, mut pong_rx) = mpsc::channel(1); + // Use a separate writer so this loop never waits on the channel it drains. + let mut physical_writer_task = tokio::spawn(async move { + let mut keepalive = tokio::time::interval_at( + tokio::time::Instant::now() + WEBSOCKET_KEEPALIVE_INTERVAL, + WEBSOCKET_KEEPALIVE_INTERVAL, + ); + keepalive.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + let mut pong_watchdog = WebSocketPongWatchdog::new(WEBSOCKET_PONG_TIMEOUT); + let pong_deadline = tokio::time::sleep(WEBSOCKET_PONG_TIMEOUT); + tokio::pin!(pong_deadline); + loop { + let message = tokio::select! { + pong = pong_rx.recv() => { + let Some(()) = pong else { + break RendezvousDisconnectReason::LocalShutdown; + }; + pong_watchdog.received_pong(); + continue; + } + _ = &mut pong_deadline, if pong_watchdog.deadline().is_some() => { + match pong_rx.try_recv() { + Ok(()) => { + pong_watchdog.received_pong(); + continue; + } + Err(tokio::sync::mpsc::error::TryRecvError::Empty) => { + break RendezvousDisconnectReason::PongTimeout; + } + Err(tokio::sync::mpsc::error::TryRecvError::Disconnected) => { + break RendezvousDisconnectReason::LocalShutdown; + } + } + } + _ = keepalive.tick(), if pong_watchdog.deadline().is_none() => { + Message::Ping(Vec::new().into()) + } + encoded = physical_outgoing_rx.recv() => { + let Some(encoded) = encoded else { + break RendezvousDisconnectReason::LocalShutdown; + }; + Message::Binary(encoded.into()) + } + }; + let is_keepalive_ping = matches!(message, Message::Ping(_)); + let write_deadline = pong_watchdog.write_deadline(tokio::time::Instant::now()); + match tokio::time::timeout_at(write_deadline, websocket_sink.send(message)).await { + Ok(Ok(())) => { + if is_keepalive_ping { + pong_watchdog.ping_sent(tokio::time::Instant::now()); + if let Some(deadline) = pong_watchdog.deadline() { + pong_deadline.as_mut().reset(deadline); + } + } + } + Ok(Err(error)) => { + warn!("Noise multiplexed environment websocket write failed: {error}"); + break RendezvousDisconnectReason::WriteError; + } + Err(_) => { + warn!("Noise multiplexed environment websocket write timed out"); + break RendezvousDisconnectReason::WriteError; + } + } + } + }); + let mut streams: HashMap> = HashMap::new(); + let mut pending_handshakes: HashMap = HashMap::new(); + let mut validation_tasks: JoinSet = JoinSet::new(); + let mut failed_handshakes = 0usize; + let mut next_validation_id = 0u64; + let mut disconnect_reason = RendezvousDisconnectReason::LocalShutdown; + + loop { + // Registry calls run separately so a slow check does not block the relay. + let frame = tokio::select! { + writer_result = &mut physical_writer_task => { + match writer_result { + Ok(reason) => disconnect_reason = reason, + Err(error) => { + warn!("Noise multiplexed environment websocket writer failed: {error}"); + disconnect_reason = RendezvousDisconnectReason::LocalShutdown; + } + } + break; + } + Some(closed_stream) = closed_stream_rx.recv() => { + // A stream ID may have been reused before this writer exits. + // Remove only the instance that sent the notification. + let is_current = streams + .get(&closed_stream.stream_id) + .is_some_and(|stream| stream.instance_id == closed_stream.instance_id); + if is_current { + streams.remove(&closed_stream.stream_id); + send_reset(&physical_outgoing_tx, closed_stream.stream_id); + } + continue; + } + validation_result = validation_tasks.join_next(), if !validation_tasks.is_empty() => { + match validation_result { + Some(Ok(validation_result)) => { + // The stream ID may have been reused while validation ran. + let is_current = pending_handshakes + .get(&validation_result.stream_id) + .is_some_and(|pending| { + pending.validation_id == validation_result.validation_id + }); + if !is_current { + continue; + } + let Some(pending) = + pending_handshakes.remove(&validation_result.stream_id) + else { + continue; + }; + if validation_result.result.is_err() { + // Validator errors may contain authorization details. + warn!( + noise_event = "authorization", + noise_outcome = "error", + noise_reason = "authorization_failed", + "Noise harness authorization failed" + ); + debug!( + stream_id = validation_result.stream_id, + "Noise harness authorization failure details" + ); + send_reset(&physical_outgoing_tx, validation_result.stream_id); + if failed_handshake_budget_exhausted(&mut failed_handshakes) { + warn!("closing Noise relay after repeated handshake failures"); + break; + } + continue; + } + if streams.len() >= MAX_ACTIVE_NOISE_RELAY_STREAMS { + warn!("Noise relay has too many active streams"); + send_reset(&physical_outgoing_tx, validation_result.stream_id); + continue; + } + + // This is the only point where the responder completes + // IK and exposes a JSON-RPC stream: Noise authenticated + // the harness key and the registry authorized it. + let (transport, response) = match pending.handshake.complete() { + Ok(completed) => completed, + Err(error) => { + warn!("failed to complete Noise relay handshake: {error}"); + send_reset(&physical_outgoing_tx, validation_result.stream_id); + if failed_handshake_budget_exhausted(&mut failed_handshakes) { + warn!("closing Noise relay after repeated handshake failures"); + break; + } + continue; + } + }; + let response = RelayMessageFrame::handshake( + validation_result.stream_id.clone(), + response, + ); + // Do not leave a half-open stream if the handshake reply + // cannot be queued immediately. + if physical_outgoing_tx + .try_send(encode_relay_message_frame(&response)) + .is_err() + { + break; + } + info!( + noise_event = "handshake", + noise_outcome = "ok", + "Noise executor handshake completed" + ); + debug!( + stream_id = validation_result.stream_id, + active_streams = streams.len() + 1, + "Noise executor stream activated" + ); + streams.insert( + validation_result.stream_id.clone(), + spawn_noise_virtual_stream( + validation_result.stream_id, + validation_result.validation_id, + handler.clone(), + physical_outgoing_tx.clone(), + closed_stream_tx.clone(), + transport, + executor_registration.clone(), + ), + ); + } + Some(Err(error)) => { + warn!("Noise relay harness key validation task failed: {error}"); + let stream_ids = pending_handshakes.keys().cloned().collect::>(); + pending_handshakes.clear(); + for stream_id in stream_ids { + send_reset(&physical_outgoing_tx, stream_id); + } + } + None => {} + } + continue; + } + incoming_message = websocket_stream.next() => match incoming_message { + Some(Ok(Message::Binary(payload))) => match decode_relay_message_frame(payload.as_ref()) { + Ok(frame) => frame, + Err(error) => { + warn!("dropping malformed Noise relay frame from harness: {error}"); + continue; + } + }, + Some(Ok(Message::Close(_))) | None => { + disconnect_reason = RendezvousDisconnectReason::PeerClose; + break; + } + Some(Ok(Message::Pong(_))) => { + let _ = pong_tx.try_send(()); + continue; + } + Some(Ok(Message::Ping(_) | Message::Frame(_))) => continue, + Some(Ok(Message::Text(_))) => { + warn!("dropping non-binary Noise relay frame from harness"); + continue; + } + Some(Err(error)) => { + debug!("Noise multiplexed environment websocket read failed: {error}"); + disconnect_reason = RendezvousDisconnectReason::ReadError; + break; + } + } + }; + + let kind = match frame.validate() { + Ok(kind) => kind, + Err(error) => { + warn!("dropping invalid Noise relay frame: {error}"); + continue; + } + }; + let stream_id = frame.stream_id.clone(); + match kind { + RelayFrameBodyKind::Handshake => { + // Reject duplicate or busy streams before paying for a hybrid + // handshake. Malformed attempts that reach cryptography are + // covered by the connection-wide failure budget below. + if streams.contains_key(&stream_id) { + send_reset(&physical_outgoing_tx, stream_id); + continue; + } + // Removing pending state makes the in-flight validation result stale. + if pending_handshakes.remove(&stream_id).is_some() { + send_reset(&physical_outgoing_tx, stream_id); + if failed_handshake_budget_exhausted(&mut failed_handshakes) { + warn!("closing Noise relay after repeated handshake failures"); + break; + } + continue; + } + if streams.len() >= MAX_ACTIVE_NOISE_RELAY_STREAMS { + warn!("Noise relay has too many active streams"); + send_reset(&physical_outgoing_tx, stream_id); + continue; + } + if validation_tasks.len() >= MAX_PENDING_HANDSHAKE_VALIDATIONS { + warn!("Noise relay has too many pending harness key validations"); + send_reset(&physical_outgoing_tx, stream_id); + continue; + } + let prologue = + noise_channel_prologue(&environment_id, &executor_registration_id, &stream_id); + let request = match frame.into_handshake_payload() { + Ok(request) => request, + Err(error) => { + warn!("failed to read Noise relay handshake frame: {error}"); + send_reset(&physical_outgoing_tx, stream_id); + continue; + } + }; + let mut pending = + match PendingResponderHandshake::read_request(&identity, &prologue, &request) { + Ok(pending) => pending, + Err(error) => { + warn!("failed to read Noise relay handshake request: {error}"); + send_reset(&physical_outgoing_tx, stream_id); + if failed_handshake_budget_exhausted(&mut failed_handshakes) { + warn!("closing Noise relay after repeated handshake failures"); + break; + } + continue; + } + }; + + // The authorization and authenticated harness key come from the + // same encrypted IK message and are validated together. + let authorization = match String::from_utf8(std::mem::take(&mut pending.payload)) { + Ok(authorization) + if authorization.len() <= MAX_HARNESS_KEY_AUTHORIZATION_BYTES => + { + Some(authorization) + } + Ok(_) => { + warn!("Noise relay handshake authorization is too long"); + None + } + Err(_) => { + warn!("Noise relay handshake authorization is not UTF-8"); + None + } + }; + let Some(authorization) = authorization else { + send_reset(&physical_outgoing_tx, stream_id); + if failed_handshake_budget_exhausted(&mut failed_handshakes) { + warn!("closing Noise relay after repeated handshake failures"); + break; + } + continue; + }; + let harness_public_key = pending.initiator_public_key.clone(); + let validation_id = next_validation_id; + next_validation_id += 1; + pending_handshakes.insert( + stream_id.clone(), + PendingHandshake { + validation_id, + handshake: pending, + }, + ); + let validator = validator.clone(); + + // Failed validation leaves no transport state and sends only a + // generic reset. + validation_tasks.spawn(async move { + let result = match timeout( + HARNESS_KEY_VALIDATION_TIMEOUT, + validator.validate_harness_key(&harness_public_key, &authorization), + ) + .await + { + Ok(result) => result, + Err(_) => Err(ExecServerError::Protocol( + "timed out validating Noise relay harness key".to_string(), + )), + }; + HarnessKeyValidationResult { + stream_id, + validation_id, + result, + } + }); + } + RelayFrameBodyKind::Data => { + // Removing pending state also makes any in-flight validation stale. + let Some(stream) = streams.get_mut(&stream_id) else { + let canceled_pending_handshake = + pending_handshakes.remove(&stream_id).is_some(); + send_reset(&physical_outgoing_tx, stream_id); + if canceled_pending_handshake + && failed_handshake_budget_exhausted(&mut failed_handshakes) + { + warn!("closing Noise relay after repeated handshake failures"); + break; + } + continue; + }; + let data = match frame.into_data() { + Ok(data) => data, + Err(error) => { + warn!("dropping malformed Noise relay data frame: {error}"); + streams.remove(&stream_id); + send_reset(&physical_outgoing_tx, stream_id); + continue; + } + }; + if let Err(error) = stream.receive_data(data) { + warn!("failed to process Noise relay payload: {error}"); + streams.remove(&stream_id); + send_reset(&physical_outgoing_tx, stream_id); + } + } + RelayFrameBodyKind::Reset => { + pending_handshakes.remove(&stream_id); + if let Some(stream) = streams.remove(&stream_id) { + // The reset reason is unauthenticated, so do not log it. + stream.disconnect(); + } + } + RelayFrameBodyKind::Ack + | RelayFrameBodyKind::Resume + | RelayFrameBodyKind::Heartbeat => {} + } + } + + for (_stream_id, stream) in streams { + stream.disconnect(); + } + // Dropping the JoinSet aborts any registry checks still running. + if !physical_writer_task.is_finished() { + physical_writer_task.abort(); + let _ = physical_writer_task.await; + } + disconnect_reason +} + +/// Charge one failed authenticated-channel attempt to this physical relay. +/// +/// Closing after a small fixed budget prevents a peer that has not been +/// authorized from triggering unbounded hybrid handshakes or registry checks. +fn failed_handshake_budget_exhausted(failed_handshakes: &mut usize) -> bool { + *failed_handshakes += 1; + *failed_handshakes >= MAX_FAILED_NOISE_HANDSHAKES +} + +/// Responder state held while registry authorization is pending. +struct PendingHandshake { + validation_id: u64, + handshake: PendingResponderHandshake, +} + +/// `validation_id` prevents an old check from completing a reused `stream_id`. +struct HarnessKeyValidationResult { + stream_id: String, + validation_id: u64, + result: Result<(), ExecServerError>, +} + +/// Queue a best-effort reset without blocking the shared websocket loop. +/// Reset reasons are relay control data and are not treated as trusted text. +fn send_reset(physical_outgoing_tx: &mpsc::Sender>, stream_id: String) { + let reset = RelayMessageFrame::reset(stream_id, NOISE_RELAY_RESET_REASON.to_string()); + let _ = physical_outgoing_tx.try_send(encode_relay_message_frame(&reset)); +} + +#[cfg(test)] +#[path = "relay_noise_tests.rs"] +mod noise_tests; + +#[cfg(test)] +mod tests { + use std::pin::Pin; + use std::sync::Arc; + use std::sync::atomic::AtomicBool; + use std::sync::atomic::Ordering; + use std::task::Context; + use std::task::Poll; + use std::time::Duration; + + use codex_exec_server_protocol::JSONRPCRequest; + use codex_exec_server_protocol::RequestId; + use futures::Sink; + use futures::Stream; + use futures::channel::mpsc as futures_mpsc; + use futures::task::AtomicWaker; + use pretty_assertions::assert_eq; + use tokio::net::TcpListener; + use tokio::time::timeout; + use tokio_tungstenite::WebSocketStream; + use tokio_tungstenite::accept_async; + use tokio_tungstenite::connect_async; + use tokio_tungstenite::tungstenite::Message; + + use super::*; + + #[tokio::test] + async fn harness_connection_sends_keepalive_and_receives_relay_data() -> anyhow::Result<()> { + let (client_websocket, mut server_websocket) = websocket_pair().await?; + let mut connection = + harness_connection_from_websocket(client_websocket, "test".to_string()); + let stream_id = read_resume_stream_id(&mut server_websocket).await?; + read_keepalive_ping(&mut server_websocket).await?; + server_websocket + .send(Message::Pong(b"keepalive".to_vec().into())) + .await?; + let message = test_jsonrpc_message(); + + server_websocket + .send(Message::Binary( + encode_relay_message_frame(&RelayMessageFrame::data( + stream_id, + /*seq*/ 0, + jsonrpc_payload(&message)?, + /*trace*/ None, + )) + .into(), + )) + .await?; + let Some(JsonRpcConnectionEvent::QueuedRequest { request, .. }) = + timeout(Duration::from_secs(1), connection.incoming_rx.recv()).await? + else { + anyhow::bail!("expected a queued JSON-RPC request"); + }; + assert_eq!(JSONRPCMessage::Request(request), message); + + drop(connection); + Ok(()) + } + + #[tokio::test] + async fn multiplexed_environment_sends_keepalive() -> anyhow::Result<()> { + let (client_websocket, mut server_websocket) = websocket_pair().await?; + let runtime_paths = crate::ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + ) + .map_err(anyhow::Error::from)?; + let environment_task = tokio::spawn(run_multiplexed_environment( + client_websocket, + ConnectionProcessor::new(runtime_paths), + "test-environment".to_string(), + "test-registration".to_string(), + NoiseChannelIdentity::generate()?, + AllowHarnessKeyValidator, + )); + + read_keepalive_ping(&mut server_websocket).await?; + + environment_task.abort(); + let _ = environment_task.await; + Ok(()) + } + + #[derive(Clone)] + struct AllowHarnessKeyValidator; + + impl HarnessKeyValidator for AllowHarnessKeyValidator { + async fn validate_harness_key( + &self, + _harness_public_key: &NoiseChannelPublicKey, + _authorization: &str, + ) -> Result<(), ExecServerError> { + Ok(()) + } + } + + #[tokio::test] + async fn send_event_with_keepalive_pings_while_incoming_queue_is_full() -> anyhow::Result<()> { + let (mut websocket, _control, mut outbound_rx) = + ControlledWebSocket::new(/*write_ready*/ true); + let (incoming_tx, mut incoming_rx) = mpsc::channel(/*buffer*/ 1); + let message = test_jsonrpc_message(); + let expected_message = message.clone(); + incoming_tx + .send(JsonRpcConnectionEvent::MalformedMessage { + reason: "first".to_string(), + }) + .await?; + let mut keepalive = tokio::time::interval_at( + tokio::time::Instant::now() + WEBSOCKET_KEEPALIVE_INTERVAL, + WEBSOCKET_KEEPALIVE_INTERVAL, + ); + keepalive.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + + let send_task = tokio::spawn(async move { + send_event_with_keepalive( + &mut websocket, + &mut keepalive, + &incoming_tx, + JsonRpcConnectionEvent::Message(message), + ) + .await + }); + + assert!(matches!( + timeout(Duration::from_secs(1), outbound_rx.next()).await?, + Some(Message::Ping(_)) + )); + assert!(matches!( + incoming_rx.recv().await, + Some(JsonRpcConnectionEvent::MalformedMessage { reason }) if reason == "first" + )); + assert!(matches!( + timeout(Duration::from_secs(1), send_task).await??, + Ok(()) + )); + assert!(matches!( + incoming_rx.recv().await, + Some(JsonRpcConnectionEvent::Message(actual)) if actual == expected_message + )); + Ok(()) + } + + #[tokio::test] + async fn harness_connection_reports_text_frames_as_malformed() -> anyhow::Result<()> { + let (client_websocket, mut server_websocket) = websocket_pair().await?; + let mut connection = + harness_connection_from_websocket(client_websocket, "test".to_string()); + + read_resume_stream_id(&mut server_websocket).await?; + server_websocket.send(Message::Text("nope".into())).await?; + assert!(matches!( + timeout(Duration::from_secs(1), connection.incoming_rx.recv()).await?, + Some(JsonRpcConnectionEvent::MalformedMessage { reason }) + if reason == "relay exec-server transport expects binary protobuf frames" + )); + + drop(connection); + Ok(()) + } + + #[tokio::test] + async fn harness_connection_reports_server_close() -> anyhow::Result<()> { + let (client_websocket, mut server_websocket) = websocket_pair().await?; + let mut connection = + harness_connection_from_websocket(client_websocket, "test".to_string()); + + read_resume_stream_id(&mut server_websocket).await?; + server_websocket.close(None).await?; + assert!(matches!( + timeout(Duration::from_secs(1), connection.incoming_rx.recv()).await?, + Some(JsonRpcConnectionEvent::Disconnected { reason: None }) + )); + + drop(connection); + Ok(()) + } + + #[tokio::test] + async fn harness_connection_keeps_outbound_frame_while_send_is_backpressured() + -> anyhow::Result<()> { + let (websocket, control, mut outbound_rx) = + ControlledWebSocket::new(/*write_ready*/ true); + let mut connection = harness_connection_from_websocket(websocket, "test".to_string()); + let Message::Binary(resume_payload) = timeout(Duration::from_secs(1), outbound_rx.next()) + .await? + .expect("resume frame") + else { + anyhow::bail!("expected relay resume frame"); + }; + let stream_id = decode_relay_message_frame(resume_payload.as_ref())?.stream_id; + let message = test_jsonrpc_message(); + + control.set_write_blocked(); + connection.outgoing_tx.send(message.clone()).await?; + control.wait_for_blocked_write().await?; + control.send_inbound(Message::Pong(b"check".to_vec().into()))?; + assert!( + timeout(Duration::from_millis(50), connection.incoming_rx.recv()) + .await + .is_err() + ); + + control.set_write_ready(); + let Message::Binary(data_payload) = timeout(Duration::from_secs(1), outbound_rx.next()) + .await? + .expect("data frame") + else { + anyhow::bail!("expected relay data frame"); + }; + let frame = decode_relay_message_frame(data_payload.as_ref())?; + assert_eq!(frame.stream_id, stream_id); + assert_eq!(frame.into_jsonrpc_message()?, message); + drop(connection); + Ok(()) + } + + async fn websocket_pair() -> anyhow::Result<( + WebSocketStream>, + WebSocketStream, + )> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let server_task = tokio::spawn(async move { + let (stream, _) = listener.accept().await?; + accept_async(stream).await.map_err(anyhow::Error::from) + }); + let (client_websocket, _) = connect_async(websocket_url).await?; + let server_websocket = server_task.await??; + Ok((client_websocket, server_websocket)) + } + + async fn read_resume_stream_id( + websocket: &mut WebSocketStream, + ) -> anyhow::Result { + let message = timeout(Duration::from_secs(1), websocket.next()) + .await? + .expect("websocket should stay open")?; + let Message::Binary(payload) = message else { + anyhow::bail!("expected relay resume frame, got {message:?}"); + }; + let frame = decode_relay_message_frame(payload.as_ref())?; + assert_eq!(frame.validate()?, RelayFrameBodyKind::Resume); + Ok(frame.stream_id) + } + + async fn read_keepalive_ping( + websocket: &mut WebSocketStream, + ) -> anyhow::Result<()> { + loop { + let Some(message) = timeout(Duration::from_secs(1), websocket.next()).await? else { + anyhow::bail!("websocket closed before keepalive ping"); + }; + match message? { + Message::Ping(_) => return Ok(()), + Message::Binary(_) | Message::Text(_) | Message::Pong(_) | Message::Frame(_) => {} + Message::Close(_) => anyhow::bail!("websocket closed before keepalive ping"), + } + } + } + + fn test_jsonrpc_message() -> JSONRPCMessage { + JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(1), + method: "test".to_string(), + params: None, + trace: None, + }) + } + + struct ControlledWebSocket { + inbound_rx: futures_mpsc::UnboundedReceiver>, + outbound_tx: futures_mpsc::UnboundedSender, + write_ready: Arc, + write_blocked: Arc, + write_blocked_waker: Arc, + write_waker: Arc, + } + + struct ControlledWebSocketHandle { + inbound_tx: futures_mpsc::UnboundedSender>, + write_ready: Arc, + write_blocked: Arc, + write_blocked_waker: Arc, + write_waker: Arc, + } + + impl ControlledWebSocket { + fn new( + write_ready: bool, + ) -> ( + Self, + ControlledWebSocketHandle, + futures_mpsc::UnboundedReceiver, + ) { + let (inbound_tx, inbound_rx) = futures_mpsc::unbounded(); + let (outbound_tx, outbound_rx) = futures_mpsc::unbounded(); + let write_ready = Arc::new(AtomicBool::new(write_ready)); + let write_blocked = Arc::new(AtomicBool::new(false)); + let write_blocked_waker = Arc::new(AtomicWaker::new()); + let write_waker = Arc::new(AtomicWaker::new()); + ( + Self { + inbound_rx, + outbound_tx, + write_ready: Arc::clone(&write_ready), + write_blocked: Arc::clone(&write_blocked), + write_blocked_waker: Arc::clone(&write_blocked_waker), + write_waker: Arc::clone(&write_waker), + }, + ControlledWebSocketHandle { + inbound_tx, + write_ready, + write_blocked, + write_blocked_waker, + write_waker, + }, + outbound_rx, + ) + } + } + + impl ControlledWebSocketHandle { + fn send_inbound(&self, message: Message) -> anyhow::Result<()> { + self.inbound_tx + .unbounded_send(Ok(message)) + .map_err(anyhow::Error::from) + } + + fn set_write_blocked(&self) { + self.write_ready.store(false, Ordering::Release); + } + + fn set_write_ready(&self) { + self.write_ready.store(true, Ordering::Release); + self.write_waker.wake(); + } + + async fn wait_for_blocked_write(&self) -> anyhow::Result<()> { + timeout( + Duration::from_secs(1), + futures::future::poll_fn(|cx| { + if self.write_blocked.load(Ordering::Acquire) { + Poll::Ready(()) + } else { + self.write_blocked_waker.register(cx.waker()); + Poll::Pending + } + }), + ) + .await?; + Ok(()) + } + } + + impl Sink for ControlledWebSocket { + type Error = std::convert::Infallible; + + fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + if self.write_ready.load(Ordering::Acquire) { + Poll::Ready(Ok(())) + } else { + self.write_blocked.store(true, Ordering::Release); + self.write_blocked_waker.wake(); + self.write_waker.register(cx.waker()); + Poll::Pending + } + } + + fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> { + self.outbound_tx + .unbounded_send(item) + .expect("test outbound receiver should stay open"); + Ok(()) + } + + fn poll_flush( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } + } + + impl Stream for ControlledWebSocket { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.inbound_rx).poll_next(cx) + } + } +} diff --git a/codex-rs/exec-server/src/relay_noise_tests.rs b/codex-rs/exec-server/src/relay_noise_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..b92d22b9bfd4679055cbb59420e323a2587ba9ae --- /dev/null +++ b/codex-rs/exec-server/src/relay_noise_tests.rs @@ -0,0 +1,553 @@ +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use anyhow::Result; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::RequestId; +use futures::SinkExt; +use futures::StreamExt; +use pretty_assertions::assert_eq; +use tokio::net::TcpListener; +use tokio::sync::Notify; +use tokio::time::timeout; +use tokio_tungstenite::accept_async; +use tokio_tungstenite::connect_async; +use tokio_tungstenite::tungstenite::Message; +use tokio_util::task::AbortOnDropHandle; + +use super::HarnessKeyValidator; +use super::MAX_FAILED_NOISE_HANDSHAKES; +use super::MAX_HARNESS_KEY_AUTHORIZATION_BYTES; +use super::RendezvousDisconnectReason; +use super::run_multiplexed_environment; +use crate::ExecServerError; +use crate::ExecServerRuntimePaths; +use crate::connection::JsonRpcConnectionEvent; +use crate::noise_channel::InitiatorHandshake; +use crate::noise_channel::NoiseChannelIdentity; +use crate::noise_channel::NoiseChannelPublicKey; +use crate::noise_channel::noise_channel_prologue; +use crate::noise_relay::NoiseHarnessConnectionArgs; +use crate::noise_relay::noise_harness_connection_from_websocket_with_readiness; +use crate::noise_relay::stream_handler::NoiseOutboundMessage; +use crate::noise_relay::stream_handler::NoiseStreamConnection; +use crate::noise_relay::stream_handler::NoiseStreamHandler; +use crate::relay::RelayFrameBodyKind; +use crate::relay::decode_relay_message_frame; +use crate::relay::encode_relay_message_frame; +use crate::relay_proto::RelayMessageFrame; +use crate::server::ConnectionProcessor; + +const ENVIRONMENT_ID: &str = "environment-1"; +const EXECUTOR_REGISTRATION_ID: &str = "registration-1"; + +#[derive(Clone)] +struct ObservedRegistration( + ConnectionProcessor, + tokio::sync::mpsc::Sender>, +); + +impl NoiseStreamHandler for ObservedRegistration { + type Incoming = JsonRpcConnectionEvent; + type Outgoing = JSONRPCMessage; + + fn decode(payload: bytes::Bytes) -> Result { + ConnectionProcessor::decode(payload) + } + fn encode(message: Self::Outgoing) -> Result { + ConnectionProcessor::encode(message) + } + async fn run_connection( + self, + connection: NoiseStreamConnection, + ) { + let registration = connection + .executor_registration + .as_ref() + .map(|registration| { + ( + registration.environment_id.clone(), + registration.executor_registration_id.clone(), + ) + }); + let _ = self.1.send(registration).await; + NoiseStreamHandler::run_connection(self.0, connection).await; + } +} + +#[tokio::test] +async fn missing_pong_disconnects_physical_relay() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let harness_connection = tokio::spawn(connect_async(websocket_url)); + let (socket, _peer_addr) = listener.accept().await?; + let environment_websocket = accept_async(socket).await?; + let (_harness_websocket, _response) = harness_connection.await??; + + let environment_task = tokio::spawn(run_multiplexed_environment( + environment_websocket, + ConnectionProcessor::new(ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?), + ENVIRONMENT_ID.to_string(), + EXECUTOR_REGISTRATION_ID.to_string(), + NoiseChannelIdentity::generate()?, + BlockingValidator { + calls: Arc::new(AtomicUsize::new(0)), + release: Arc::new(Notify::new()), + }, + )); + + assert_eq!( + timeout(Duration::from_secs(1), environment_task).await??, + RendezvousDisconnectReason::PongTimeout + ); + Ok(()) +} + +#[tokio::test] +async fn pong_keeps_physical_relay_connected() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let harness_connection = tokio::spawn(connect_async(websocket_url)); + let (socket, _peer_addr) = listener.accept().await?; + let environment_websocket = accept_async(socket).await?; + let (mut harness_websocket, _response) = harness_connection.await??; + + let environment_task = tokio::spawn(run_multiplexed_environment( + environment_websocket, + ConnectionProcessor::new(ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?), + ENVIRONMENT_ID.to_string(), + EXECUTOR_REGISTRATION_ID.to_string(), + NoiseChannelIdentity::generate()?, + BlockingValidator { + calls: Arc::new(AtomicUsize::new(0)), + release: Arc::new(Notify::new()), + }, + )); + + timeout(Duration::from_secs(1), async { + let mut pings = 0; + while pings < 6 { + match harness_websocket.next().await { + Some(Ok(Message::Ping(payload))) => { + harness_websocket.send(Message::Pong(payload)).await?; + pings += 1; + } + Some(Ok(Message::Pong(_) | Message::Frame(_))) => {} + Some(Ok(message)) => anyhow::bail!("expected keepalive ping, got {message:?}"), + Some(Err(error)) => return Err(error.into()), + None => anyhow::bail!("environment disconnected before six keepalive pings"), + } + } + Ok::<_, anyhow::Error>(()) + }) + .await??; + harness_websocket.close(None).await?; + assert_eq!( + timeout(Duration::from_secs(1), environment_task).await??, + RendezvousDisconnectReason::PeerClose + ); + Ok(()) +} + +#[derive(Clone)] +struct BlockingValidator { + calls: Arc, + release: Arc, +} + +impl HarnessKeyValidator for BlockingValidator { + fn validate_harness_key( + &self, + _harness_public_key: &NoiseChannelPublicKey, + _authorization: &str, + ) -> impl std::future::Future> + Send { + let calls = Arc::clone(&self.calls); + let release = Arc::clone(&self.release); + async move { + calls.fetch_add(1, Ordering::SeqCst); + release.notified().await; + Ok(()) + } + } +} + +#[tokio::test] +async fn processor_exit_resets_noise_harness_stream() -> Result<()> { + let (registration_tx, mut registration_rx) = tokio::sync::mpsc::channel(1); + let listener = TcpListener::bind("127.0.0.1:0").await?; + let connecting = tokio::spawn(connect_async(format!("ws://{}", listener.local_addr()?))); + let (socket, _) = listener.accept().await?; + let environment_websocket = accept_async(socket).await?; + let (harness_websocket, _) = connecting.await??; + let identity = NoiseChannelIdentity::generate()?; + let release = Arc::new(Notify::new()); + release.notify_one(); + let environment_task = AbortOnDropHandle::new(tokio::spawn(run_multiplexed_environment( + environment_websocket, + ObservedRegistration( + ConnectionProcessor::new(ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?), + registration_tx, + ), + ENVIRONMENT_ID.to_string(), + EXECUTOR_REGISTRATION_ID.to_string(), + identity.clone(), + BlockingValidator { + calls: Arc::new(AtomicUsize::new(0)), + release, + }, + ))); + let mut connection = noise_harness_connection_from_websocket_with_readiness( + harness_websocket, + NoiseHarnessConnectionArgs { + connection_label: "processor exit test".to_string(), + environment_id: ENVIRONMENT_ID.to_string(), + executor_registration_id: EXECUTOR_REGISTRATION_ID.to_string(), + identity: NoiseChannelIdentity::generate()?, + responder_public_key: identity.public_key(), + harness_key_authorization: "authorization".to_string(), + }, + ) + .connection; + // Valid JSON reaches the processor; an unsolicited response closes it and + // aborts its writer. The physical relay must still deliver the reset. + connection + .outgoing_tx + .send(JSONRPCMessage::Response(JSONRPCResponse { + id: RequestId::Integer(1), + result: serde_json::Value::Null, + })) + .await?; + assert_eq!( + timeout(Duration::from_secs(1), registration_rx.recv()).await?, + Some(Some(( + ENVIRONMENT_ID.to_string(), + EXECUTOR_REGISTRATION_ID.to_string() + ))), + ); + assert!(matches!( + timeout(Duration::from_secs(1), connection.incoming_rx.recv()).await?, + Some(JsonRpcConnectionEvent::Disconnected { reason: Some(reason) }) + if reason == "Noise relay stream reset" + )); + for task in connection.task_handles { + task.abort(); + let _ = task.await; + } + environment_task.abort(); + let _ = environment_task.await; + Ok(()) +} + +#[tokio::test] +async fn pending_harness_key_validation_does_not_block_new_handshakes() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let harness_connection = tokio::spawn(connect_async(websocket_url)); + let (socket, _peer_addr) = listener.accept().await?; + let environment_websocket = accept_async(socket).await?; + let (mut harness_websocket, _response) = harness_connection.await??; + + let environment_identity = NoiseChannelIdentity::generate()?; + let harness_identity = NoiseChannelIdentity::generate()?; + let calls = Arc::new(AtomicUsize::new(0)); + let environment_task = tokio::spawn(run_multiplexed_environment( + environment_websocket, + ConnectionProcessor::new(ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?), + ENVIRONMENT_ID.to_string(), + EXECUTOR_REGISTRATION_ID.to_string(), + environment_identity.clone(), + BlockingValidator { + calls: Arc::clone(&calls), + release: Arc::new(Notify::new()), + }, + )); + + for stream_id in ["stream-1", "stream-2"] { + let prologue = noise_channel_prologue(ENVIRONMENT_ID, EXECUTOR_REGISTRATION_ID, stream_id); + let (_handshake, request) = InitiatorHandshake::start( + &harness_identity, + &environment_identity.public_key(), + &prologue, + b"authorization", + )?; + let frame = RelayMessageFrame::handshake(stream_id.to_string(), request); + harness_websocket + .send(Message::Binary(encode_relay_message_frame(&frame).into())) + .await?; + } + + timeout(Duration::from_secs(1), async { + while calls.load(Ordering::SeqCst) != 2 { + tokio::task::yield_now().await; + } + }) + .await?; + + harness_websocket.close(None).await?; + timeout(Duration::from_secs(1), environment_task).await??; + Ok(()) +} + +#[tokio::test] +async fn duplicate_handshakes_exhaust_failure_budget() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let harness_connection = tokio::spawn(connect_async(websocket_url)); + let (socket, _peer_addr) = listener.accept().await?; + let environment_websocket = accept_async(socket).await?; + let (mut harness_websocket, _response) = harness_connection.await??; + + let environment_identity = NoiseChannelIdentity::generate()?; + let harness_identity = NoiseChannelIdentity::generate()?; + let calls = Arc::new(AtomicUsize::new(0)); + let release = Arc::new(Notify::new()); + let environment_task = tokio::spawn(run_multiplexed_environment( + environment_websocket, + ConnectionProcessor::new(ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?), + ENVIRONMENT_ID.to_string(), + EXECUTOR_REGISTRATION_ID.to_string(), + environment_identity.clone(), + BlockingValidator { + calls: Arc::clone(&calls), + release: Arc::clone(&release), + }, + )); + + let stream_id = "stream-1"; + let prologue = noise_channel_prologue(ENVIRONMENT_ID, EXECUTOR_REGISTRATION_ID, stream_id); + let (_handshake, request) = InitiatorHandshake::start( + &harness_identity, + &environment_identity.public_key(), + &prologue, + b"authorization", + )?; + let frame = RelayMessageFrame::handshake(stream_id.to_string(), request); + let encoded = encode_relay_message_frame(&frame); + harness_websocket + .send(Message::Binary(encoded.clone().into())) + .await?; + timeout(Duration::from_secs(1), async { + while calls.load(Ordering::SeqCst) != 1 { + tokio::task::yield_now().await; + } + }) + .await?; + + for attempt in 1..MAX_FAILED_NOISE_HANDSHAKES { + if attempt > 1 { + harness_websocket + .send(Message::Binary(encoded.clone().into())) + .await?; + timeout(Duration::from_secs(1), async { + while calls.load(Ordering::SeqCst) != attempt { + tokio::task::yield_now().await; + } + }) + .await?; + } + harness_websocket + .send(Message::Binary(encoded.clone().into())) + .await?; + let payload = timeout(Duration::from_secs(1), async { + loop { + match harness_websocket.next().await { + Some(Ok(Message::Binary(payload))) => break Ok(payload), + Some(Ok(Message::Ping(_) | Message::Pong(_) | Message::Frame(_))) => {} + Some(Ok(message)) => anyhow::bail!("expected reset frame, got {message:?}"), + Some(Err(error)) => break Err(error.into()), + None => anyhow::bail!("environment closed before sending reset"), + } + } + }) + .await??; + let reset = decode_relay_message_frame(payload.as_ref())?; + assert_eq!(reset.stream_id, stream_id); + assert_eq!(reset.validate()?, RelayFrameBodyKind::Reset); + } + + harness_websocket + .send(Message::Binary(encoded.clone().into())) + .await?; + timeout(Duration::from_secs(1), async { + while calls.load(Ordering::SeqCst) != MAX_FAILED_NOISE_HANDSHAKES { + tokio::task::yield_now().await; + } + }) + .await?; + harness_websocket + .send(Message::Binary(encoded.into())) + .await?; + timeout(Duration::from_secs(1), environment_task).await??; + release.notify_waiters(); + Ok(()) +} + +#[tokio::test] +async fn oversized_harness_authorization_is_rejected_before_validation() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let harness_connection = tokio::spawn(connect_async(websocket_url)); + let (socket, _peer_addr) = listener.accept().await?; + let environment_websocket = accept_async(socket).await?; + let (mut harness_websocket, _response) = harness_connection.await??; + + let environment_identity = NoiseChannelIdentity::generate()?; + let harness_identity = NoiseChannelIdentity::generate()?; + let calls = Arc::new(AtomicUsize::new(0)); + let environment_task = tokio::spawn(run_multiplexed_environment( + environment_websocket, + ConnectionProcessor::new(ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?), + ENVIRONMENT_ID.to_string(), + EXECUTOR_REGISTRATION_ID.to_string(), + environment_identity.clone(), + BlockingValidator { + calls: Arc::clone(&calls), + release: Arc::new(Notify::new()), + }, + )); + + let stream_id = "stream-1"; + let prologue = noise_channel_prologue(ENVIRONMENT_ID, EXECUTOR_REGISTRATION_ID, stream_id); + let oversized_authorization = vec![b'a'; MAX_HARNESS_KEY_AUTHORIZATION_BYTES + 1]; + let (_handshake, request) = InitiatorHandshake::start( + &harness_identity, + &environment_identity.public_key(), + &prologue, + &oversized_authorization, + )?; + let frame = RelayMessageFrame::handshake(stream_id.to_string(), request); + harness_websocket + .send(Message::Binary(encode_relay_message_frame(&frame).into())) + .await?; + + let Message::Binary(payload) = timeout(Duration::from_secs(1), harness_websocket.next()) + .await? + .ok_or_else(|| anyhow::anyhow!("environment closed before sending reset"))?? + else { + anyhow::bail!("expected binary reset frame"); + }; + let reset = decode_relay_message_frame(payload.as_ref())?; + assert_eq!(reset.validate()?, RelayFrameBodyKind::Reset); + assert_eq!(calls.load(Ordering::SeqCst), 0); + + harness_websocket.close(None).await?; + timeout(Duration::from_secs(1), environment_task).await??; + Ok(()) +} + +#[tokio::test] +async fn repeated_malformed_handshakes_close_the_physical_relay() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let harness_connection = tokio::spawn(connect_async(websocket_url)); + let (socket, _peer_addr) = listener.accept().await?; + let environment_websocket = accept_async(socket).await?; + let (mut harness_websocket, _response) = harness_connection.await??; + + let environment_identity = NoiseChannelIdentity::generate()?; + let harness_identity = NoiseChannelIdentity::generate()?; + let environment_task = tokio::spawn(run_multiplexed_environment( + environment_websocket, + ConnectionProcessor::new(ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?), + ENVIRONMENT_ID.to_string(), + EXECUTOR_REGISTRATION_ID.to_string(), + environment_identity.clone(), + BlockingValidator { + calls: Arc::new(AtomicUsize::new(0)), + release: Arc::new(Notify::new()), + }, + )); + + for attempt in 0..MAX_FAILED_NOISE_HANDSHAKES { + let stream_id = format!("malformed-{attempt}"); + let prologue = noise_channel_prologue(ENVIRONMENT_ID, EXECUTOR_REGISTRATION_ID, &stream_id); + let (_handshake, mut request) = InitiatorHandshake::start( + &harness_identity, + &environment_identity.public_key(), + &prologue, + b"authorization", + )?; + let last_byte = request.last_mut().expect("handshake request is not empty"); + *last_byte ^= 1; + let frame = RelayMessageFrame::handshake(stream_id, request); + harness_websocket + .send(Message::Binary(encode_relay_message_frame(&frame).into())) + .await?; + } + + timeout(Duration::from_secs(1), environment_task).await??; + Ok(()) +} + +#[tokio::test] +async fn repeated_early_data_during_validation_closes_the_physical_relay() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let harness_connection = tokio::spawn(connect_async(websocket_url)); + let (socket, _peer_addr) = listener.accept().await?; + let environment_websocket = accept_async(socket).await?; + let (mut harness_websocket, _response) = harness_connection.await??; + + let environment_identity = NoiseChannelIdentity::generate()?; + let harness_identity = NoiseChannelIdentity::generate()?; + let environment_task = tokio::spawn(run_multiplexed_environment( + environment_websocket, + ConnectionProcessor::new(ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?), + ENVIRONMENT_ID.to_string(), + EXECUTOR_REGISTRATION_ID.to_string(), + environment_identity.clone(), + BlockingValidator { + calls: Arc::new(AtomicUsize::new(0)), + release: Arc::new(Notify::new()), + }, + )); + + for attempt in 0..MAX_FAILED_NOISE_HANDSHAKES { + let stream_id = format!("early-data-{attempt}"); + let prologue = noise_channel_prologue(ENVIRONMENT_ID, EXECUTOR_REGISTRATION_ID, &stream_id); + let (_handshake, request) = InitiatorHandshake::start( + &harness_identity, + &environment_identity.public_key(), + &prologue, + b"authorization", + )?; + for frame in [ + RelayMessageFrame::handshake(stream_id.clone(), request), + RelayMessageFrame::data(stream_id, /*seq*/ 0, vec![0], /*trace*/ None), + ] { + harness_websocket + .send(Message::Binary(encode_relay_message_frame(&frame).into())) + .await?; + } + } + + timeout(Duration::from_secs(1), environment_task).await??; + Ok(()) +} diff --git a/codex-rs/exec-server/src/relay_proto.rs b/codex-rs/exec-server/src/relay_proto.rs new file mode 100644 index 0000000000000000000000000000000000000000..da7cb5296d46bb7c76a5baeb452909608fbb4c14 --- /dev/null +++ b/codex-rs/exec-server/src/relay_proto.rs @@ -0,0 +1,9 @@ +#[path = "proto/codex.exec_server.relay.v1.rs"] +mod generated; + +pub(crate) use generated::RelayData; +pub(crate) use generated::RelayHandshake; +pub(crate) use generated::RelayMessageFrame; +pub(crate) use generated::RelayReset; +pub(crate) use generated::RelayResume; +pub(crate) use generated::relay_message_frame; diff --git a/codex-rs/exec-server/src/remote.rs b/codex-rs/exec-server/src/remote.rs new file mode 100644 index 0000000000000000000000000000000000000000..0d2711b7a8a4c4284220cc2e296dd04e678cc879 --- /dev/null +++ b/codex-rs/exec-server/src/remote.rs @@ -0,0 +1,1318 @@ +use std::sync::Arc; +use std::time::Duration; +use std::time::Instant; + +use codex_api::AuthProvider; +use codex_api::SharedAuthProvider; +use codex_http_client::ClientRouteClass; +use codex_http_client::HttpClientFactory; +use codex_http_client::HttpResponse; +use codex_http_client::RouteAwareClientPool; +use futures::FutureExt; +use http::HeaderMap; +use http::HeaderName; +use http::HeaderValue; +use http::StatusCode; +use serde::Deserialize; +use tokio::time::sleep; +use tokio::time::timeout_at; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tracing::debug; +use tracing::info; +use tracing::warn; + +use codex_utils_rustls_provider::ensure_rustls_crypto_provider; +use codex_websocket_client::WebSocketConnection; +use codex_websocket_client::WebSocketConnector; +use codex_websocket_client::WebSocketTlsMode; + +use crate::EnvironmentRegistryConnectRequest; +use crate::EnvironmentRegistryConnectResponse; +use crate::EnvironmentRegistryHarnessKeyValidationRequest; +use crate::EnvironmentRegistryHarnessKeyValidationResponse; +use crate::EnvironmentRegistryRegistrationRequest; +use crate::EnvironmentRegistryRegistrationResponse; +use crate::ExecServerError; +use crate::ExecServerRuntimePaths; +use crate::ExecServerTelemetry; +use crate::NoiseChannelIdentity; +use crate::NoiseChannelPublicKey; +use crate::NoiseRendezvousConnectBundle; +use crate::NoiseRendezvousConnectProvider; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::forward::Forwarder; +use crate::noise_relay::noise_relay_websocket_config; +use crate::noise_relay::stream_handler::NoiseStreamHandler; +use crate::relay::HarnessKeyValidator; +use crate::relay::run_multiplexed_environment; +use crate::server::ConnectionProcessor; +use crate::server::RequestDispatchMode; +use crate::trace_context::current_rendezvous_headers; +use crate::trace_context::current_trace_context_headers; + +#[path = "remote/direct.rs"] +mod direct; + +use direct::run_direct_environment; + +const ERROR_BODY_PREVIEW_BYTES: usize = 4096; +const NOISE_RELAY_SECURITY_PROFILE: &str = "noise_hybrid_ik_v1"; + +mod registration_retry; + +/// Wire transport used after registering a remote exec-server. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +/// The transport used to connect a remote exec-server environment. +pub enum RemoteEnvironmentTransport { + #[default] + Noise, + Direct, +} + +#[derive(Clone)] +struct EnvironmentRegistryClient { + base_url: String, + auth_provider: SharedAuthProvider, + http: RouteAwareClientPool, + connect_timeout: Duration, + telemetry: ExecServerTelemetry, +} + +impl std::fmt::Debug for EnvironmentRegistryClient { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("EnvironmentRegistryClient") + .field("base_url", &self.base_url) + .field("auth_provider", &"") + .finish_non_exhaustive() + } +} + +impl EnvironmentRegistryClient { + #[cfg(test)] + fn new(base_url: String, auth_provider: SharedAuthProvider) -> Result { + Self::new_with_telemetry( + base_url, + auth_provider, + ExecServerTelemetry::default(), + HttpClientFactory::new(codex_http_client::OutboundProxyPolicy::ReqwestDefault), + ) + } + + fn new_with_telemetry( + base_url: String, + auth_provider: SharedAuthProvider, + telemetry: ExecServerTelemetry, + http_client_factory: HttpClientFactory, + ) -> Result { + let base_url = normalize_base_url(base_url)?; + Ok(Self { + base_url, + auth_provider, + http: RouteAwareClientPool::new_without_redirects_or_request_logging( + http_client_factory, + ClientRouteClass::Api, + ), + connect_timeout: DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + telemetry, + }) + } + + /// Register the executor public key and obtain the rendezvous allocation. + /// The returned registration ID is included in each stream's Noise prologue. + #[tracing::instrument( + name = "codex.exec_server.remote.register", + skip_all, + fields( + otel.kind = "client", + otel.name = "codex.exec_server.remote.register", + result = tracing::field::Empty, + ) + )] + async fn register_environment( + &self, + environment_id: &str, + executor_public_key: &NoiseChannelPublicKey, + ) -> Result { + let started_at = Instant::now(); + let response = self + .register_environment_inner(environment_id, executor_public_key) + .await; + let result = if response.is_ok() { "success" } else { "error" }; + tracing::Span::current().record("result", result); + self.telemetry + .remote_registration_completed(result, started_at.elapsed()); + response + } + + async fn register_environment_inner( + &self, + environment_id: &str, + executor_public_key: &NoiseChannelPublicKey, + ) -> Result { + let deadline = tokio::time::Instant::now() + self.connect_timeout; + let url = endpoint_url( + &self.base_url, + &format!("/cloud/environment/{environment_id}/register"), + ); + let body = EnvironmentRegistryRegistrationRequest { + security_profile: NOISE_RELAY_SECURITY_PROFILE.to_string(), + executor_public_key: executor_public_key.clone(), + }; + let response = timeout_at(deadline, async { + self.http + .post(url) + .headers(self.resolve_auth_headers().await?) + .headers(current_trace_context_headers()) + .json(&body) + .send() + .await + .map_err(ExecServerError::EnvironmentRegistryRequest) + }) + .await + .unwrap_or_else(|_| { + Err(ExecServerError::EnvironmentRegistryRequest( + codex_http_client::RouteAwareRequestError::Timeout, + )) + })?; + let status = response.status(); + // Read diagnostics within the same attempt budget, preserving a known error status. + let response: EnvironmentRegistryRegistrationResponse = + timeout_at(deadline, self.parse_json_response(response)) + .await + .unwrap_or_else(|_| { + Err(match status { + StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => { + environment_registry_auth_error(status, "response body timed out") + } + status if !status.is_success() => { + environment_registry_http_error(status, "response body timed out") + } + _ => ExecServerError::EnvironmentRegistryRequest( + codex_http_client::RouteAwareRequestError::Timeout, + ), + }) + })?; + if response.environment_id != environment_id { + return Err(ExecServerError::Protocol( + "environment registry returned a different environment id".to_string(), + )); + } + if response.security_profile != NOISE_RELAY_SECURITY_PROFILE { + return Err(ExecServerError::Protocol(format!( + "environment registry returned unsupported security profile `{}`", + response.security_profile + ))); + } + info!( + noise_event = "registration", + noise_outcome = "ok", + security_profile = NOISE_RELAY_SECURITY_PROFILE, + "Noise executor registration completed" + ); + debug!( + environment_id = response.environment_id, + executor_registration_id = response.executor_registration_id, + "Noise executor registration details" + ); + Ok(response) + } + + /// Authorize one Noise harness key and obtain the full rendezvous bundle. + #[tracing::instrument( + name = "codex.exec_server.remote.environment_registry.connect", + skip_all, + fields( + otel.kind = "client", + otel.name = "codex.exec_server.remote.environment_registry.connect", + environment_id = %environment_id, + ) + )] + async fn connect_environment( + &self, + environment_id: &str, + harness_public_key: NoiseChannelPublicKey, + ) -> Result { + let url = endpoint_url( + &self.base_url, + &format!("/cloud/environment/{environment_id}/connect"), + ); + let body = EnvironmentRegistryConnectRequest { harness_public_key }; + let response = self + .http + .post(url) + .headers(self.resolve_auth_headers().await?) + .headers(current_trace_context_headers()) + .json(&body) + .timeout(self.connect_timeout) + .send() + .await?; + let response: EnvironmentRegistryConnectResponse = + self.parse_json_response(response).await?; + if response.environment_id != environment_id { + return Err(ExecServerError::Protocol( + "environment registry returned a different environment id".to_string(), + )); + } + if response.security_profile != NOISE_RELAY_SECURITY_PROFILE { + return Err(ExecServerError::Protocol(format!( + "environment registry returned unsupported security profile `{}`", + response.security_profile + ))); + } + if response.url.trim().is_empty() + || response.executor_registration_id.trim().is_empty() + || response.harness_key_authorization.trim().is_empty() + { + return Err(ExecServerError::Protocol( + "environment registry returned incomplete Noise connection data".to_string(), + )); + } + Ok(NoiseRendezvousConnectBundle { + websocket_url: response.url, + environment_id: response.environment_id, + executor_registration_id: response.executor_registration_id, + executor_public_key: response.executor_public_key, + harness_key_authorization: response.harness_key_authorization, + }) + } + + async fn resolve_auth_headers(&self) -> Result { + self.auth_provider + .resolve_auth_headers() + .await + .map_err(|error| { + ExecServerError::EnvironmentRegistryAuth(format!( + "failed to resolve environment registry authentication: {error}" + )) + }) + } + + async fn parse_json_response(&self, response: HttpResponse) -> Result + where + R: for<'de> Deserialize<'de>, + { + if response.status().is_success() { + let body = response + .text() + .await + .map_err(|error| ExecServerError::EnvironmentRegistryRequest(error.into()))?; + return serde_json::from_str(&body).map_err(ExecServerError::Json); + } + + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + if matches!(status, StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN) { + return Err(environment_registry_auth_error(status, &body)); + } + + Err(environment_registry_http_error(status, &body)) + } +} + +#[derive(Clone)] +struct RegistryHarnessKeyValidator { + client: EnvironmentRegistryClient, + environment_id: String, + executor_registration_id: String, +} + +impl HarnessKeyValidator for RegistryHarnessKeyValidator { + /// Authorize the harness key recovered from the first IK message. + /// Noise proves key possession; the registry decides whether that key may use + /// this executor. The authorization token and public key are checked together. + #[tracing::instrument( + name = "codex.exec_server.remote.environment_registry.validate_harness_key", + skip_all, + fields( + otel.kind = "client", + otel.name = "codex.exec_server.remote.environment_registry.validate_harness_key", + environment_id = %self.environment_id, + executor_registration_id = %self.executor_registration_id, + ) + )] + async fn validate_harness_key( + &self, + harness_public_key: &NoiseChannelPublicKey, + authorization: &str, + ) -> Result<(), ExecServerError> { + let environment_id = &self.environment_id; + let url = endpoint_url( + &self.client.base_url, + &format!("/cloud/environment/{environment_id}/validate"), + ); + let body = EnvironmentRegistryHarnessKeyValidationRequest { + executor_registration_id: self.executor_registration_id.clone(), + harness_public_key: harness_public_key.clone(), + harness_key_authorization: authorization.to_string(), + }; + let response = self + .client + .http + .post(url) + .headers(self.client.resolve_auth_headers().await?) + .headers(current_trace_context_headers()) + .json(&body) + .send() + .await?; + let status = response.status(); + if !status.is_success() { + // The request contains the short-lived authorization. Do not include + // a response body that might echo it in logs or error chains. + if matches!(status, StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN) { + return Err(ExecServerError::EnvironmentRegistryAuth(format!( + "environment registry harness key validation authentication failed ({status})" + ))); + } + return Err(ExecServerError::EnvironmentRegistryHttp { + status, + code: None, + message: "environment registry harness key validation failed".to_string(), + }); + } + let response = response + .json::() + .await + .map_err(|error| ExecServerError::EnvironmentRegistryRequest(error.into()))?; + if !response.valid { + return Err(ExecServerError::Protocol( + "environment registry rejected Noise relay harness key".to_string(), + )); + } + Ok(()) + } +} + +/// Noise connection configuration for a Codex harness. +/// +/// Configuration stays inert until the effective outbound HTTP policy is known. +/// Its connection provider then holds the authenticated registry client so every +/// reconnect receives fresh URL and harness-key authorization material. +#[derive(Clone)] +pub(crate) struct NoiseRendezvousEnvironmentConfig { + base_url: String, + environment_id: String, + auth_provider: SharedAuthProvider, +} + +impl std::fmt::Debug for NoiseRendezvousEnvironmentConfig { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("NoiseRendezvousEnvironmentConfig") + .field("base_url", &"") + .field("environment_id", &self.environment_id) + .field("auth_provider", &"") + .finish() + } +} + +impl NoiseRendezvousEnvironmentConfig { + pub(crate) fn new( + base_url: String, + environment_id: String, + bearer_token: String, + chatgpt_account_id: Option, + ) -> Result { + let base_url = normalize_base_url(base_url)?; + let environment_id = normalize_environment_id(environment_id)?; + let auth_provider = static_bearer_auth_provider(bearer_token, chatgpt_account_id)?; + Ok(Self { + base_url, + environment_id, + auth_provider, + }) + } + + pub(crate) fn into_connect_provider( + self, + http_client_factory: HttpClientFactory, + ) -> Result, ExecServerError> { + let client = EnvironmentRegistryClient::new_with_telemetry( + self.base_url, + self.auth_provider, + ExecServerTelemetry::default(), + http_client_factory, + )?; + Ok(Arc::new(EnvironmentRegistryNoiseConnectProvider { + client, + environment_id: self.environment_id, + })) + } +} + +#[derive(Clone, Debug)] +struct EnvironmentRegistryNoiseConnectProvider { + client: EnvironmentRegistryClient, + environment_id: String, +} + +impl NoiseRendezvousConnectProvider for EnvironmentRegistryNoiseConnectProvider { + fn connect_bundle( + &self, + harness_public_key: NoiseChannelPublicKey, + ) -> futures::future::BoxFuture<'_, Result> { + async move { + self.client + .connect_environment(&self.environment_id, harness_public_key) + .await + } + .boxed() + } +} + +#[derive(Clone)] +struct StaticBearerAuthProvider { + authorization: HeaderValue, + chatgpt_account_id: Option, +} + +impl std::fmt::Debug for StaticBearerAuthProvider { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("StaticBearerAuthProvider") + .field("authorization", &"") + .field( + "chatgpt_account_id", + &self.chatgpt_account_id.as_ref().map(|_| ""), + ) + .finish() + } +} + +impl AuthProvider for StaticBearerAuthProvider { + fn add_auth_headers(&self, headers: &mut HeaderMap) { + headers.insert(http::header::AUTHORIZATION, self.authorization.clone()); + if let Some(chatgpt_account_id) = &self.chatgpt_account_id { + headers.insert( + HeaderName::from_static("chatgpt-account-id"), + chatgpt_account_id.clone(), + ); + } + } +} + +fn static_bearer_auth_provider( + bearer_token: String, + chatgpt_account_id: Option, +) -> Result { + let bearer_token = bearer_token.trim(); + if bearer_token.is_empty() { + return Err(ExecServerError::EnvironmentRegistryConfig( + "environment registry bearer token is required".to_string(), + )); + } + let authorization = + HeaderValue::try_from(format!("Bearer {bearer_token}")).map_err(|error| { + ExecServerError::EnvironmentRegistryConfig(format!( + "environment registry bearer token is not a valid HTTP header: {error}" + )) + })?; + let chatgpt_account_id = chatgpt_account_id + .as_deref() + .map(str::trim) + .filter(|account_id| !account_id.is_empty()) + .map(|account_id| { + HeaderValue::try_from(account_id).map_err(|error| { + ExecServerError::EnvironmentRegistryConfig(format!( + "ChatGPT account id is not a valid HTTP header: {error}" + )) + }) + }) + .transpose()?; + Ok(Arc::new(StaticBearerAuthProvider { + authorization, + chatgpt_account_id, + })) +} + +/// Configuration for registering an exec-server for remote use. +#[derive(Clone)] +pub struct RemoteEnvironmentConfig { + pub base_url: String, + pub environment_id: String, + pub name: String, + pub request_dispatch_mode: RequestDispatchMode, + transport: RemoteEnvironmentTransport, + auth_provider: SharedAuthProvider, + telemetry: ExecServerTelemetry, + http_client_factory: HttpClientFactory, +} + +impl std::fmt::Debug for RemoteEnvironmentConfig { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RemoteEnvironmentConfig") + .field("base_url", &self.base_url) + .field("environment_id", &self.environment_id) + .field("name", &self.name) + .field("request_dispatch_mode", &self.request_dispatch_mode) + .field("transport", &self.transport) + .field("auth_provider", &"") + .finish() + } +} + +impl RemoteEnvironmentConfig { + /// Creates a remote environment configuration using the default Noise transport. + pub fn new( + base_url: String, + environment_id: String, + auth_provider: SharedAuthProvider, + http_client_factory: HttpClientFactory, + ) -> Result { + Self::new_with_transport( + base_url, + environment_id, + RemoteEnvironmentTransport::Noise, + auth_provider, + http_client_factory, + ) + } + + /// Creates a remote environment configuration using an explicit transport. + pub fn new_with_transport( + base_url: String, + environment_id: String, + transport: RemoteEnvironmentTransport, + auth_provider: SharedAuthProvider, + http_client_factory: HttpClientFactory, + ) -> Result { + let environment_id = normalize_environment_id(environment_id)?; + Ok(Self { + base_url, + environment_id, + name: "codex-exec-server".to_string(), + request_dispatch_mode: RequestDispatchMode::Inline, + transport, + auth_provider, + telemetry: ExecServerTelemetry::default(), + http_client_factory, + }) + } + + pub fn with_telemetry(mut self, telemetry: ExecServerTelemetry) -> Self { + self.telemetry = telemetry; + self + } +} + +/// Register an exec-server for remote use and serve requests over its configured transport. +/// +/// In Noise mode, the executor identity is generated once per process and reused across +/// reconnects. The registration and rendezvous URL are also reused until +/// rendezvous rejects the URL, at which point the next attempt registers again. +/// The websocket carries cleartext routing metadata and encrypted payloads. +/// +/// Direct mode reuses its registration across reconnects. A WebSocket handshake +/// conflict refreshes the registration; other permanent client errors stop the runner. +pub async fn run_remote_environment( + config: RemoteEnvironmentConfig, + runtime_paths: ExecServerRuntimePaths, +) -> Result<(), ExecServerError> { + run_remote_environment_until_shutdown(config, runtime_paths, std::future::pending()).await +} + +/// Serve a remote environment until its owner requests graceful shutdown. +/// +/// Active sessions and their processes are drained before this function returns. +pub async fn run_remote_environment_until_shutdown( + config: RemoteEnvironmentConfig, + runtime_paths: ExecServerRuntimePaths, + shutdown: F, +) -> Result<(), ExecServerError> +where + F: std::future::Future, +{ + let processor = ConnectionProcessor::new_with_telemetry( + runtime_paths, + config.telemetry.clone(), + config.http_client_factory.clone(), + config.request_dispatch_mode, + ); + + let result = match config.transport { + RemoteEnvironmentTransport::Noise => { + run_remote_transport(config, shutdown, |config, client| { + run_remote_environment_connections(config, client, processor.clone()) + }) + .await + } + RemoteEnvironmentTransport::Direct => { + run_remote_transport(config, shutdown, |config, client| { + run_direct_environment(config, client, processor.clone()) + }) + .await + } + }; + processor.shutdown().await; + result +} + +/// Register a remote environment backed by an independently owned WebSocket executor. +pub async fn run_remote_environment_forward_until_shutdown( + config: RemoteEnvironmentConfig, + websocket_url: String, + shutdown: F, +) -> Result<(), ExecServerError> +where + F: std::future::Future, +{ + // Forwarder implements the Noise stream bridge. Direct forwarding needs a + // separate Direct-compatible bridge, so reject it rather than use the + // Noise-specific path. + // Remove this guard when a Direct-compatible forwarder is added. + if config.transport == RemoteEnvironmentTransport::Direct { + return Err(ExecServerError::EnvironmentRegistryConfig( + "direct exec-server transport does not support forwarding".to_string(), + )); + } + let forwarder = Forwarder::new( + websocket_url, + &config.http_client_factory, + config.telemetry.clone(), + )?; + run_remote_transport(config, shutdown, |config, client| { + run_remote_environment_connections(config, client, forwarder) + }) + .await +} + +async fn run_remote_transport( + config: RemoteEnvironmentConfig, + shutdown: F, + run_loop: R, +) -> Result<(), ExecServerError> +where + F: std::future::Future, + R: FnOnce(RemoteEnvironmentConfig, EnvironmentRegistryClient) -> T, + T: std::future::Future>, +{ + ensure_rustls_crypto_provider(); + let client = EnvironmentRegistryClient::new_with_telemetry( + config.base_url.clone(), + config.auth_provider.clone(), + config.telemetry.clone(), + config.http_client_factory.clone(), + )?; + let run = run_loop(config, client); + tokio::pin!(run, shutdown); + tokio::select! { + result = &mut run => result, + _ = &mut shutdown => Ok(()), + } +} + +async fn run_remote_environment_connections( + config: RemoteEnvironmentConfig, + client: EnvironmentRegistryClient, + handler: H, +) -> Result<(), ExecServerError> { + let identity = NoiseChannelIdentity::generate().map_err(|error| { + ExecServerError::Protocol(format!("failed to generate Noise relay identity: {error}")) + })?; + let mut backoff = Duration::from_secs(1); + let mut response = client + .register_environment_with_retry(&config.environment_id, &identity.public_key()) + .await?; + + loop { + match connect_rendezvous( + &response.url, + &config.telemetry, + &config.http_client_factory, + ) + .await + { + Ok(websocket) => { + backoff = Duration::from_secs(1); + let executor_registration_id = response.executor_registration_id.clone(); + info!( + noise_event = "rendezvous_connection", + noise_outcome = "ok", + "Noise executor connected to rendezvous" + ); + let disconnect_reason = run_multiplexed_environment( + websocket, + handler.clone(), + response.environment_id.clone(), + executor_registration_id.clone(), + identity.clone(), + RegistryHarnessKeyValidator { + client: client.clone(), + environment_id: config.environment_id.clone(), + executor_registration_id, + }, + ) + .await; + info!( + noise_event = "rendezvous_connection", + noise_outcome = "disconnected", + noise_reason = disconnect_reason.as_str(), + "Noise executor disconnected from rendezvous" + ); + config + .telemetry + .remote_reconnect(disconnect_reason.as_str()); + } + Err(error) => { + let registration_rejected = matches!( + &error, + tokio_tungstenite::tungstenite::Error::Http(response) + if response.status().is_client_error() + ); + warn!( + noise_event = "rendezvous_connection", + noise_outcome = "error", + noise_reason = "websocket_error", + "Noise executor failed to connect to rendezvous" + ); + debug!(error = %error, "Noise executor rendezvous connection error"); + if registration_rejected { + config.telemetry.remote_reconnect("registration_rejected"); + response = client + .register_environment_with_retry( + &config.environment_id, + &identity.public_key(), + ) + .await?; + } else { + config.telemetry.remote_reconnect("connect_failed"); + } + } + } + + sleep(backoff).await; + backoff = (backoff * 2).min(Duration::from_secs(30)); + } +} + +#[tracing::instrument( + name = "codex.exec_server.remote.rendezvous.connect", + skip_all, + fields( + otel.kind = "client", + otel.name = "codex.exec_server.remote.rendezvous.connect", + result = tracing::field::Empty, + ) +)] +async fn connect_rendezvous( + url: &str, + telemetry: &ExecServerTelemetry, + http_client_factory: &HttpClientFactory, +) -> Result { + let started_at = Instant::now(); + let result = async { + let mut request = url.into_client_request()?; + request.headers_mut().extend(current_rendezvous_headers()); + let connector = WebSocketConnector::new_with_tls_mode( + http_client_factory, + WebSocketTlsMode::TungsteniteDefault, + ) + .map_err(|error| tokio_tungstenite::tungstenite::Error::Io(std::io::Error::other(error)))?; + connector + .with_tcp_nodelay() + .connect(request, noise_relay_websocket_config()) + .await + .map(|(websocket, _)| websocket) + } + .await; + let result_name = if result.is_ok() { "success" } else { "error" }; + tracing::Span::current().record("result", result_name); + telemetry.remote_rendezvous_completed(result_name, started_at.elapsed()); + result +} + +fn normalize_environment_id(environment_id: String) -> Result { + let environment_id = environment_id.trim().to_string(); + if environment_id.is_empty() { + return Err(ExecServerError::EnvironmentRegistryConfig( + "environment id is required for remote exec-server registration".to_string(), + )); + } + Ok(environment_id) +} + +#[derive(Deserialize)] +struct RegistryErrorBody { + error: Option, +} + +#[derive(Deserialize)] +struct RegistryError { + code: Option, + message: Option, +} + +fn normalize_base_url(base_url: String) -> Result { + let trimmed = base_url.trim().trim_end_matches('/').to_string(); + if trimmed.is_empty() { + return Err(ExecServerError::EnvironmentRegistryConfig( + "environment registry base URL is required".to_string(), + )); + } + Ok(trimmed) +} + +fn endpoint_url(base_url: &str, path: &str) -> String { + format!("{base_url}/{}", path.trim_start_matches('/')) +} + +fn environment_registry_auth_error(status: StatusCode, body: &str) -> ExecServerError { + let message = registry_error_message(body).unwrap_or_else(|| "empty error body".to_string()); + ExecServerError::EnvironmentRegistryAuth(format!( + "environment registry authentication failed ({status}): {message}" + )) +} + +fn environment_registry_http_error(status: StatusCode, body: &str) -> ExecServerError { + let parsed = serde_json::from_str::(body).ok(); + let (code, message) = parsed + .and_then(|body| body.error) + .map(|error| { + ( + error.code, + error.message.unwrap_or_else(|| { + preview_error_body(body).unwrap_or_else(|| "empty error body".to_string()) + }), + ) + }) + .unwrap_or_else(|| { + ( + None, + preview_error_body(body) + .unwrap_or_else(|| "empty or malformed error body".to_string()), + ) + }); + ExecServerError::EnvironmentRegistryHttp { + status, + code, + message, + } +} + +fn registry_error_message(body: &str) -> Option { + serde_json::from_str::(body) + .ok() + .and_then(|body| body.error) + .and_then(|error| error.message) + .or_else(|| preview_error_body(body)) +} + +fn preview_error_body(body: &str) -> Option { + let trimmed = body.trim(); + if trimmed.is_empty() { + return None; + } + Some(trimmed.chars().take(ERROR_BODY_PREVIEW_BYTES).collect()) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use codex_api::AuthProvider; + use codex_http_client::OutboundProxyPolicy; + use http::HeaderMap; + use http::HeaderValue; + use opentelemetry::trace::TracerProvider as _; + use opentelemetry_sdk::trace::SdkTracerProvider; + use pretty_assertions::assert_eq; + use tokio::io::AsyncReadExt; + use tokio::io::AsyncWriteExt; + use tokio::net::TcpListener; + use tracing::Instrument; + use tracing_subscriber::prelude::*; + use wiremock::Mock; + use wiremock::MockServer; + use wiremock::ResponseTemplate; + use wiremock::matchers::body_partial_json; + use wiremock::matchers::header; + use wiremock::matchers::header_regex; + use wiremock::matchers::method; + use wiremock::matchers::path; + + use super::*; + + #[derive(Debug)] + struct StaticRegistryAuthProvider; + + impl AuthProvider for StaticRegistryAuthProvider { + fn add_auth_headers(&self, _headers: &mut HeaderMap) {} + + fn resolve_auth_headers(&self) -> codex_api::AuthHeadersFuture<'_> { + Box::pin(async { + let mut headers = HeaderMap::new(); + let _ = headers.insert( + http::header::AUTHORIZATION, + HeaderValue::from_static("Bearer registry-token"), + ); + let _ = headers.insert( + "ChatGPT-Account-ID", + HeaderValue::from_static("workspace-123"), + ); + Ok(headers) + }) + } + } + + fn static_registry_auth_provider() -> SharedAuthProvider { + Arc::new(StaticRegistryAuthProvider) + } + + #[tokio::test(flavor = "current_thread")] + async fn register_environment_posts_with_auth_provider_headers() { + let provider = SdkTracerProvider::builder().build(); + let tracer = provider.tracer("exec-server-test"); + let subscriber = + tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer)); + let _guard = subscriber.set_default(); + tracing::callsite::rebuild_interest_cache(); + let server = MockServer::start().await; + let executor_public_key = NoiseChannelIdentity::generate() + .expect("identity") + .public_key(); + Mock::given(method("POST")) + .and(path("/cloud/environment/environment-requested/register")) + .and(header("authorization", "Bearer registry-token")) + .and(header("chatgpt-account-id", "workspace-123")) + .and(header_regex( + "traceparent", + "^00-[0-9a-f]{32}-[0-9a-f]{16}-0[01]$", + )) + .and(body_partial_json(serde_json::json!({ + "security_profile": NOISE_RELAY_SECURITY_PROFILE, + "executor_public_key": executor_public_key.clone(), + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "environment_id": "environment-requested", + "url": "wss://rendezvous.test/cloud-agent/default/ws/environment/environment-requested?role=environment&sig=abc", + "security_profile": NOISE_RELAY_SECURITY_PROFILE, + "executor_registration_id": "registration-1", + }))) + .mount(&server) + .await; + let client = EnvironmentRegistryClient::new(server.uri(), static_registry_auth_provider()) + .expect("client"); + + let response = client + .register_environment("environment-requested", &executor_public_key) + .instrument(tracing::info_span!("remote-operation")) + .await + .expect("register environment"); + + assert_eq!( + response, + EnvironmentRegistryRegistrationResponse { + environment_id: "environment-requested".to_string(), + url: "wss://rendezvous.test/cloud-agent/default/ws/environment/environment-requested?role=environment&sig=abc".to_string(), + security_profile: NOISE_RELAY_SECURITY_PROFILE.to_string(), + executor_registration_id: "registration-1".to_string(), + } + ); + } + + #[tokio::test] + async fn noise_connect_provider_requests_and_validates_a_full_bundle() { + let server = MockServer::start().await; + let harness_public_key = NoiseChannelIdentity::generate() + .expect("identity") + .public_key(); + let executor_public_key = NoiseChannelIdentity::generate() + .expect("identity") + .public_key(); + Mock::given(method("POST")) + .and(path("/cloud/environment/environment-requested/connect")) + .and(header("authorization", "Bearer registry-token")) + .and(header("chatgpt-account-id", "workspace-123")) + .and(body_partial_json(serde_json::json!({ + "harness_public_key": harness_public_key.clone(), + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "environment_id": "environment-requested", + "url": "wss://rendezvous.test/cloud-agent/default/ws/environment/environment-requested?role=harness&sig=abc", + "security_profile": NOISE_RELAY_SECURITY_PROFILE, + "executor_registration_id": "registration-1", + "executor_public_key": executor_public_key.clone(), + "harness_key_authorization": "authorization-1", + }))) + .mount(&server) + .await; + let config = NoiseRendezvousEnvironmentConfig::new( + server.uri(), + "environment-requested".to_string(), + "registry-token".to_string(), + Some("workspace-123".to_string()), + ) + .expect("noise configuration"); + + let bundle = config + .into_connect_provider(HttpClientFactory::new( + codex_http_client::OutboundProxyPolicy::ReqwestDefault, + )) + .expect("Noise connect provider") + .connect_bundle(harness_public_key) + .await + .expect("Noise connect bundle"); + + assert_eq!( + bundle.websocket_url, + "wss://rendezvous.test/cloud-agent/default/ws/environment/environment-requested?role=harness&sig=abc" + ); + assert_eq!(bundle.environment_id, "environment-requested"); + assert_eq!(bundle.executor_registration_id, "registration-1"); + assert_eq!(bundle.executor_public_key, executor_public_key); + assert_eq!(bundle.harness_key_authorization, "authorization-1"); + } + + #[tokio::test] + async fn connect_environment_times_out_when_registry_stalls() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/cloud/environment/environment-requested/connect")) + .respond_with(ResponseTemplate::new(200).set_delay(Duration::from_secs(1))) + .mount(&server) + .await; + let mut client = + EnvironmentRegistryClient::new(server.uri(), static_registry_auth_provider()) + .expect("client"); + client.connect_timeout = Duration::from_millis(50); + let harness_public_key = NoiseChannelIdentity::generate() + .expect("identity") + .public_key(); + + let error = match client + .connect_environment("environment-requested", harness_public_key) + .await + { + Ok(_) => panic!("stalled connect response should time out"), + Err(error) => error, + }; + + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryRequest(error) if error.is_timeout() + )); + } + + #[tokio::test] + async fn connect_environment_times_out_when_registry_response_body_stalls() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("registry listener should bind"); + let registry_url = format!( + "http://{}", + listener + .local_addr() + .expect("registry listener should have an address") + ); + tokio::spawn(async move { + let (mut stream, _) = listener + .accept() + .await + .expect("registry request should connect"); + stream + .write_all( + b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 256\r\n\r\n{", + ) + .await + .expect("registry response headers should write"); + sleep(Duration::from_secs(1)).await; + }); + let mut client = + EnvironmentRegistryClient::new(registry_url, static_registry_auth_provider()) + .expect("client"); + client.connect_timeout = Duration::from_millis(50); + let harness_public_key = NoiseChannelIdentity::generate() + .expect("identity") + .public_key(); + + let error = match client + .connect_environment("environment-requested", harness_public_key) + .await + { + Ok(_) => panic!("stalled connect response body should time out"), + Err(error) => error, + }; + + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryRequest(error) if error.is_timeout() + )); + } + + #[tokio::test] + async fn connect_environment_retries_interrupted_registry_response_bodies() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("registry listener should bind"); + let registry_url = format!( + "http://{}", + listener + .local_addr() + .expect("registry listener should have an address") + ); + tokio::spawn(async move { + let (mut stream, _) = listener + .accept() + .await + .expect("registry request should connect"); + let mut request = [0_u8; 4096]; + let bytes_read = stream + .read(&mut request) + .await + .expect("registry request should arrive before the response"); + assert_ne!(bytes_read, 0, "registry request should not be empty"); + stream + .write_all( + b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 256\r\n\r\n{", + ) + .await + .expect("registry response headers should write"); + stream + .shutdown() + .await + .expect("registry connection should close"); + }); + let client = EnvironmentRegistryClient::new(registry_url, static_registry_auth_provider()) + .expect("client"); + let harness_public_key = NoiseChannelIdentity::generate() + .expect("identity") + .public_key(); + + let error = client + .connect_environment("environment-requested", harness_public_key) + .await + .err() + .expect("interrupted response body must fail"); + + assert!( + crate::client::is_retryable_registry_error(&error), + "interrupted registry response body should be retryable: {error:?}" + ); + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryRequest(_) + )); + } + + #[tokio::test] + async fn connect_environment_does_not_retry_malformed_successful_responses() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/cloud/environment/environment-requested/connect")) + .respond_with(ResponseTemplate::new(200).set_body_string("{")) + .mount(&server) + .await; + let client = EnvironmentRegistryClient::new(server.uri(), static_registry_auth_provider()) + .expect("client"); + let harness_public_key = NoiseChannelIdentity::generate() + .expect("identity") + .public_key(); + + let error = client + .connect_environment("environment-requested", harness_public_key) + .await + .err() + .expect("malformed response must fail"); + + assert!(!crate::client::is_retryable_registry_error(&error)); + assert!(matches!(error, ExecServerError::Json(_))); + } + + #[tokio::test] + async fn register_environment_does_not_follow_redirects_with_auth_headers() { + let server = MockServer::start().await; + let executor_public_key = NoiseChannelIdentity::generate() + .expect("identity") + .public_key(); + Mock::given(method("POST")) + .and(path("/cloud/environment/environment-requested/register")) + .and(header("authorization", "Bearer registry-token")) + .respond_with( + ResponseTemplate::new(302) + .insert_header("location", format!("{}/redirect-target", server.uri())), + ) + .mount(&server) + .await; + Mock::given(path("/redirect-target")) + .and(header("authorization", "Bearer registry-token")) + .respond_with(ResponseTemplate::new(200)) + .expect(0) + .mount(&server) + .await; + let client = EnvironmentRegistryClient::new(server.uri(), static_registry_auth_provider()) + .expect("client"); + + let error = client + .register_environment("environment-requested", &executor_public_key) + .await + .expect_err("redirect response should not be followed"); + + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryHttp { + status: StatusCode::FOUND, + .. + } + )); + } + + #[test] + fn remote_environment_config_preserves_http_client_factory_policy() { + let config = RemoteEnvironmentConfig::new( + "https://registry.example".to_string(), + "env-1".to_string(), + static_registry_auth_provider(), + HttpClientFactory::new(OutboundProxyPolicy::RespectSystemProxy), + ) + .expect("config"); + + assert_eq!( + config.http_client_factory.outbound_proxy_policy(), + OutboundProxyPolicy::RespectSystemProxy + ); + } + + #[test] + fn remote_environment_config_new_defaults_to_noise_transport() { + let config = RemoteEnvironmentConfig::new( + "https://registry.example".to_string(), + "env-1".to_string(), + static_registry_auth_provider(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .expect("config"); + + assert_eq!(config.transport, RemoteEnvironmentTransport::Noise); + } + + #[test] + fn remote_environment_config_new_with_transport_preserves_direct_transport() { + let config = RemoteEnvironmentConfig::new_with_transport( + "https://registry.example".to_string(), + "env-1".to_string(), + RemoteEnvironmentTransport::Direct, + static_registry_auth_provider(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .expect("config"); + + assert_eq!(config.transport, RemoteEnvironmentTransport::Direct); + } + + #[test] + fn debug_output_redacts_auth_provider() { + let config = RemoteEnvironmentConfig::new( + "https://registry.example".to_string(), + "env-1".to_string(), + static_registry_auth_provider(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .expect("config"); + + let debug = format!("{config:?}"); + + assert!(debug.contains("")); + assert!(!debug.contains("workspace-123")); + } +} + +#[cfg(test)] +#[path = "remote/noise_tests.rs"] +mod noise_tests; diff --git a/codex-rs/exec-server/src/remote/direct.rs b/codex-rs/exec-server/src/remote/direct.rs new file mode 100644 index 0000000000000000000000000000000000000000..02365f58bb111a79605eeebe5f3c75c6f6e19707 --- /dev/null +++ b/codex-rs/exec-server/src/remote/direct.rs @@ -0,0 +1,277 @@ +use std::time::Instant; + +use codex_api::AuthError; +use codex_api::AuthProvider; +use codex_http_client::Request; +use codex_websocket_client::WebSocketConnection; +use codex_websocket_client::WebSocketConnector; +use http::Method; +use http::StatusCode; +use serde::Deserialize; +use serde::Serialize; +use tokio::time::sleep; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tracing::info; +use tracing::warn; + +use super::EnvironmentRegistryClient; +use super::RemoteEnvironmentConfig; +use crate::ExecServerError; +use crate::client::is_retryable_recovery_error; +use crate::client::registry_recovery_retry_delay; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::client_transport::authenticate_websocket_request; +use crate::client_transport::connect_websocket_request; +use crate::connection::JsonRpcConnection; +use crate::server::ConnectionProcessor; +use crate::telemetry::ConnectionTransport; +use crate::trace_context::current_trace_context_headers; + +const DIRECT_TRANSPORT: &str = "direct_jsonrpc_v1"; + +#[derive(Debug, Serialize)] +struct DirectRegistrationRequest { + transport: &'static str, +} + +#[derive(Debug, Deserialize)] +struct DirectRegistrationResponse { + environment_id: String, + transport: String, + registration_id: String, + url: String, +} + +impl EnvironmentRegistryClient { + #[tracing::instrument( + name = "codex.exec_server.remote.register", + skip_all, + fields( + otel.kind = "client", + otel.name = "codex.exec_server.remote.register", + result = tracing::field::Empty, + ) + )] + async fn register_direct_environment( + &self, + environment_id: &str, + ) -> Result { + let started_at = Instant::now(); + let response = async { + let url = super::endpoint_url( + &self.base_url, + &format!("/cloud/environment/{environment_id}/direct/register"), + ); + let request = Request::new(Method::POST, url) + .with_json(&DirectRegistrationRequest { + transport: DIRECT_TRANSPORT, + }) + .into_prepared() + .map_err(ExecServerError::EnvironmentRegistryConfig)?; + let request = self + .auth_provider + .apply_auth(request) + .await + .map_err(direct_auth_error)?; + let prepared = request + .prepare_body_for_send() + .map_err(ExecServerError::EnvironmentRegistryConfig)?; + let response = self + .http + .request(request.method, request.url) + .headers(prepared.headers) + .headers(current_trace_context_headers()) + .body(prepared.body.unwrap_or_default()) + .timeout(self.connect_timeout) + .send() + .await?; + let response: DirectRegistrationResponse = self.parse_json_response(response).await?; + if response.environment_id != environment_id { + return Err(ExecServerError::Protocol( + "environment registry returned a different environment id".to_string(), + )); + } + if response.transport != DIRECT_TRANSPORT { + return Err(ExecServerError::Protocol(format!( + "environment registry returned unsupported direct transport `{}`", + response.transport + ))); + } + if response.registration_id.trim().is_empty() || response.url.trim().is_empty() { + return Err(ExecServerError::Protocol( + "environment registry returned incomplete direct connection data".to_string(), + )); + } + require_tls_or_loopback(&response.url, "wss")?; + Ok(response) + } + .await; + let result = if response.is_ok() { "success" } else { "error" }; + tracing::Span::current().record("result", result); + self.telemetry + .remote_registration_completed(result, started_at.elapsed()); + response + } +} + +pub(super) async fn run_direct_environment( + config: RemoteEnvironmentConfig, + client: EnvironmentRegistryClient, + processor: ConnectionProcessor, +) -> Result<(), ExecServerError> { + require_tls_or_loopback(&config.base_url, "https")?; + let mut retry_attempt = 0; + let mut registration = client + .register_direct_environment(&config.environment_id) + .await?; + + loop { + match connect_direct( + ®istration.url, + config.auth_provider.as_ref(), + &config.http_client_factory, + ) + .await + { + Ok(websocket) => { + retry_attempt = 0; + info!( + environment_id = registration.environment_id, + registration_id = registration.registration_id, + "direct exec-server connected" + ); + processor + .run_connection( + JsonRpcConnection::from_websocket( + websocket, + format!( + "direct exec-server websocket {}", + websocket_diagnostic_url(®istration.url) + ), + ), + ConnectionTransport::WebSocket, + ) + .await; + config.telemetry.remote_reconnect("disconnected"); + } + Err(error) + if is_retryable_recovery_error(&error) + && !matches!( + &error, + ExecServerError::WebSocketConnect { + source: tokio_tungstenite::tungstenite::Error::Http(response), + .. + } if response.status().is_client_error() + && !matches!( + response.status(), + StatusCode::REQUEST_TIMEOUT + | StatusCode::CONFLICT + | StatusCode::TOO_MANY_REQUESTS + ) + ) => + { + // A handshake conflict rejects this registration; other transient failures + // reconnect using the existing URL. + if matches!( + &error, + ExecServerError::WebSocketConnect { + source: tokio_tungstenite::tungstenite::Error::Http(response), + .. + } if response.status() == StatusCode::CONFLICT + ) { + registration = client + .register_direct_environment(&config.environment_id) + .await?; + } + warn!("direct exec-server connection failed; retrying"); + config.telemetry.remote_reconnect("connect_failed"); + } + Err(error) => return Err(error), + } + + sleep(registry_recovery_retry_delay( + &config.environment_id, + retry_attempt, + )) + .await; + retry_attempt = retry_attempt.saturating_add(1); + } +} + +async fn connect_direct( + url: &str, + auth_provider: &dyn AuthProvider, + http_client_factory: &codex_http_client::HttpClientFactory, +) -> Result { + let mut request = + url.into_client_request() + .map_err(|source| ExecServerError::WebSocketConnect { + url: websocket_diagnostic_url(url), + source, + })?; + request + .headers_mut() + .extend(current_trace_context_headers()); + authenticate_websocket_request(&mut request, auth_provider) + .await + .map_err(direct_auth_error)?; + let authenticated_url = request.uri().to_string(); + require_tls_or_loopback(&authenticated_url, "wss")?; + let connector = WebSocketConnector::new(http_client_factory) + .map_err(|error| ExecServerError::WebSocketConfiguration(error.to_string()))? + .with_tcp_nodelay(); + connect_websocket_request( + request, + websocket_diagnostic_url(&authenticated_url), + connector, + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + /*use_loopback_direct*/ false, + ) + .await +} + +pub(super) fn require_tls_or_loopback( + url: &str, + secure_scheme: &str, +) -> Result<(), ExecServerError> { + let parsed = url::Url::parse(url).map_err(|error| { + ExecServerError::EnvironmentRegistryConfig(format!("invalid remote endpoint URL: {error}")) + })?; + if parsed.scheme() == secure_scheme { + return Ok(()); + } + + let loopback = match parsed.host() { + Some(url::Host::Domain(host)) => host.eq_ignore_ascii_case("localhost"), + Some(url::Host::Ipv4(host)) => host.is_loopback(), + Some(url::Host::Ipv6(host)) => host.is_loopback(), + None => false, + }; + if loopback + && secure_scheme + .strip_suffix('s') + .is_some_and(|scheme| parsed.scheme() == scheme) + { + return Ok(()); + } + + Err(ExecServerError::EnvironmentRegistryConfig(format!( + "remote transport requires {secure_scheme} for non-loopback endpoints" + ))) +} + +fn websocket_diagnostic_url(url: &str) -> String { + url.split(['?', '#']).next().unwrap_or(url).to_string() +} + +fn direct_auth_error(error: AuthError) -> ExecServerError { + let message = format!("failed to resolve environment registry authentication: {error}"); + match error { + AuthError::Build(_) => ExecServerError::EnvironmentRegistryAuth(message), + AuthError::Transient(_) => ExecServerError::Disconnected(message), + } +} + +#[cfg(test)] +#[path = "direct_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/remote/direct_tests.rs b/codex-rs/exec-server/src/remote/direct_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..d0cbff0d08490dd6cd40d6d2615d9b5a9d716e4a --- /dev/null +++ b/codex-rs/exec-server/src/remote/direct_tests.rs @@ -0,0 +1,449 @@ +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; +use std::time::Instant; + +use anyhow::Result; +use codex_api::AuthProvider; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use http::HeaderMap; +use http::HeaderValue; +use pretty_assertions::assert_eq; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::time::timeout; +use tokio_tungstenite::accept_hdr_async; +use tokio_tungstenite::tungstenite::handshake::server::Request as HandshakeRequest; +use tokio_tungstenite::tungstenite::handshake::server::Response as HandshakeResponse; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use super::*; +use crate::ExecServerRuntimePaths; +use crate::RemoteEnvironmentTransport; + +#[derive(Debug)] +struct StaticAuthProvider; + +impl AuthProvider for StaticAuthProvider { + fn add_auth_headers(&self, headers: &mut HeaderMap) { + headers.insert( + http::header::AUTHORIZATION, + HeaderValue::from_static("AWS4-HMAC-SHA256 test-signature"), + ); + } + + fn apply_auth( + &self, + mut request: codex_http_client::Request, + ) -> codex_api::AuthProviderFuture<'_> { + Box::pin(async move { + if request.method == http::Method::GET { + assert_eq!(request.headers.len(), 1); + assert!(request.headers.contains_key(http::header::HOST)); + assert!(request.url.starts_with("http")); + } + self.add_auth_headers(&mut request.headers); + Ok(request) + }) + } +} + +#[derive(Debug)] +struct QueryAuthProvider; + +impl AuthProvider for QueryAuthProvider { + fn add_auth_headers(&self, _headers: &mut HeaderMap) {} + + fn apply_auth( + &self, + mut request: codex_http_client::Request, + ) -> codex_api::AuthProviderFuture<'_> { + Box::pin(async move { + let mut url = url::Url::parse(&request.url) + .map_err(|error| codex_api::AuthError::Build(error.to_string()))?; + url.query_pairs_mut().append_pair("auth", "signed-query"); + request.url = url.into(); + Ok(request) + }) + } +} + +#[tokio::test] +async fn direct_websocket_signs_only_host_and_preserves_handshake_headers() -> Result<()> { + let mut request = "wss://executor.example.com/connect".into_client_request()?; + request + .headers_mut() + .insert("traceparent", HeaderValue::from_static("trace-context")); + + authenticate_websocket_request(&mut request, &StaticAuthProvider).await?; + + assert!(request.headers().contains_key("sec-websocket-key")); + assert_eq!(request.headers()["traceparent"], "trace-context"); + assert_eq!( + request.headers()[http::header::AUTHORIZATION], + "AWS4-HMAC-SHA256 test-signature" + ); + Ok(()) +} + +#[tokio::test] +async fn direct_websocket_uses_authenticated_url_query_parameters() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let url = format!("ws://{}/connect?existing=value", listener.local_addr()?); + let acceptor = tokio::spawn(async move { + let (socket, _) = listener.accept().await?; + let callback = |request: &HandshakeRequest, response: HandshakeResponse| { + assert_eq!(request.uri().path(), "/connect"); + assert_eq!( + request.uri().query(), + Some("existing=value&auth=signed-query") + ); + Ok(response) + }; + accept_hdr_async(socket, callback).await.map(|_| ()) + }); + + let connection = connect_direct( + &url, + &QueryAuthProvider, + &HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .await?; + drop(connection); + acceptor.await??; + Ok(()) +} + +#[test] +fn direct_endpoints_require_tls_except_on_loopback() { + for (url, scheme, allowed) in [ + ("https://registry.example.com", "https", true), + ("wss://executor.example.com/connect", "wss", true), + ("http://127.0.0.1:8080", "https", true), + ("ws://localhost:8080/connect", "wss", true), + ("http://registry.example.com", "https", false), + ("ws://executor.example.com/connect", "wss", false), + ] { + assert_eq!( + require_tls_or_loopback(url, scheme).is_ok(), + allowed, + "{url}" + ); + } +} + +#[test] +fn direct_authentication_retries_only_transient_credential_failures() { + for (error, retryable) in [ + (AuthError::Transient("temporary".to_string()), true), + (AuthError::Build("invalid".to_string()), false), + ] { + assert_eq!( + is_retryable_recovery_error(&direct_auth_error(error)), + retryable + ); + } +} + +#[tokio::test] +async fn direct_registration_validates_connection_data() -> Result<()> { + for (field, invalid_value) in [ + ("environment_id", "different-environment"), + ("transport", "noise_hybrid_ik_v1"), + ("registration_id", ""), + ("url", ""), + ("url", "ws://executor.example.com/connect"), + ] { + let registry = MockServer::start().await; + let mut response = serde_json::json!({ + "environment_id": "environment-requested", + "transport": DIRECT_TRANSPORT, + "registration_id": "registration-1", + "url": "wss://executor.example.com/connect", + }); + response[field] = serde_json::json!(invalid_value); + Mock::given(method("POST")) + .and(path( + "/cloud/environment/environment-requested/direct/register", + )) + .respond_with(ResponseTemplate::new(200).set_body_json(response)) + .expect(1) + .mount(®istry) + .await; + let client = EnvironmentRegistryClient::new(registry.uri(), Arc::new(StaticAuthProvider))?; + let error = client + .register_direct_environment("environment-requested") + .await + .expect_err("invalid registration must be rejected"); + if invalid_value.starts_with("ws://") { + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryConfig(_) + )); + } else { + assert!(matches!(error, ExecServerError::Protocol(_))); + } + registry.verify().await; + } + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn direct_registration_uses_proxy_policy_without_logging_secrets() -> Result<()> { + use tracing_subscriber::Layer; + use tracing_subscriber::layer::SubscriberExt; + + let mut log_file = tempfile::tempfile()?; + let writer = log_file.try_clone()?; + let subscriber = tracing_subscriber::registry().with( + tracing_subscriber::fmt::layer() + .with_ansi(false) + .with_writer(move || writer.try_clone().expect("clone log file")) + .with_filter( + tracing_subscriber::filter::Targets::new() + .with_target("codex_http_client", tracing::Level::TRACE) + .with_target("codex_exec_server", tracing::Level::TRACE), + ), + ); + let _guard = tracing::subscriber::set_default(subscriber); + tracing::debug!(target: "codex_exec_server", "direct registry log capture sentinel"); + + let proxy = MockServer::start().await; + let registry_url = "http://direct-registry-proxy.invalid/registry-path-secret"; + let request_url = + format!("{registry_url}/cloud/environment/environment-requested/direct/register"); + codex_http_client::cache_system_proxy_route_for_test(&request_url, proxy.uri()); + Mock::given(method("POST")) + .and(path( + "/registry-path-secret/cloud/environment/environment-requested/direct/register", + )) + .and(header("authorization", "AWS4-HMAC-SHA256 test-signature")) + .respond_with( + ResponseTemplate::new(200) + .insert_header("set-cookie", "session=registry-cookie-secret") + .insert_header( + "location", + "https://registry.example/?token=registry-location-secret", + ) + .set_body_json(serde_json::json!({ + "environment_id": "environment-requested", + "transport": DIRECT_TRANSPORT, + "registration_id": "registration-1", + "url": "wss://executor.example/connect?token=websocket-query-secret", + })), + ) + .expect(1) + .mount(&proxy) + .await; + let client = EnvironmentRegistryClient::new_with_telemetry( + registry_url.to_string(), + Arc::new(StaticAuthProvider), + crate::ExecServerTelemetry::default(), + HttpClientFactory::new(OutboundProxyPolicy::RespectSystemProxy), + )?; + let response = client + .register_direct_environment("environment-requested") + .await?; + assert_eq!( + response.url, + "wss://executor.example/connect?token=websocket-query-secret" + ); + let requests = proxy + .received_requests() + .await + .expect("record proxy requests"); + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].url.as_str(), request_url); + assert_eq!(requests[0].body, br#"{"transport":"direct_jsonrpc_v1"}"#); + proxy.verify().await; + + std::io::Seek::rewind(&mut log_file)?; + let mut logs = String::new(); + std::io::Read::read_to_string(&mut log_file, &mut logs)?; + assert!(logs.contains("direct registry log capture sentinel")); + for secret in [ + "test-signature", + "registry-path-secret", + "registry-cookie-secret", + "registry-location-secret", + "websocket-query-secret", + ] { + assert!(!logs.contains(secret), "registry logs exposed {secret}"); + } + Ok(()) +} + +#[tokio::test] +async fn direct_registration_failure_stops_initial_and_conflict_attempts() -> Result<()> { + for successful_registrations in [0, 1] { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}/connect", listener.local_addr()?); + let registry = MockServer::start().await; + let registration_count = AtomicUsize::new(0); + Mock::given(method("POST")) + .and(path( + "/cloud/environment/environment-requested/direct/register", + )) + .respond_with(move |_: &wiremock::Request| { + let attempt = registration_count.fetch_add(1, Ordering::Relaxed); + if attempt < successful_registrations { + ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "environment_id": "environment-requested", + "transport": DIRECT_TRANSPORT, + "registration_id": "registration-1", + "url": websocket_url, + })) + } else { + ResponseTemplate::new(503) + } + }) + .expect((successful_registrations + 1) as u64) + .mount(®istry) + .await; + let config = RemoteEnvironmentConfig::new_with_transport( + registry.uri(), + "environment-requested".to_string(), + RemoteEnvironmentTransport::Direct, + Arc::new(StaticAuthProvider), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )?; + let runtime_paths = ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?; + let task = tokio::spawn(crate::run_remote_environment(config, runtime_paths)); + if successful_registrations == 1 { + let (mut socket, _) = timeout(Duration::from_secs(5), listener.accept()).await??; + let mut request = [0; 4096]; + let _ = socket.read(&mut request).await?; + socket + .write_all(b"HTTP/1.1 409 Conflict\r\nContent-Length: 0\r\n\r\n") + .await?; + socket.shutdown().await?; + } + let error = timeout(Duration::from_secs(5), task) + .await?? + .expect_err("registration failure should stop the runner"); + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryHttp { + status: StatusCode::SERVICE_UNAVAILABLE, + .. + } + )); + registry.verify().await; + } + Ok(()) +} + +#[tokio::test] +async fn direct_websocket_reuses_registration_and_stops_on_permanent_errors() -> Result<()> { + for (status, should_retry) in [ + (Some(StatusCode::BAD_REQUEST), false), + (Some(StatusCode::UNAUTHORIZED), false), + (Some(StatusCode::FORBIDDEN), false), + (Some(StatusCode::NOT_FOUND), false), + (Some(StatusCode::METHOD_NOT_ALLOWED), false), + (Some(StatusCode::GONE), false), + (Some(StatusCode::REQUEST_TIMEOUT), true), + (Some(StatusCode::CONFLICT), true), + (Some(StatusCode::TOO_MANY_REQUESTS), true), + (Some(StatusCode::INTERNAL_SERVER_ERROR), true), + (Some(StatusCode::SERVICE_UNAVAILABLE), true), + (None, true), // Close the TCP socket without sending an HTTP response. + ] { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let registry = MockServer::start().await; + let expected_calls = if status == Some(StatusCode::CONFLICT) { + 2 + } else { + 1 + }; + let registration_count = AtomicUsize::new(0); + Mock::given(method("POST")) + .and(path( + "/cloud/environment/environment-requested/direct/register", + )) + .and(header("authorization", "AWS4-HMAC-SHA256 test-signature")) + .respond_with(move |_: &wiremock::Request| { + let registration = registration_count.fetch_add(1, Ordering::Relaxed) + 1; + ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "environment_id": "environment-requested", + "transport": DIRECT_TRANSPORT, + "registration_id": format!("registration-{registration}"), + "url": format!("{websocket_url}/registration-{registration}"), + })) + }) + .expect(expected_calls) + .mount(®istry) + .await; + let config = RemoteEnvironmentConfig::new_with_transport( + registry.uri(), + "environment-requested".to_string(), + RemoteEnvironmentTransport::Direct, + Arc::new(StaticAuthProvider), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )?; + let runtime_paths = ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?; + let task = tokio::spawn(crate::run_remote_environment(config, runtime_paths)); + + let (mut socket, _) = timeout(Duration::from_secs(5), listener.accept()).await??; + let mut request = [0; 4096]; + let _ = socket.read(&mut request).await?; + if let Some(status) = status { + let response = format!( + "HTTP/1.1 {} {}\r\nContent-Length: 0\r\n\r\n", + status.as_u16(), + status.canonical_reason().unwrap_or_default() + ); + socket.write_all(response.as_bytes()).await?; + } + socket.shutdown().await?; + drop(socket); + + if should_retry { + let expected_path = format!("/registration-{expected_calls}"); + let check_registration = |request: &HandshakeRequest, response: HandshakeResponse| { + assert_eq!(request.uri().path(), expected_path); + Ok(response) + }; + let (socket, _) = timeout(Duration::from_secs(5), listener.accept()).await??; + let websocket = accept_hdr_async(socket, &check_registration).await?; + let retry_delay = + registry_recovery_retry_delay("environment-requested", /*attempt*/ 0); + let retry_started = Instant::now(); + drop(websocket); + + let (socket, _) = timeout(Duration::from_secs(5), listener.accept()).await??; + assert!(retry_started.elapsed() >= retry_delay); + let _websocket = accept_hdr_async(socket, &check_registration).await?; + registry.verify().await; + task.abort(); + let _ = task.await; + } else { + let error = timeout(Duration::from_secs(5), task) + .await?? + .expect_err("permanent WebSocket client error should be terminal"); + assert!(matches!( + error, + ExecServerError::WebSocketConnect { + source: tokio_tungstenite::tungstenite::Error::Http(response), + .. + } if Some(response.status()) == status + )); + } + } + Ok(()) +} diff --git a/codex-rs/exec-server/src/remote/noise_tests.rs b/codex-rs/exec-server/src/remote/noise_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..c59328227a5e1f87dc3027cb4bb42698a34eab48 --- /dev/null +++ b/codex-rs/exec-server/src/remote/noise_tests.rs @@ -0,0 +1,355 @@ +use std::io::Write; +use std::sync::Arc; +use std::sync::Mutex; +use std::time::Duration; + +use anyhow::Result; +use codex_api::AuthProvider; +use codex_api::SharedAuthProvider; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_http_client::cache_system_proxy_route_for_test; +use http::HeaderMap; +use http::HeaderValue; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::time::timeout; +use tokio_tungstenite::accept_async; +use tracing_subscriber::Layer; +use tracing_subscriber::layer::SubscriberExt; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::body_partial_json; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; +use wiremock::matchers::query_param; + +use super::*; + +const HARNESS_KEY_AUTHORIZATION: &str = "authorization-that-must-not-leak"; + +#[derive(Debug)] +struct StaticRegistryAuthProvider; + +impl AuthProvider for StaticRegistryAuthProvider { + fn add_auth_headers(&self, headers: &mut HeaderMap) { + let _ = headers.insert( + http::header::AUTHORIZATION, + HeaderValue::from_static("Bearer registry-token"), + ); + } +} + +fn static_registry_auth_provider() -> SharedAuthProvider { + Arc::new(StaticRegistryAuthProvider) +} + +#[tokio::test(flavor = "current_thread")] +async fn registry_requests_do_not_log_sensitive_urls_or_response_headers() -> Result<()> { + let log_buffer = Arc::new(Mutex::new(Vec::new())); + let writer_buffer = Arc::clone(&log_buffer); + let subscriber = tracing_subscriber::registry().with( + tracing_subscriber::fmt::layer() + .with_ansi(false) + .with_writer(move || RegistryLogWriter(Arc::clone(&writer_buffer))) + .with_filter( + tracing_subscriber::filter::Targets::new() + .with_target("codex_http_client", tracing::Level::TRACE) + .with_target("codex_exec_server", tracing::Level::TRACE), + ), + ); + let _guard = tracing::subscriber::set_default(subscriber); + tracing::debug!(target: "codex_exec_server", "registry log capture sentinel"); + + let server = MockServer::start().await; + let harness_public_key = NoiseChannelIdentity::generate()?.public_key(); + let executor_public_key = NoiseChannelIdentity::generate()?.public_key(); + for (operation, response, cookie_secret, location_secret) in [ + ( + "register", + serde_json::json!({ + "environment_id": "environment-requested", + "url": "wss://rendezvous.test/environment", + "security_profile": NOISE_RELAY_SECURITY_PROFILE, + "executor_registration_id": "registration-1", + }), + "register-cookie-secret", + "register-location-secret", + ), + ( + "connect", + serde_json::json!({ + "environment_id": "environment-requested", + "url": "wss://rendezvous.test/harness", + "security_profile": NOISE_RELAY_SECURITY_PROFILE, + "executor_registration_id": "registration-1", + "executor_public_key": executor_public_key.clone(), + "harness_key_authorization": HARNESS_KEY_AUTHORIZATION, + }), + "connect-cookie-secret", + "connect-location-secret", + ), + ( + "validate", + serde_json::json!({ "valid": true }), + "validate-cookie-secret", + "validate-location-secret", + ), + ] { + Mock::given(method("POST")) + .and(path("/registry-path-secret")) + .and(query_param( + "registry_token", + format!( + "registry-query-secret/cloud/environment/environment-requested/{operation}" + ), + )) + .respond_with( + ResponseTemplate::new(200) + .insert_header("set-cookie", format!("session={cookie_secret}")) + .insert_header( + "location", + format!("https://registry.example/private?token={location_secret}"), + ) + .set_body_json(response), + ) + .expect(1) + .mount(&server) + .await; + } + + let registry_url = + server + .uri() + .replacen("http://", "http://registry-user:registry-password@", 1); + let registry_url = + format!("{registry_url}/registry-path-secret?registry_token=registry-query-secret"); + let client = EnvironmentRegistryClient::new(registry_url, static_registry_auth_provider())?; + client + .register_environment("environment-requested", &executor_public_key) + .await?; + client + .connect_environment("environment-requested", harness_public_key.clone()) + .await?; + RegistryHarnessKeyValidator { + client, + environment_id: "environment-requested".to_string(), + executor_registration_id: "registration-1".to_string(), + } + .validate_harness_key(&harness_public_key, HARNESS_KEY_AUTHORIZATION) + .await?; + + let logs = String::from_utf8(log_buffer.lock().expect("log buffer lock").clone())?; + assert!(logs.contains("registry log capture sentinel")); + for secret in [ + "registry-user", + "registry-password", + "registry-path-secret", + "registry-query-secret", + "registry-token", + HARNESS_KEY_AUTHORIZATION, + "register-cookie-secret", + "register-location-secret", + "connect-cookie-secret", + "connect-location-secret", + "validate-cookie-secret", + "validate-location-secret", + ] { + assert!(!logs.contains(secret), "logs exposed {secret}:\n{logs}"); + } + + Ok(()) +} + +#[tokio::test] +async fn reconnect_reuses_registration_until_url_is_rejected() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let rendezvous_url = format!("ws://{}", listener.local_addr()?); + let registry = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/cloud/environment/environment-requested/register")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "environment_id": "environment-requested", + "url": rendezvous_url, + "security_profile": NOISE_RELAY_SECURITY_PROFILE, + "executor_registration_id": "registration-1", + }))) + .expect(2) + .mount(®istry) + .await; + let config = RemoteEnvironmentConfig::new( + registry.uri(), + "environment-requested".to_string(), + static_registry_auth_provider(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )?; + let environment_task = tokio::spawn(run_remote_environment( + config, + ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + )?, + )); + + let (first_socket, _peer_addr) = timeout(Duration::from_secs(5), listener.accept()).await??; + let mut first_websocket = accept_async(first_socket).await?; + first_websocket.close(None).await?; + + // An ordinary disconnect retries the same URL without registering again. + let (mut rejected_socket, _peer_addr) = + timeout(Duration::from_secs(5), listener.accept()).await??; + let mut request = [0u8; 4096]; + let _ = rejected_socket.read(&mut request).await?; + rejected_socket + .write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\n\r\n") + .await?; + rejected_socket.shutdown().await?; + + // The 4xx response discards the old registration before this attempt. + let (third_socket, _peer_addr) = timeout(Duration::from_secs(5), listener.accept()).await??; + let _third_websocket = accept_async(third_socket).await?; + registry.verify().await; + + environment_task.abort(); + let _ = environment_task.await; + Ok(()) +} + +#[tokio::test] +async fn noise_connect_provider_uses_supplied_system_proxy_policy() -> Result<()> { + let proxy = MockServer::start().await; + let registry_url = "http://registry-policy-proxy.test"; + let request_url = format!("{registry_url}/cloud/environment/environment-requested/connect"); + cache_system_proxy_route_for_test(&request_url, proxy.uri()); + + let harness_public_key = NoiseChannelIdentity::generate()?.public_key(); + let executor_public_key = NoiseChannelIdentity::generate()?.public_key(); + Mock::given(method("POST")) + .and(path("/cloud/environment/environment-requested/connect")) + .and(header("authorization", "Bearer registry-token")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "environment_id": "environment-requested", + "url": "wss://rendezvous.test/cloud-agent/default/ws/environment/environment-requested", + "security_profile": NOISE_RELAY_SECURITY_PROFILE, + "executor_registration_id": "registration-1", + "executor_public_key": executor_public_key.clone(), + "harness_key_authorization": HARNESS_KEY_AUTHORIZATION, + }))) + .expect(1) + .mount(&proxy) + .await; + + let provider = NoiseRendezvousEnvironmentConfig::new( + registry_url.to_string(), + "environment-requested".to_string(), + "registry-token".to_string(), + /*chatgpt_account_id*/ None, + )? + .into_connect_provider(HttpClientFactory::new( + OutboundProxyPolicy::RespectSystemProxy, + ))?; + let bundle = timeout( + Duration::from_secs(5), + provider.connect_bundle(harness_public_key), + ) + .await??; + let requests = proxy + .received_requests() + .await + .expect("proxy request recording should be enabled"); + + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].url.as_str(), request_url); + assert_eq!(bundle.executor_public_key, executor_public_key); + + Ok(()) +} + +#[tokio::test] +async fn validate_harness_key_requires_explicit_valid_response() { + let server = MockServer::start().await; + let harness_public_key = NoiseChannelIdentity::generate() + .expect("identity") + .public_key(); + Mock::given(method("POST")) + .and(path("/cloud/environment/environment-requested/validate")) + .and(header("authorization", "Bearer registry-token")) + .and(body_partial_json(serde_json::json!({ + "executor_registration_id": "registration-1", + "harness_public_key": harness_public_key.clone(), + "harness_key_authorization": HARNESS_KEY_AUTHORIZATION, + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "valid": false, + }))) + .mount(&server) + .await; + let client = EnvironmentRegistryClient::new(server.uri(), static_registry_auth_provider()) + .expect("client"); + + let error = RegistryHarnessKeyValidator { + client, + environment_id: "environment-requested".to_string(), + executor_registration_id: "registration-1".to_string(), + } + .validate_harness_key(&harness_public_key, HARNESS_KEY_AUTHORIZATION) + .await + .expect_err("a false validation response must fail closed"); + + assert!(matches!( + error, + ExecServerError::Protocol(message) + if message == "environment registry rejected Noise relay harness key" + )); +} + +#[tokio::test] +async fn validate_harness_key_does_not_expose_error_body() { + let server = MockServer::start().await; + let harness_public_key = NoiseChannelIdentity::generate() + .expect("identity") + .public_key(); + Mock::given(method("POST")) + .and(path("/cloud/environment/environment-requested/validate")) + .respond_with(ResponseTemplate::new(500).set_body_string(HARNESS_KEY_AUTHORIZATION)) + .mount(&server) + .await; + let client = EnvironmentRegistryClient::new(server.uri(), static_registry_auth_provider()) + .expect("client"); + + let error = RegistryHarnessKeyValidator { + client, + environment_id: "environment-requested".to_string(), + executor_registration_id: "registration-1".to_string(), + } + .validate_harness_key(&harness_public_key, HARNESS_KEY_AUTHORIZATION) + .await + .expect_err("validation HTTP error should fail closed"); + + let display = error.to_string(); + assert!(!display.contains(HARNESS_KEY_AUTHORIZATION)); + assert!(matches!( + error, + ExecServerError::EnvironmentRegistryHttp { message, .. } + if message == "environment registry harness key validation failed" + )); +} + +struct RegistryLogWriter(Arc>>); + +impl Write for RegistryLogWriter { + fn write(&mut self, buffer: &[u8]) -> std::io::Result { + self.0 + .lock() + .expect("log buffer lock") + .extend_from_slice(buffer); + Ok(buffer.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} diff --git a/codex-rs/exec-server/src/remote/registration_retry.rs b/codex-rs/exec-server/src/remote/registration_retry.rs new file mode 100644 index 0000000000000000000000000000000000000000..1a2dc989da947393a1191ea4200ab8346825925c --- /dev/null +++ b/codex-rs/exec-server/src/remote/registration_retry.rs @@ -0,0 +1,54 @@ +//! Retry only explicit registration conflicts after the registry's write-retry loop has finished. +//! Ambiguous failures must not be replayed: a timed-out request can still replace a newer registration. +//! The enclosing remote-transport future owns cancellation; retries spawn no background work. + +use http::StatusCode; +use tokio::time::sleep; +use tracing::warn; + +use super::EnvironmentRegistryClient; +use crate::EnvironmentRegistryRegistrationResponse; +use crate::ExecServerError; +use crate::NoiseChannelPublicKey; +use crate::client::registry_recovery_retry_delay; + +impl EnvironmentRegistryClient { + pub(super) async fn register_environment_with_retry( + &self, + environment_id: &str, + executor_public_key: &NoiseChannelPublicKey, + ) -> Result { + // Competing executors for the same environment must not retry in lockstep. + let retry_key = uuid::Uuid::new_v4().to_string(); + let mut attempt = 0_u32; + loop { + match self + .register_environment(environment_id, executor_public_key) + .await + { + Ok(response) => return Ok(response), + Err(ExecServerError::EnvironmentRegistryHttp { status, code, .. }) + if status == StatusCode::SERVICE_UNAVAILABLE + && code.as_deref() == Some("registration_conflict") => + { + let delay = registry_recovery_retry_delay(&retry_key, attempt); + attempt = attempt.saturating_add(1); + // Do not log response bodies or transport errors: they can contain credentials. + warn!( + noise_event = "registration", + noise_outcome = "retry", + retry_attempt = attempt, + retry_delay_ms = delay.as_millis() as u64, + "Noise executor retrying registry registration conflict" + ); + sleep(delay).await; + } + Err(error) => return Err(error), + } + } + } +} + +#[cfg(test)] +#[path = "registration_retry_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/remote/registration_retry_tests.rs b/codex-rs/exec-server/src/remote/registration_retry_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..2c280c494798e522cd4ef4fcb4bd2cd055092d60 --- /dev/null +++ b/codex-rs/exec-server/src/remote/registration_retry_tests.rs @@ -0,0 +1,178 @@ +//! Only confirmed conflicts permit replay; ambiguous registration results remain terminal. + +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_api::AuthProvider; +use codex_http_client::RouteAwareRequestError; +use http::HeaderMap; +use pretty_assertions::assert_eq; +use tokio::io::AsyncBufReadExt; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::io::BufReader; +use tokio::net::TcpListener; +use tokio::sync::Notify; +use tokio::time::advance; +use tokio::time::sleep; +use tokio::time::timeout; +use tokio_util::task::AbortOnDropHandle; +use tracing::instrument::WithSubscriber; +use tracing_subscriber::prelude::*; + +use super::EnvironmentRegistryClient; +use super::ExecServerError; +use crate::NoiseChannelIdentity; + +const ENVIRONMENT_ID: &str = "registration-retry-test"; +const CONFLICT_BODY: &str = + r#"{"error":{"code":"registration_conflict","message":"registration unavailable"}}"#; +const ERROR_BODY: &str = + r#"{"error":{"code":"registration_denied","message":"registration unavailable"}}"#; +const SUCCESS_BODY: &str = r#"{"environment_id":"registration-retry-test","url":"ws://localhost/relay","security_profile":"noise_hybrid_ik_v1","executor_registration_id":"committed-registration"}"#; + +#[derive(Debug, Default)] +struct RegistrationAuthProvider { + calls: AtomicUsize, +} + +impl AuthProvider for RegistrationAuthProvider { + fn add_auth_headers(&self, _headers: &mut HeaderMap) {} + + fn resolve_auth_headers(&self) -> codex_api::AuthHeadersFuture<'_> { + self.calls.fetch_add(1, Ordering::Relaxed); + Box::pin(async { Ok(HeaderMap::new()) }) + } +} + +struct RetryObserved(Arc); + +impl tracing_subscriber::Layer for RetryObserved { + fn on_event(&self, event: &tracing::Event<'_>, _: tracing_subscriber::layer::Context<'_, S>) { + if event.metadata().target() == "codex_exec_server::remote::registration_retry" + && event.metadata().fields().field("retry_attempt").is_some() + { + self.0.notify_one(); + } + } +} + +#[test_case::test_case(200, "not JSON", Duration::ZERO; "malformed_success")] +#[test_case::test_case(200, SUCCESS_BODY, Duration::from_secs(60); "committed_success_body_timeout")] +#[test_case::test_case(503, CONFLICT_BODY, Duration::ZERO; "confirmed_conflict_backoff_is_cancellable")] +#[test_case::test_case(503, CONFLICT_BODY, Duration::from_secs(60); "conflict_code_not_received")] +#[test_case::test_case(503, "not JSON", Duration::ZERO; "malformed_unavailable")] +#[test_case::test_case(503, "{}", Duration::ZERO; "missing_conflict_code")] +#[test_case::test_case(503, ERROR_BODY, Duration::ZERO; "different_error_code")] +#[test_case::test_case(502, CONFLICT_BODY, Duration::ZERO; "gateway_error")] +#[test_case::test_case(408, CONFLICT_BODY, Duration::ZERO; "request_timeout")] +#[test_case::test_case(429, CONFLICT_BODY, Duration::ZERO; "too_many_requests")] +#[test_case::test_case(401, ERROR_BODY, Duration::from_secs(60); "unauthorized_stalled_body")] +#[test_case::test_case(403, ERROR_BODY, Duration::from_secs(60); "forbidden_stalled_body")] +#[test_case::test_case(404, ERROR_BODY, Duration::from_secs(60); "environment_deleted_stalled_body")] +#[test_case::test_case(401, ERROR_BODY, Duration::from_millis(50); "delayed_unauthorized_details")] +#[test_case::test_case(403, ERROR_BODY, Duration::from_millis(50); "delayed_forbidden_details")] +#[test_case::test_case(404, ERROR_BODY, Duration::from_millis(50); "delayed_error_details")] +#[tokio::test] +async fn registration_requires_a_confirmed_conflict_before_replay( + status: u16, + body: &'static str, + body_delay: Duration, +) -> anyhow::Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let registry_url = format!("http://{}", listener.local_addr()?); + let _server = AbortOnDropHandle::new(tokio::spawn(async move { + let (mut stream, _) = listener.accept().await?; + // Drain the complete request before responding so unread bytes cannot cause a reset. + let mut reader = BufReader::new(&mut stream); + let mut content_length = 0; + loop { + let mut line = String::new(); + anyhow::ensure!(reader.read_line(&mut line).await? > 0, "missing headers"); + if line == "\r\n" { + break; + } + if let Some((name, value)) = line.split_once(':') + && name.eq_ignore_ascii_case("content-length") + { + content_length = value.trim().parse::()?; + } + } + anyhow::ensure!(content_length <= 16_384, "registration request too large"); + reader.read_exact(&mut vec![0; content_length]).await?; + drop(reader); + + let headers = format!( + "HTTP/1.1 {status} Registration Result\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n", + body.len() + ); + stream.write_all(headers.as_bytes()).await?; + sleep(body_delay).await; + stream.write_all(body.as_bytes()).await?; + std::future::pending::<()>().await; + anyhow::Ok(()) + })); + let auth = Arc::new(RegistrationAuthProvider::default()); + let mut client = EnvironmentRegistryClient::new(registry_url, auth.clone())?; + client.connect_timeout = Duration::from_millis(500); + let key = NoiseChannelIdentity::generate()?.public_key(); + let retry = Arc::new(Notify::new()); + let subscriber = tracing_subscriber::registry().with(RetryObserved(retry.clone())); + let mut registration = Box::pin( + client + .register_environment_with_retry(ENVIRONMENT_ID, &key) + .with_subscriber(subscriber), + ); + + if status == 503 && body == CONFLICT_BODY && body_delay.is_zero() { + timeout(Duration::from_secs(1), async { + tokio::select! { + _ = retry.notified() => anyhow::Ok(()), + _ = &mut registration => anyhow::bail!("a confirmed conflict must enter backoff"), + } + }) + .await??; + // The retry event fires in the same poll that enters backoff, without a timing guess. + drop(registration); + tokio::time::pause(); + advance(Duration::from_secs(60)).await; + tokio::task::yield_now().await; + assert_eq!(auth.calls.load(Ordering::Relaxed), 1); + return Ok(()); + } + + let error = timeout(Duration::from_secs(1), registration) + .await? + .expect_err("an unconfirmed registration outcome must not be replayed"); + assert_eq!(auth.calls.load(Ordering::Relaxed), 1); + + match error { + ExecServerError::Json(_) if status == 200 && body_delay.is_zero() => {} + ExecServerError::EnvironmentRegistryRequest(RouteAwareRequestError::Timeout) + if status == 200 && body_delay > client.connect_timeout => {} + ExecServerError::EnvironmentRegistryAuth(message) if matches!(status, 401 | 403) => { + if body_delay < client.connect_timeout { + assert!(message.ends_with(": registration unavailable")); + } + } + ExecServerError::EnvironmentRegistryHttp { + status: actual, + code, + message, + } if !matches!(status, 200 | 401 | 403) => { + let expected_code = match (body_delay < client.connect_timeout, body) { + (true, ERROR_BODY) => Some("registration_denied"), + (true, CONFLICT_BODY) => Some("registration_conflict"), + _ => None, + }; + assert_eq!((actual.as_u16(), code.as_deref()), (status, expected_code)); + if expected_code.is_some() { + assert_eq!(message, "registration unavailable"); + } + } + _ => anyhow::bail!("unexpected registration failure kind"), + } + Ok(()) +} diff --git a/codex-rs/exec-server/src/remote_file_stream.rs b/codex-rs/exec-server/src/remote_file_stream.rs new file mode 100644 index 0000000000000000000000000000000000000000..107dea51d7d551c67ccfaa680426b3ee805725b2 --- /dev/null +++ b/codex-rs/exec-server/src/remote_file_stream.rs @@ -0,0 +1,121 @@ +use bytes::Bytes; +use codex_utils_path_uri::PathUri; +use tokio::io; +use uuid::Uuid; + +use super::map_remote_error; +use crate::ExecServerClient; +use crate::FILE_READ_CHUNK_SIZE; +use crate::FileSystemReadStream; +use crate::FileSystemResult; +use crate::FileSystemSandboxContext; +use crate::protocol::FS_READ_BLOCK_METHOD; +use crate::protocol::FsCloseParams; +use crate::protocol::FsOpenParams; +use crate::protocol::FsReadBlockParams; + +struct FileReadRegistration { + client: ExecServerClient, + handle_id: String, + runtime: Option, + active: bool, +} + +pub(super) async fn open( + client: ExecServerClient, + path: PathUri, + sandbox: Option, +) -> FileSystemResult { + let registration = FileReadRegistration { + client, + handle_id: Uuid::new_v4().simple().to_string(), + runtime: tokio::runtime::Handle::try_current().ok(), + active: true, + }; + registration + .client + .fs_open(FsOpenParams { + handle_id: registration.handle_id.clone(), + path, + sandbox, + }) + .await + .map_err(map_remote_error)?; + Ok(FileSystemReadStream::new(futures::stream::try_unfold( + Some((registration, 0_u64)), + |state| async move { + let Some((mut registration, offset)) = state else { + return Ok(None); + }; + let response = registration + .client + .fs_read_block(FsReadBlockParams { + handle_id: registration.handle_id.clone(), + offset, + len: FILE_READ_CHUNK_SIZE, + }) + .await + .map_err(map_remote_error)?; + let chunk = Bytes::from(response.chunk.into_inner()); + if chunk.len() > FILE_READ_CHUNK_SIZE { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!( + "{FS_READ_BLOCK_METHOD} returned {} bytes, maximum is {}", + chunk.len(), + FILE_READ_CHUNK_SIZE + ), + )); + } + if response.eof { + if registration + .client + .fs_close(FsCloseParams { + handle_id: registration.handle_id.clone(), + }) + .await + .is_ok() + { + registration.active = false; + } + return if chunk.is_empty() { + Ok(None) + } else { + Ok(Some((chunk, None))) + }; + } + if chunk.is_empty() { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("{FS_READ_BLOCK_METHOD} returned an empty non-terminal block"), + )); + } + let next_offset = offset.checked_add(chunk.len() as u64).ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidData, + format!("{FS_READ_BLOCK_METHOD} offset overflowed after {offset} bytes"), + ) + })?; + Ok(Some((chunk, Some((registration, next_offset))))) + }, + ))) +} + +impl Drop for FileReadRegistration { + fn drop(&mut self) { + if !self.active { + return; + } + let client = self.client.clone(); + let handle_id = self.handle_id.clone(); + let runtime = self + .runtime + .clone() + .or_else(|| tokio::runtime::Handle::try_current().ok()); + if let Some(runtime) = runtime { + runtime.spawn(async move { + let _ = client.fs_close(FsCloseParams { handle_id }).await; + }); + } + } +} diff --git a/codex-rs/exec-server/src/remote_file_system.rs b/codex-rs/exec-server/src/remote_file_system.rs new file mode 100644 index 0000000000000000000000000000000000000000..d574cece8c7ebf187db99938000907d9fdba9197 --- /dev/null +++ b/codex-rs/exec-server/src/remote_file_system.rs @@ -0,0 +1,534 @@ +use std::collections::HashMap; +use std::sync::Arc; + +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD; +use codex_utils_path_uri::PathUri; +use tokio::io; +use tokio::sync::Mutex; +use tokio::sync::OnceCell; +use tracing::trace; + +use crate::CopyOptions; +use crate::CreateDirectoryOptions; +use crate::ExecServerError; +use crate::ExecutorFileSystem; +use crate::ExecutorFileSystemFuture; +use crate::FileMetadata; +use crate::FileSystemReadStream; +use crate::FileSystemResult; +use crate::FileSystemSandboxContext; +use crate::GetMetadataOptions; +use crate::ReadDirectoryEntry; +use crate::ReadFileOptions; +use crate::RemoveOptions; +use crate::WalkOptions; +use crate::WalkOutcome; +use crate::WriteFileOptions; +use crate::client::LazyRemoteExecServerClient; +use crate::protocol::FsCanonicalizeParams; +use crate::protocol::FsCopyParams; +use crate::protocol::FsCreateDirectoryParams; +use crate::protocol::FsGetMetadataParams; +use crate::protocol::FsReadDirectoryParams; +use crate::protocol::FsReadFileParams; +use crate::protocol::FsRemoveParams; +use crate::protocol::FsWalkParams; +use crate::protocol::FsWriteFileParams; + +const INVALID_REQUEST_ERROR_CODE: i64 = -32600; +const NOT_FOUND_ERROR_CODE: i64 = -32004; + +#[path = "remote_file_stream.rs"] +mod file_stream; + +type InFlightMetadataRequest = OnceCell>>; + +pub(crate) struct RemoteFileSystem { + client: LazyRemoteExecServerClient, + metadata_requests: Mutex>>, +} + +impl RemoteFileSystem { + pub(crate) fn new(client: LazyRemoteExecServerClient) -> Self { + trace!("remote fs new"); + Self { + client, + metadata_requests: Mutex::new(HashMap::new()), + } + } + + async fn canonicalize( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + trace!("remote fs canonicalize"); + let client = self.client.get().await.map_err(map_remote_error)?; + let response = client + .fs_canonicalize(FsCanonicalizeParams { + path: path.clone(), + sandbox: remote_sandbox_context(sandbox), + }) + .await + .map_err(map_remote_error)?; + Ok(response.path) + } + + async fn read_file( + &self, + path: &PathUri, + options: ReadFileOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult> { + trace!("remote fs read_file"); + let client = self.client.get().await.map_err(map_remote_error)?; + let response = client + .fs_read_file(FsReadFileParams { + path: path.clone(), + follow_symlinks: (!options.follow_symlinks).then_some(false), + sandbox: remote_sandbox_context(sandbox), + }) + .await + .map_err(map_remote_error)?; + STANDARD.decode(response.data_base64).map_err(|err| { + io::Error::new( + io::ErrorKind::InvalidData, + format!("remote fs/readFile returned invalid base64 dataBase64: {err}"), + ) + }) + } + + async fn read_file_stream( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + trace!("remote fs read_file_stream"); + let client = self.client.get().await.map_err(map_remote_error)?; + file_stream::open(client, path.clone(), remote_sandbox_context(sandbox)).await + } + + async fn write_file( + &self, + path: &PathUri, + contents: Vec, + options: WriteFileOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + trace!("remote fs write_file"); + let client = self.client.get().await.map_err(map_remote_error)?; + let result = client + .fs_write_file(FsWriteFileParams { + path: path.clone(), + data_base64: STANDARD.encode(contents), + follow_symlinks: (!options.follow_symlinks).then_some(false), + sandbox: remote_sandbox_context(sandbox), + }) + .await; + self.metadata_requests.lock().await.clear(); + result.map_err(map_remote_error)?; + Ok(()) + } + + async fn create_directory( + &self, + path: &PathUri, + options: CreateDirectoryOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + trace!("remote fs create_directory"); + let client = self.client.get().await.map_err(map_remote_error)?; + let result = client + .fs_create_directory(FsCreateDirectoryParams { + path: path.clone(), + recursive: Some(options.recursive), + follow_symlinks: (!options.follow_symlinks).then_some(false), + sandbox: remote_sandbox_context(sandbox), + }) + .await; + self.metadata_requests.lock().await.clear(); + result.map_err(map_remote_error)?; + Ok(()) + } + + /// Shares identical unsandboxed metadata requests only while their RPC remains in flight. + async fn get_metadata( + &self, + path: &PathUri, + options: GetMetadataOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + if sandbox.is_some() || !options.follow_symlinks { + return self.get_metadata_uncached(path, options, sandbox).await; + } + + let request = { + let mut requests = self.metadata_requests.lock().await; + requests.retain(|_, in_flight| Arc::strong_count(in_flight) > 1); + Arc::clone(requests.entry(path.clone()).or_default()) + }; + let result = match request + .get_or_init(|| async { + self.get_metadata_uncached(path, options, /*sandbox*/ None) + .await + .map_err(Arc::new) + }) + .await + { + Ok(metadata) => Ok(metadata.clone()), + Err(error) => Err(io::Error::new(error.kind(), error.to_string())), + }; + + let mut requests = self.metadata_requests.lock().await; + if requests + .get(path) + .is_some_and(|in_flight| Arc::ptr_eq(in_flight, &request)) + { + requests.remove(path); + } + + result + } + + /// Sends a fresh metadata request without sharing another caller's sandbox or result. + async fn get_metadata_uncached( + &self, + path: &PathUri, + options: GetMetadataOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + trace!("remote fs get_metadata"); + let client = self.client.get().await.map_err(map_remote_error)?; + let response = client + .fs_get_metadata(FsGetMetadataParams { + path: path.clone(), + follow_symlinks: (!options.follow_symlinks).then_some(false), + sandbox: remote_sandbox_context(sandbox), + }) + .await + .map_err(map_remote_error)?; + Ok(FileMetadata { + is_directory: response.is_directory, + is_file: response.is_file, + is_symlink: response.is_symlink, + size: response.size, + created_at_ms: response.created_at_ms, + modified_at_ms: response.modified_at_ms, + }) + } + + async fn read_directory( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult> { + trace!("remote fs read_directory"); + let client = self.client.get().await.map_err(map_remote_error)?; + let response = client + .fs_read_directory(FsReadDirectoryParams { + path: path.clone(), + sandbox: remote_sandbox_context(sandbox), + }) + .await + .map_err(map_remote_error)?; + Ok(response + .entries + .into_iter() + .map(|entry| ReadDirectoryEntry { + file_name: entry.file_name, + is_directory: entry.is_directory, + is_file: entry.is_file, + }) + .collect()) + } + + async fn walk( + &self, + path: &PathUri, + options: WalkOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + trace!("remote fs walk"); + let client = self.client.get().await.map_err(map_remote_error)?; + client + .fs_walk(FsWalkParams { + path: path.clone(), + options, + sandbox: remote_sandbox_context(sandbox), + }) + .await + .map_err(map_remote_error) + } + + async fn remove( + &self, + path: &PathUri, + options: RemoveOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + trace!("remote fs remove"); + let client = self.client.get().await.map_err(map_remote_error)?; + let result = client + .fs_remove(FsRemoveParams { + path: path.clone(), + recursive: Some(options.recursive), + force: Some(options.force), + follow_symlinks: (!options.follow_symlinks).then_some(false), + sandbox: remote_sandbox_context(sandbox), + }) + .await; + self.metadata_requests.lock().await.clear(); + result.map_err(map_remote_error)?; + Ok(()) + } + + async fn copy( + &self, + source_path: &PathUri, + destination_path: &PathUri, + options: CopyOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + trace!("remote fs copy"); + let client = self.client.get().await.map_err(map_remote_error)?; + let result = client + .fs_copy(FsCopyParams { + source_path: source_path.clone(), + destination_path: destination_path.clone(), + recursive: options.recursive, + sandbox: remote_sandbox_context(sandbox), + }) + .await; + self.metadata_requests.lock().await.clear(); + result.map_err(map_remote_error)?; + Ok(()) + } +} + +impl ExecutorFileSystem for RemoteFileSystem { + fn canonicalize<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, PathUri> { + Box::pin(RemoteFileSystem::canonicalize(self, path, sandbox)) + } + + fn read_file<'a>( + &'a self, + path: &'a PathUri, + options: ReadFileOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, Vec> { + Box::pin(RemoteFileSystem::read_file(self, path, options, sandbox)) + } + + fn read_file_stream<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, FileSystemReadStream> { + Box::pin(RemoteFileSystem::read_file_stream(self, path, sandbox)) + } + + fn write_file<'a>( + &'a self, + path: &'a PathUri, + contents: Vec, + options: WriteFileOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(RemoteFileSystem::write_file( + self, path, contents, options, sandbox, + )) + } + + fn create_directory<'a>( + &'a self, + path: &'a PathUri, + options: CreateDirectoryOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(RemoteFileSystem::create_directory( + self, path, options, sandbox, + )) + } + + fn get_metadata<'a>( + &'a self, + path: &'a PathUri, + options: GetMetadataOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, FileMetadata> { + Box::pin(RemoteFileSystem::get_metadata(self, path, options, sandbox)) + } + + fn read_directory<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, Vec> { + Box::pin(RemoteFileSystem::read_directory(self, path, sandbox)) + } + + fn walk<'a>( + &'a self, + path: &'a PathUri, + options: WalkOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, WalkOutcome> { + Box::pin(RemoteFileSystem::walk(self, path, options, sandbox)) + } + + fn remove<'a>( + &'a self, + path: &'a PathUri, + options: RemoveOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(RemoteFileSystem::remove(self, path, options, sandbox)) + } + + fn copy<'a>( + &'a self, + source_path: &'a PathUri, + destination_path: &'a PathUri, + options: CopyOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(RemoteFileSystem::copy( + self, + source_path, + destination_path, + options, + sandbox, + )) + } +} + +fn remote_sandbox_context( + sandbox: Option<&FileSystemSandboxContext>, +) -> Option { + sandbox + .cloned() + .map(FileSystemSandboxContext::drop_cwd_if_unused) +} + +fn map_remote_error(error: ExecServerError) -> io::Error { + match error { + ExecServerError::Server { code, message } if code == NOT_FOUND_ERROR_CODE => { + io::Error::new(io::ErrorKind::NotFound, message) + } + ExecServerError::Server { code, message } if code == INVALID_REQUEST_ERROR_CODE => { + io::Error::new(io::ErrorKind::InvalidInput, message) + } + ExecServerError::Server { message, .. } => io::Error::other(message), + ExecServerError::Closed | ExecServerError::Disconnected(_) => { + io::Error::new(io::ErrorKind::BrokenPipe, "exec-server transport closed") + } + _ => io::Error::other(error.to_string()), + } +} + +#[cfg(all(test, any(unix, windows)))] +#[path = "remote_file_system_path_uri_tests.rs"] +mod path_uri_tests; + +#[cfg(test)] +mod tests { + use codex_protocol::models::PermissionProfile; + use codex_protocol::permissions::FileSystemAccessMode; + use codex_protocol::permissions::FileSystemPath; + use codex_protocol::permissions::FileSystemSandboxEntry; + use codex_protocol::permissions::FileSystemSandboxPolicy; + use codex_protocol::permissions::FileSystemSpecialPath; + use codex_protocol::permissions::NetworkSandboxPolicy; + use codex_utils_absolute_path::AbsolutePathBuf; + use codex_utils_path_uri::PathUri; + use pretty_assertions::assert_eq; + + use super::*; + + #[test] + fn remote_sandbox_context_drops_unused_cwd() { + let policy = FileSystemSandboxPolicy::restricted(vec![FileSystemSandboxEntry { + path: FileSystemPath::Path { + path: absolute_test_path("remote-root").into(), + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }]); + let permissions = + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted); + let sandbox_context = FileSystemSandboxContext::from_permission_profile_with_cwd( + permissions, + path_uri("host-checkout"), + ); + + let remote_context = + remote_sandbox_context(Some(&sandbox_context)).expect("remote sandbox context"); + + assert_eq!(remote_context.cwd, None); + } + + #[test] + fn remote_sandbox_context_preserves_required_cwd() { + let policy = FileSystemSandboxPolicy::restricted(vec![FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + access: FileSystemAccessMode::Write, + missing_path_behavior: None, + }]); + let permissions = + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted); + let cwd = path_uri("host-checkout"); + let sandbox_context = + FileSystemSandboxContext::from_permission_profile_with_cwd(permissions, cwd.clone()); + + let remote_context = + remote_sandbox_context(Some(&sandbox_context)).expect("remote sandbox context"); + + assert_eq!(remote_context.cwd, Some(cwd)); + } + + #[test] + fn transport_errors_map_to_broken_pipe() { + let errors = [ + ExecServerError::Closed, + ExecServerError::Disconnected("exec-server transport disconnected".to_string()), + ]; + + let mapped_errors = errors + .into_iter() + .map(|error| { + let error = map_remote_error(error); + (error.kind(), error.to_string()) + }) + .collect::>(); + + assert_eq!( + mapped_errors, + vec![ + ( + io::ErrorKind::BrokenPipe, + "exec-server transport closed".to_string() + ), + ( + io::ErrorKind::BrokenPipe, + "exec-server transport closed".to_string() + ), + ] + ); + } + + fn absolute_test_path(name: &str) -> AbsolutePathBuf { + let path = std::env::temp_dir().join(name); + AbsolutePathBuf::from_absolute_path(&path).expect("absolute path") + } + + fn path_uri(name: &str) -> PathUri { + PathUri::from_abs_path(&absolute_test_path(name)) + } +} diff --git a/codex-rs/exec-server/src/remote_file_system_path_uri_tests.rs b/codex-rs/exec-server/src/remote_file_system_path_uri_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4956f4bfdb7263790ad6360955df61c7ab166e44 --- /dev/null +++ b/codex-rs/exec-server/src/remote_file_system_path_uri_tests.rs @@ -0,0 +1,662 @@ +#![allow(clippy::expect_used)] + +use codex_exec_server_protocol::JSONRPCError; +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_protocol::models::PermissionProfile; +use codex_protocol::permissions::FileSystemAccessMode; +use codex_protocol::permissions::FileSystemPath; +use codex_protocol::permissions::FileSystemSandboxEntry; +use codex_protocol::permissions::FileSystemSandboxPolicy; +use codex_protocol::permissions::FileSystemSpecialPath; +use codex_protocol::permissions::NetworkSandboxPolicy; +use codex_utils_path_uri::PathUri; +use futures::SinkExt; +use futures::StreamExt; +use pretty_assertions::assert_eq; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::sync::oneshot; +use tokio::time::Duration; +use tokio::time::timeout; +use tokio_tungstenite::WebSocketStream; +use tokio_tungstenite::accept_async; +use tokio_tungstenite::tungstenite::Message; + +use super::*; +use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT; +use crate::client_api::ExecServerTransportParams; +use crate::protocol::FS_COPY_METHOD; +use crate::protocol::FS_CREATE_DIRECTORY_METHOD; +use crate::protocol::FS_GET_METADATA_METHOD; +use crate::protocol::FS_READ_FILE_METHOD; +use crate::protocol::FS_REMOVE_METHOD; +use crate::protocol::FS_WRITE_FILE_METHOD; +use crate::protocol::FsGetMetadataParams; +use crate::protocol::FsGetMetadataResponse; +use crate::protocol::FsReadFileParams; +use crate::protocol::FsReadFileResponse; +use crate::protocol::INITIALIZE_METHOD; +use crate::protocol::INITIALIZED_METHOD; +use crate::protocol::InitializeResponse; + +#[tokio::test] +async fn remote_file_system_sends_path_and_sandbox_cwd_uris_without_native_conversion() { + let (websocket_url, captured_params, server) = + record_read_file_params(/*expected_requests*/ 2).await; + let file_system = RemoteFileSystem::new(LazyRemoteExecServerClient::new( + ExecServerTransportParams::websocket_url( + websocket_url, + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + ), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )); + let paths = vec![ + PathUri::parse("file:///C:/Users/Alice/src/main.rs").expect("valid drive URI"), + PathUri::parse("file://server/share/src/main.rs").expect("valid UNC URI"), + ]; + let sandbox_cwd = non_native_cwd(); + let policy = FileSystemSandboxPolicy::restricted(vec![FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + access: FileSystemAccessMode::Write, + missing_path_behavior: None, + }]); + let sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + sandbox_cwd, + ); + + for path in &paths { + assert_eq!( + file_system + .read_file(path, Default::default(), Some(&sandbox)) + .await + .expect("remote read should succeed"), + Vec::::new() + ); + } + + let expected_params = paths + .into_iter() + .map(|path| FsReadFileParams { + path, + follow_symlinks: None, + sandbox: Some(sandbox.clone()), + }) + .collect::>(); + assert_eq!( + captured_params.await.expect("captured params"), + expected_params + ); + server.await.expect("recording server should succeed"); +} + +#[tokio::test] +async fn concurrent_remote_metadata_requests_share_only_in_flight_results() { + let (abandoned_request_tx, abandoned_request_rx) = oneshot::channel(); + let (websocket_url, captured_params, server) = record_metadata_params(vec![ + MetadataResponse::Immediate(Ok(metadata_response(/*size*/ 42))), + MetadataResponse::Abandoned(abandoned_request_tx), + MetadataResponse::Immediate(Ok(metadata_response(/*size*/ 43))), + ]) + .await; + let file_system = std::sync::Arc::new(RemoteFileSystem::new(LazyRemoteExecServerClient::new( + ExecServerTransportParams::websocket_url( + websocket_url, + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + ), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ))); + let path = PathUri::parse("file:///workspace/project/AGENTS.md").expect("valid path URI"); + + let (first, second) = tokio::join!( + file_system.get_metadata(&path, Default::default(), /*sandbox*/ None), + file_system.get_metadata(&path, Default::default(), /*sandbox*/ None), + ); + let expected = FileMetadata { + is_directory: false, + is_file: true, + is_symlink: false, + size: 42, + created_at_ms: 10, + modified_at_ms: 20, + }; + assert_eq!(first.expect("first metadata request"), expected); + assert_eq!(second.expect("second metadata request"), expected); + assert!(file_system.metadata_requests.lock().await.is_empty()); + + let initializing_file_system = std::sync::Arc::clone(&file_system); + let initializing_path = path.clone(); + let initializer = tokio::spawn(async move { + initializing_file_system + .get_metadata( + &initializing_path, + Default::default(), + /*sandbox*/ None, + ) + .await + }); + abandoned_request_rx + .await + .expect("server should receive the abandoned metadata request"); + + let follower = file_system.get_metadata(&path, Default::default(), /*sandbox*/ None); + tokio::pin!(follower); + assert!(futures::poll!(follower.as_mut()).is_pending()); + assert_eq!( + file_system + .metadata_requests + .lock() + .await + .get(&path) + .map(std::sync::Arc::strong_count), + Some(3) + ); + + initializer.abort(); + assert!( + initializer + .await + .expect_err("metadata initializer should be aborted") + .is_cancelled() + ); + + let mut refreshed = expected; + refreshed.size = 43; + assert_eq!( + follower + .await + .expect("waiting metadata request should retry the abandoned initializer"), + refreshed + ); + assert!(file_system.metadata_requests.lock().await.is_empty()); + assert_eq!( + captured_params.await.expect("captured metadata requests"), + vec![ + FsGetMetadataParams { + path: path.clone(), + follow_symlinks: None, + sandbox: None, + }; + 3 + ] + ); + server.await.expect("metadata server should succeed"); +} + +#[tokio::test] +async fn concurrent_remote_metadata_errors_are_shared_but_retried() { + let (websocket_url, captured_params, server) = record_metadata_params(vec![ + MetadataResponse::Immediate(Err(JSONRPCErrorError { + code: NOT_FOUND_ERROR_CODE, + data: None, + message: "metadata not found".to_string(), + })), + MetadataResponse::Immediate(Ok(metadata_response(/*size*/ 42))), + ]) + .await; + let file_system = RemoteFileSystem::new(LazyRemoteExecServerClient::new( + ExecServerTransportParams::websocket_url( + websocket_url, + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + ), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )); + let path = PathUri::parse("file:///workspace/project/AGENTS.md").expect("valid path URI"); + + let (first, second) = tokio::join!( + file_system.get_metadata(&path, Default::default(), /*sandbox*/ None), + file_system.get_metadata(&path, Default::default(), /*sandbox*/ None), + ); + assert_eq!( + first.expect_err("first metadata error").kind(), + std::io::ErrorKind::NotFound + ); + assert_eq!( + second.expect_err("shared metadata error").kind(), + std::io::ErrorKind::NotFound + ); + assert_eq!( + file_system + .get_metadata(&path, Default::default(), /*sandbox*/ None) + .await + .expect("failed metadata request should be retried") + .size, + 42 + ); + assert_eq!( + captured_params.await.expect("captured metadata requests"), + vec![ + FsGetMetadataParams { + path: path.clone(), + follow_symlinks: None, + sandbox: None, + }, + FsGetMetadataParams { + path, + follow_symlinks: None, + sandbox: None, + }, + ] + ); + server.await.expect("metadata server should succeed"); +} + +#[tokio::test] +async fn remote_metadata_starts_fresh_after_intervening_filesystem_mutation() { + let path = PathUri::parse("file:///workspace/project/AGENTS.md").expect("valid path URI"); + let source_path = + PathUri::parse("file:///workspace/project/source.md").expect("valid source path URI"); + + for (mutation, mutation_response) in [ + (MetadataMutation::Write, Ok(())), + (MetadataMutation::CreateDirectory, Ok(())), + (MetadataMutation::Remove, Ok(())), + (MetadataMutation::Copy, Ok(())), + ( + MetadataMutation::Copy, + Err(JSONRPCErrorError { + code: INVALID_REQUEST_ERROR_CODE, + data: None, + message: "mutation partially failed".to_string(), + }), + ), + ] { + let mutation_fails = mutation_response.is_err(); + let (websocket_url, captured_params, server) = record_metadata_params(vec![ + MetadataResponse::Immediate(Ok(metadata_response(/*size*/ 42))), + MetadataResponse::Mutation(mutation_response), + MetadataResponse::Immediate(Ok(metadata_response(/*size*/ 43))), + ]) + .await; + let file_system = RemoteFileSystem::new(LazyRemoteExecServerClient::new( + ExecServerTransportParams::websocket_url( + websocket_url, + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + ), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )); + file_system + .client + .get() + .await + .expect("remote filesystem client should connect"); + + let stale_request = + file_system.get_metadata(&path, Default::default(), /*sandbox*/ None); + tokio::pin!(stale_request); + assert!(futures::poll!(stale_request.as_mut()).is_pending()); + + let result = match mutation { + MetadataMutation::Write => { + file_system + .write_file( + &path, + b"updated".to_vec(), + Default::default(), + /*sandbox*/ None, + ) + .await + } + MetadataMutation::CreateDirectory => { + file_system + .create_directory( + &path, + CreateDirectoryOptions { + recursive: true, + follow_symlinks: true, + }, + /*sandbox*/ None, + ) + .await + } + MetadataMutation::Remove => { + file_system + .remove( + &path, + RemoveOptions { + recursive: true, + force: true, + follow_symlinks: true, + }, + /*sandbox*/ None, + ) + .await + } + MetadataMutation::Copy => { + file_system + .copy( + &source_path, + &path, + CopyOptions { recursive: true }, + /*sandbox*/ None, + ) + .await + } + }; + if mutation_fails { + assert_eq!( + result + .expect_err("remote filesystem mutation should fail") + .kind(), + std::io::ErrorKind::InvalidInput + ); + } else { + result.expect("remote filesystem mutation should succeed"); + } + + let (refreshed, stale) = tokio::join!( + file_system.get_metadata(&path, Default::default(), /*sandbox*/ None), + stale_request.as_mut(), + ); + assert_eq!(stale.expect("original metadata request").size, 42); + assert_eq!(refreshed.expect("post-mutation metadata request").size, 43); + assert_eq!( + captured_params.await.expect("captured metadata requests"), + vec![ + FsGetMetadataParams { + path: path.clone(), + follow_symlinks: None, + sandbox: None, + }; + 2 + ] + ); + server.await.expect("metadata server should succeed"); + } +} + +#[tokio::test] +async fn remote_metadata_requests_do_not_cross_path_or_sandbox_boundaries() { + let metadata = metadata_response(/*size*/ 42); + let (websocket_url, captured_params, server) = record_metadata_params(vec![ + MetadataResponse::Immediate(Ok(metadata.clone())), + MetadataResponse::Immediate(Ok(metadata.clone())), + MetadataResponse::Immediate(Ok(metadata.clone())), + MetadataResponse::Immediate(Ok(metadata)), + ]) + .await; + let file_system = RemoteFileSystem::new(LazyRemoteExecServerClient::new( + ExecServerTransportParams::websocket_url( + websocket_url, + DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT, + ), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )); + let first_path = PathUri::parse("file:///workspace/project/AGENTS.md").expect("valid path URI"); + let second_path = PathUri::parse("file:///workspace/project/SKILL.md").expect("valid path URI"); + let sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions( + &FileSystemSandboxPolicy::restricted(vec![FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + access: FileSystemAccessMode::Write, + missing_path_behavior: None, + }]), + NetworkSandboxPolicy::Restricted, + ), + non_native_cwd(), + ); + + let (first, second) = tokio::join!( + file_system.get_metadata(&first_path, Default::default(), /*sandbox*/ None), + file_system.get_metadata(&second_path, Default::default(), /*sandbox*/ None), + ); + first.expect("metadata for first path"); + second.expect("metadata for second path"); + + let (first, second) = tokio::join!( + file_system.get_metadata(&first_path, Default::default(), Some(&sandbox)), + file_system.get_metadata(&first_path, Default::default(), Some(&sandbox)), + ); + first.expect("first sandboxed metadata request"); + second.expect("second sandboxed metadata request"); + + let captured_params = captured_params.await.expect("captured metadata requests"); + assert_eq!( + captured_params + .iter() + .filter(|params| params.path == first_path && params.sandbox.is_none()) + .count(), + 1 + ); + assert_eq!( + captured_params + .iter() + .filter(|params| params.path == second_path && params.sandbox.is_none()) + .count(), + 1 + ); + assert_eq!( + captured_params + .iter() + .filter(|params| params.path == first_path && params.sandbox.as_ref() == Some(&sandbox)) + .count(), + 2 + ); + server.await.expect("metadata server should succeed"); +} + +async fn record_read_file_params( + expected_requests: usize, +) -> ( + String, + oneshot::Receiver>, + tokio::task::JoinHandle<()>, +) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!("ws://{}", listener.local_addr().expect("listener address")); + let (captured_params_tx, captured_params_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("listener should accept"); + let mut websocket = accept_async(stream) + .await + .expect("websocket handshake should succeed"); + complete_websocket_initialize(&mut websocket).await; + + let mut captured_params = Vec::with_capacity(expected_requests); + for _ in 0..expected_requests { + let request = match read_jsonrpc_websocket(&mut websocket).await { + JSONRPCMessage::Request(request) if request.method == FS_READ_FILE_METHOD => { + request + } + other => panic!("expected fs/readFile request, got {other:?}"), + }; + let params: FsReadFileParams = + serde_json::from_value(request.params.expect("fs/readFile params should exist")) + .expect("fs/readFile params should deserialize"); + captured_params.push(params); + write_jsonrpc_websocket( + &mut websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(FsReadFileResponse { + data_base64: String::new(), + }) + .expect("fs/readFile response should serialize"), + }), + ) + .await; + } + captured_params_tx + .send(captured_params) + .expect("captured params receiver should stay open"); + }); + + (websocket_url, captured_params_rx, server) +} + +enum MetadataResponse { + Immediate(Result), + Abandoned(oneshot::Sender<()>), + Mutation(Result<(), JSONRPCErrorError>), +} + +fn metadata_response(size: u64) -> FsGetMetadataResponse { + FsGetMetadataResponse { + is_directory: false, + is_file: true, + is_symlink: false, + size, + created_at_ms: 10, + modified_at_ms: 20, + } +} + +#[derive(Clone, Copy)] +enum MetadataMutation { + Write, + CreateDirectory, + Remove, + Copy, +} + +async fn record_metadata_params( + responses: Vec, +) -> ( + String, + oneshot::Receiver>, + tokio::task::JoinHandle<()>, +) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!("ws://{}", listener.local_addr().expect("listener address")); + let (captured_params_tx, captured_params_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("listener should accept"); + let mut websocket = accept_async(stream) + .await + .expect("websocket handshake should succeed"); + complete_websocket_initialize(&mut websocket).await; + + let mut captured_params = Vec::with_capacity(responses.len()); + for response in responses { + let request = match read_jsonrpc_websocket(&mut websocket).await { + JSONRPCMessage::Request(request) => request, + other => panic!("expected filesystem request, got {other:?}"), + }; + if matches!(response, MetadataResponse::Mutation(_)) { + assert!(matches!( + request.method.as_str(), + FS_WRITE_FILE_METHOD + | FS_CREATE_DIRECTORY_METHOD + | FS_REMOVE_METHOD + | FS_COPY_METHOD + )); + } else { + assert_eq!(request.method, FS_GET_METADATA_METHOD); + let params: FsGetMetadataParams = serde_json::from_value( + request.params.expect("fs/getMetadata params should exist"), + ) + .expect("fs/getMetadata params should deserialize"); + captured_params.push(params); + } + let response = match response { + MetadataResponse::Immediate(Ok(metadata)) => { + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(metadata) + .expect("fs/getMetadata response should serialize"), + }) + } + MetadataResponse::Immediate(Err(error)) + | MetadataResponse::Mutation(Err(error)) => JSONRPCMessage::Error(JSONRPCError { + error, + id: request.id, + }), + MetadataResponse::Mutation(Ok(())) => JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::json!({}), + }), + MetadataResponse::Abandoned(received_tx) => { + received_tx + .send(()) + .expect("abandoned request observer should stay open"); + continue; + } + }; + write_jsonrpc_websocket(&mut websocket, response).await; + } + captured_params_tx + .send(captured_params) + .expect("captured params receiver should stay open"); + }); + + (websocket_url, captured_params_rx, server) +} + +fn non_native_cwd() -> PathUri { + #[cfg(unix)] + let uri = "file://server/share/checkout"; + #[cfg(windows)] + let uri = "file:///usr/local/checkout"; + + PathUri::parse(uri).expect("non-native cwd URI") +} + +async fn complete_websocket_initialize(websocket: &mut WebSocketStream) { + let request = match read_jsonrpc_websocket(websocket).await { + JSONRPCMessage::Request(request) if request.method == INITIALIZE_METHOD => request, + other => panic!("expected initialize request, got {other:?}"), + }; + write_jsonrpc_websocket( + websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(InitializeResponse { + session_id: "session-1".to_string(), + environment_info: None, + }) + .expect("initialize response should serialize"), + }), + ) + .await; + + match read_jsonrpc_websocket(websocket).await { + JSONRPCMessage::Notification(notification) if notification.method == INITIALIZED_METHOD => { + } + other => panic!("expected initialized notification, got {other:?}"), + } +} + +async fn read_jsonrpc_websocket(websocket: &mut WebSocketStream) -> JSONRPCMessage { + loop { + match timeout(Duration::from_secs(1), websocket.next()) + .await + .expect("json-rpc 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("json-rpc text frame should parse"); + } + Message::Binary(bytes) => { + return serde_json::from_slice(bytes.as_ref()) + .expect("json-rpc binary frame should parse"); + } + Message::Ping(_) | Message::Pong(_) => {} + other => panic!("expected json-rpc websocket frame, got {other:?}"), + } + } +} + +async fn write_jsonrpc_websocket( + websocket: &mut WebSocketStream, + message: JSONRPCMessage, +) { + let encoded = serde_json::to_string(&message).expect("json-rpc should serialize"); + websocket + .send(Message::Text(encoded.into())) + .await + .expect("json-rpc websocket frame should write"); +} diff --git a/codex-rs/exec-server/src/remote_process.rs b/codex-rs/exec-server/src/remote_process.rs new file mode 100644 index 0000000000000000000000000000000000000000..4c2d4c8e63514a3f8157f1b0275662f5407fcd68 --- /dev/null +++ b/codex-rs/exec-server/src/remote_process.rs @@ -0,0 +1,137 @@ +use std::sync::Arc; + +use codex_network_proxy::NetworkPolicyDecider; +use tokio::sync::watch; +use tracing::trace; + +use crate::ExecBackend; +use crate::ExecBackendFuture; +use crate::ExecProcess; +use crate::ExecProcessEventReceiver; +use crate::ExecProcessFuture; +use crate::StartedExecProcess; +use crate::client::LazyRemoteExecServerClient; +use crate::client::Session; +use crate::process::sandbox_type_from_protocol; +use crate::protocol::ExecParams; +use crate::protocol::ProcessSignal; +use crate::protocol::ReadResponse; +use crate::protocol::WriteResponse; + +#[derive(Clone)] +pub(crate) struct RemoteProcess { + client: LazyRemoteExecServerClient, +} + +struct RemoteExecProcess { + session: Session, +} + +impl RemoteProcess { + pub(crate) fn new(client: LazyRemoteExecServerClient) -> Self { + trace!("remote process new"); + Self { client } + } + + async fn start( + &self, + params: ExecParams, + network_policy_decider: Option>, + ) -> Result { + let client = self.client.get().await?; + let session = client.start_process(params, network_policy_decider).await?; + let sandbox_type = sandbox_type_from_protocol(session.sandbox_type()); + + Ok(StartedExecProcess { + process: Arc::new(RemoteExecProcess { session }), + sandbox_type, + }) + } +} + +impl ExecBackend for RemoteProcess { + fn start(&self, params: ExecParams) -> ExecBackendFuture<'_> { + Box::pin(RemoteProcess::start( + self, params, /*network_policy_decider*/ None, + )) + } + + fn start_with_network_policy_decider( + &self, + params: ExecParams, + decider: Arc, + ) -> ExecBackendFuture<'_> { + Box::pin(RemoteProcess::start(self, params, Some(decider))) + } +} + +impl RemoteExecProcess { + async fn read( + &self, + after_seq: Option, + max_bytes: Option, + wait_ms: Option, + ) -> Result { + self.session.read(after_seq, max_bytes, wait_ms).await + } + + async fn write(&self, chunk: Vec) -> Result { + trace!("exec process write"); + self.session.write(chunk).await + } + + async fn signal(&self, signal: ProcessSignal) -> Result<(), crate::ExecServerError> { + trace!("exec process signal"); + self.session.signal(signal).await + } + + async fn terminate(&self) -> Result<(), crate::ExecServerError> { + trace!("exec process terminate"); + self.session.terminate().await + } +} + +impl ExecProcess for RemoteExecProcess { + fn process_id(&self) -> &crate::ProcessId { + self.session.process_id() + } + + fn subscribe_wake(&self) -> watch::Receiver { + self.session.subscribe_wake() + } + + fn subscribe_events(&self) -> ExecProcessEventReceiver { + self.session.subscribe_events() + } + + fn read( + &self, + after_seq: Option, + max_bytes: Option, + wait_ms: Option, + ) -> ExecProcessFuture<'_, ReadResponse> { + Box::pin(RemoteExecProcess::read(self, after_seq, max_bytes, wait_ms)) + } + + fn write(&self, chunk: Vec) -> ExecProcessFuture<'_, WriteResponse> { + Box::pin(RemoteExecProcess::write(self, chunk)) + } + + fn signal(&self, signal: ProcessSignal) -> ExecProcessFuture<'_, ()> { + Box::pin(RemoteExecProcess::signal(self, signal)) + } + + fn terminate(&self) -> ExecProcessFuture<'_, ()> { + Box::pin(RemoteExecProcess::terminate(self)) + } +} + +impl Drop for RemoteExecProcess { + fn drop(&mut self) { + self.session.cancel_network_policy_decisions(); + let session = self.session.clone(); + tokio::spawn(async move { + session.unregister().await; + }); + } +} diff --git a/codex-rs/exec-server/src/resolved_capability.rs b/codex-rs/exec-server/src/resolved_capability.rs new file mode 100644 index 0000000000000000000000000000000000000000..078b2b982b215331e2f1c781bd87ffb4f28f91ee --- /dev/null +++ b/codex-rs/exec-server/src/resolved_capability.rs @@ -0,0 +1,187 @@ +use std::collections::HashMap; +use std::fmt; +use std::sync::Arc; + +use codex_protocol::capabilities::CapabilityRootLocation; +use codex_protocol::capabilities::SelectedCapabilityRoot; + +use crate::Environment; +use crate::EnvironmentManager; + +/// A selected capability root paired with its currently ready environment handle. +/// +/// Environment IDs have stable identity and contents. This process-local value must not be +/// persisted: it only keeps the current connection handle alive while one model step uses the +/// stable environment. +#[derive(Clone)] +pub struct ResolvedSelectedCapabilityRoot { + selected_root: SelectedCapabilityRoot, + environment: Arc, +} + +/// A passive view of selected capability roots and unavailable environments. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct SelectedCapabilityRootsStatus { + /// Selected roots whose environments are ready. + pub ready_roots: Vec, + /// Missing environments and terminal connection failures. + pub warnings: Vec, +} + +impl ResolvedSelectedCapabilityRoot { + pub fn selected_root(&self) -> &SelectedCapabilityRoot { + &self.selected_root + } + + pub fn environment(&self) -> &Arc { + &self.environment + } +} + +impl EnvironmentManager { + /// Inspects selected roots without starting or waiting for an environment. + /// + /// Starting or recovering environments are omitted. Missing environments and terminal + /// connection failures are returned as warnings so read-only catalog clients can distinguish + /// them from an empty catalog. + /// + /// Environment IDs are stable identities, so callers can safely resolve a returned root by + /// ID when they read it. + pub fn inspect_selected_capability_roots( + &self, + selected_roots: &[SelectedCapabilityRoot], + ) -> SelectedCapabilityRootsStatus { + let (candidates, mut warnings) = { + let environments = self + .environments + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let mut candidates = Vec::with_capacity(selected_roots.len()); + let mut warnings = Vec::new(); + for selected_root in selected_roots { + let CapabilityRootLocation::Environment { environment_id, .. } = + &selected_root.location; + let Some(environment) = environments.get(environment_id) else { + warnings.push(format!( + "selected capability root `{}` references unavailable environment `{environment_id}`", + selected_root.id + )); + continue; + }; + candidates.push((selected_root.clone(), Arc::clone(environment))); + } + (candidates, warnings) + }; + let mut readiness = HashMap::new(); + for (selected_root, environment) in &candidates { + let CapabilityRootLocation::Environment { environment_id, .. } = + &selected_root.location; + if readiness.contains_key(environment_id) { + continue; + } + let ready = match environment.readiness_result() { + Some(Ok(())) => true, + Some(Err(error)) => { + warnings.push(format!( + "selected capability environment `{environment_id}` is unavailable: {error}" + )); + false + } + None => false, + }; + readiness.insert(environment_id.clone(), ready); + } + + let ready_roots = candidates + .into_iter() + .filter(|(selected_root, _)| { + let CapabilityRootLocation::Environment { environment_id, .. } = + &selected_root.location; + readiness.get(environment_id).copied().unwrap_or(false) + }) + .map(|(selected_root, _)| selected_root) + .collect(); + SelectedCapabilityRootsStatus { + ready_roots, + warnings, + } + } + + /// Resolves selected roots whose stable environments are ready for the current model step. + /// + /// Environment identity comes from the selected root's stable environment ID. A ready + /// environment captured for the step carries its exact process-local handle so readiness and + /// execution cannot come from different registry snapshots. Missing, starting, or failed + /// environments are omitted. A lazy environment is started for a later step. + #[tracing::instrument(name = "capability_roots.resolve", skip_all)] + pub async fn resolve_selected_capability_roots( + &self, + selected_roots: &[SelectedCapabilityRoot], + captured_environments: &HashMap>>, + ) -> Vec { + let candidates = { + let environments = self + .environments + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner); + selected_roots + .iter() + .filter_map(|selected_root| { + let CapabilityRootLocation::Environment { environment_id, .. } = + &selected_root.location; + let (environment, already_ready) = + match captured_environments.get(environment_id) { + Some(Some(environment)) => (Arc::clone(environment), true), + Some(None) => return None, + None => (Arc::clone(environments.get(environment_id)?), false), + }; + Some(( + ResolvedSelectedCapabilityRoot { + selected_root: selected_root.clone(), + environment, + }, + already_ready, + )) + }) + .collect::>() + }; + + let mut readiness = HashMap::new(); + for (candidate, already_ready) in &candidates { + let CapabilityRootLocation::Environment { environment_id, .. } = + &candidate.selected_root().location; + if readiness.contains_key(environment_id) { + continue; + } + let environment = candidate.environment(); + let ready = if *already_ready { + true + } else if environment.startup_finished() { + environment.wait_until_ready().await.is_ok() + } else { + Environment::start_connecting_for_use(environment); + false + }; + readiness.insert(environment_id.clone(), ready); + } + + candidates + .into_iter() + .map(|(candidate, _)| candidate) + .filter(|candidate| { + let CapabilityRootLocation::Environment { environment_id, .. } = + &candidate.selected_root().location; + readiness.get(environment_id).copied().unwrap_or(false) + }) + .collect() + } +} + +impl fmt::Debug for ResolvedSelectedCapabilityRoot { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ResolvedSelectedCapabilityRoot") + .field("selected_root", &self.selected_root) + .finish_non_exhaustive() + } +} diff --git a/codex-rs/exec-server/src/rpc.rs b/codex-rs/exec-server/src/rpc.rs new file mode 100644 index 0000000000000000000000000000000000000000..5047b065620f8d072301f94a6cecbcc829b8faed --- /dev/null +++ b/codex-rs/exec-server/src/rpc.rs @@ -0,0 +1,1317 @@ +use std::collections::HashMap; +use std::collections::HashSet; +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::AtomicI64; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCError; +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCNotification; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::RequestId; +use codex_otel::MetricsClient; +use codex_protocol::protocol::W3cTraceContext; +use serde::Serialize; +use serde::de::DeserializeOwned; +use serde_json::Value; +use tokio::sync::Mutex; +use tokio::sync::OwnedSemaphorePermit; +use tokio::sync::Semaphore; +use tokio::sync::SemaphorePermit; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::sync::watch; +use tokio::task::JoinHandle; +use tokio::time::timeout; + +use crate::client_telemetry::record_client_request; +use crate::connection::JsonRpcConnection; +use crate::connection::JsonRpcConnectionEvent; +use crate::connection::JsonRpcTransport; +use crate::rpc_server_requests::RpcServerRequestSender; + +#[cfg(test)] +#[path = "rpc_client_metrics_tests.rs"] +mod client_metrics_tests; + +pub(crate) const SESSION_ALREADY_ATTACHED_ERROR_CODE: i64 = -32010; +const MAX_IN_FLIGHT_REGULAR_CALLS: usize = 1024; +const RESERVED_CLEANUP_CALLS: usize = 1; +const RESERVED_OUTBOUND_CONTROL_MESSAGES: usize = 16; + +#[derive(Debug)] +pub(crate) enum RpcCallError { + /// The underlying JSON-RPC transport closed before this call completed. + Closed, + /// The response bytes were valid JSON-RPC but not the expected result type. + Json(serde_json::Error), + /// The executor returned a JSON-RPC error response for this call. + Server(JSONRPCErrorError), + /// The executor did not return a response before the caller's deadline. + TimedOut { method: String, timeout: Duration }, + /// The client already has the maximum number of regular RPC calls in flight. + PendingRequestLimitExceeded { limit: usize }, +} + +type PendingRequest = oneshot::Sender>; +type BoxFuture = Pin + Send + 'static>>; +type RequestRoute = Box< + dyn Fn(Arc, JSONRPCRequest) -> BoxFuture> + Send + Sync, +>; +type NotificationRoute = + Box, JSONRPCNotification) -> BoxFuture> + Send + Sync>; + +enum RpcCallTimeout { + None, + After(Duration), +} + +#[derive(Debug)] +pub(crate) enum RpcClientEvent { + Request { + request: JSONRPCRequest, + request_span: tracing::Span, + }, + Notification(JSONRPCNotification), + Disconnected { + reason: Option, + }, +} + +pub(crate) enum RpcInboundRequestAdmissionError { + InvalidRequestId, + DuplicateRequestId, + AtCapacity, +} + +pub(crate) struct RpcInboundRequestGuard { + request_id: RequestId, + request_ids: Arc>>, + _call_slot: OwnedSemaphorePermit, +} + +impl Drop for RpcInboundRequestGuard { + fn drop(&mut self) { + self.request_ids + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .remove(&self.request_id); + } +} + +#[derive(Debug, Clone, PartialEq)] +pub(crate) enum RpcServerOutboundMessage { + Request(JSONRPCRequest), + Response { + request_id: RequestId, + result: Value, + }, + Error { + request_id: RequestId, + error: JSONRPCErrorError, + }, + Notification(JSONRPCNotification), +} + +#[derive(Clone)] +pub(crate) struct RpcNotificationSender { + outgoing_tx: mpsc::Sender, + requests: RpcServerRequestSender, +} + +impl RpcNotificationSender { + pub(crate) fn new(outgoing_tx: mpsc::Sender) -> Self { + let requests = RpcServerRequestSender::new(outgoing_tx.clone()); + Self { + outgoing_tx, + requests, + } + } + + pub(crate) fn request_sender(&self) -> RpcServerRequestSender { + self.requests.clone() + } + + pub(crate) async fn response( + &self, + request_id: RequestId, + result: Value, + ) -> Result<(), JSONRPCErrorError> { + self.outgoing_tx + .send(RpcServerOutboundMessage::Response { request_id, result }) + .await + .map_err(|_| internal_error("RPC connection closed while sending response".into())) + } + + pub(crate) async fn notify( + &self, + method: &str, + params: &P, + ) -> Result<(), JSONRPCErrorError> { + let params = serde_json::to_value(params).map_err(|err| internal_error(err.to_string()))?; + self.outgoing_tx + .send(RpcServerOutboundMessage::Notification( + JSONRPCNotification { + method: method.to_string(), + params: Some(params), + }, + )) + .await + .map_err(|_| internal_error("RPC connection closed while sending notification".into())) + } + + pub(crate) fn try_notify(&self, method: &str, params: &P) -> bool { + let Ok(permit) = self.outgoing_tx.try_reserve() else { + return false; + }; + if self.outgoing_tx.capacity() < RESERVED_OUTBOUND_CONTROL_MESSAGES { + return false; + } + let Ok(params) = serde_json::to_value(params) else { + return false; + }; + permit.send(RpcServerOutboundMessage::Notification( + JSONRPCNotification { + method: method.to_string(), + params: Some(params), + }, + )); + true + } +} + +pub(crate) struct RpcRouter { + request_routes: HashMap<&'static str, RequestRoute>, + notification_routes: HashMap<&'static str, NotificationRoute>, +} + +impl Default for RpcRouter { + fn default() -> Self { + Self { + request_routes: HashMap::new(), + notification_routes: HashMap::new(), + } + } +} + +impl RpcRouter +where + S: Send + Sync + 'static, +{ + pub(crate) fn new() -> Self { + Self::default() + } + + pub(crate) fn request(&mut self, method: &'static str, handler: F) + where + P: DeserializeOwned + Send + 'static, + R: Serialize + Send + 'static, + F: Fn(Arc, P) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + self.request_with_trace(method, move |state, params, _trace| handler(state, params)); + } + + /// Supplies the incoming W3C carrier to handlers that need it without requiring a trace exporter. + pub(crate) fn request_with_trace(&mut self, method: &'static str, handler: F) + where + P: DeserializeOwned + Send + 'static, + R: Serialize + Send + 'static, + F: Fn(Arc, P, Option) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + self.request_routes.insert( + method, + Box::new(move |state, request| { + let trace = request.trace; + let request_id = request.id; + let params = request.params; + let response = + decode_request_params::

(params).map(|params| handler(state, params, trace)); + Box::pin(async move { + let response = match response { + Ok(response) => response.await, + Err(error) => { + return Some(RpcServerOutboundMessage::Error { request_id, error }); + } + }; + Some(match response { + Ok(result) => match serde_json::to_value(result) { + Ok(result) => RpcServerOutboundMessage::Response { request_id, result }, + Err(err) => RpcServerOutboundMessage::Error { + request_id, + error: internal_error(err.to_string()), + }, + }, + Err(error) => RpcServerOutboundMessage::Error { request_id, error }, + }) + }) + }), + ); + } + + pub(crate) fn request_with_id(&mut self, method: &'static str, handler: F) + where + P: DeserializeOwned + Send + 'static, + F: Fn(Arc, RequestId, P) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + self.request_routes.insert( + method, + Box::new(move |state, request| { + let request_id = request.id; + let params = decode_request_params::

(request.params) + .map(|params| handler(state, request_id.clone(), params)); + Box::pin(async move { + let response = match params { + Ok(response) => response.await, + Err(error) => { + return Some(RpcServerOutboundMessage::Error { request_id, error }); + } + }; + match response { + Ok(()) => None, + Err(error) => Some(RpcServerOutboundMessage::Error { request_id, error }), + } + }) + }), + ); + } + + pub(crate) fn notification(&mut self, method: &'static str, handler: F) + where + P: DeserializeOwned + Send + 'static, + F: Fn(Arc, P) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + self.notification_routes.insert( + method, + Box::new(move |state, notification| { + let params = decode_notification_params::

(notification.params) + .map(|params| handler(state, params)); + Box::pin(async move { + let handler = match params { + Ok(handler) => handler, + Err(err) => return Err(err), + }; + handler.await + }) + }), + ); + } + + pub(crate) fn request_route(&self, method: &str) -> Option<(&'static str, &RequestRoute)> { + self.request_routes + .get_key_value(method) + .map(|(&method, route)| (method, route)) + } + + pub(crate) fn notification_route(&self, method: &str) -> Option<&NotificationRoute> { + self.notification_routes.get(method) + } +} + +pub(crate) struct RpcClient { + metrics: Option, + write_tx: mpsc::Sender, + pending: Arc>>, + inbound_request_ids: Arc>>, + // Shared transport state from `JsonRpcConnection`. Calls use this to fail + // immediately when the socket closes, even if no JSON-RPC error response + // can be delivered for their request id. + disconnected_rx: watch::Receiver, + closed: Arc, + shared_call_slots: Semaphore, + cleanup_call_slots: Semaphore, + next_request_id: AtomicI64, + transport_tasks: Vec>, + transport: JsonRpcTransport, + reader_task: JoinHandle<()>, +} + +impl RpcClient { + pub(crate) fn new(connection: JsonRpcConnection) -> (Self, mpsc::Receiver) { + let JsonRpcConnection { + outgoing_tx: write_tx, + mut incoming_rx, + disconnected_rx, + task_handles: transport_tasks, + transport, + } = connection; + let pending = Arc::new(Mutex::new(HashMap::::new())); + let closed = Arc::new(AtomicBool::new(false)); + let (event_tx, event_rx) = mpsc::channel(128); + + let pending_for_reader = Arc::clone(&pending); + let closed_for_reader = Arc::clone(&closed); + let transport_for_reader = transport.clone(); + let reader_task = tokio::spawn(async move { + let disconnect_reason = loop { + let Some(event) = incoming_rx.recv().await else { + break None; + }; + match event { + JsonRpcConnectionEvent::Message(message) => { + if let Err(err) = + handle_server_message(&pending_for_reader, &event_tx, message).await + { + let _ = err; + break None; + } + } + JsonRpcConnectionEvent::QueuedRequest { + request, + request_span, + .. + } => { + if event_tx + .send(RpcClientEvent::Request { + request, + request_span, + }) + .await + .is_err() + { + break None; + } + } + JsonRpcConnectionEvent::MalformedMessage { reason } => { + let _ = reason; + break None; + } + JsonRpcConnectionEvent::Disconnected { reason } => { + break reason; + } + } + }; + + closed_for_reader.store(true, Ordering::Release); + drain_pending(&pending_for_reader).await; + let _ = event_tx + .send(RpcClientEvent::Disconnected { + reason: disconnect_reason, + }) + .await; + transport_for_reader.terminate(); + }); + + ( + Self { + metrics: codex_otel::global(), + write_tx, + pending, + inbound_request_ids: Arc::new(StdMutex::new(HashSet::new())), + disconnected_rx, + closed, + shared_call_slots: Semaphore::new(MAX_IN_FLIGHT_REGULAR_CALLS), + cleanup_call_slots: Semaphore::new(RESERVED_CLEANUP_CALLS), + next_request_id: AtomicI64::new(1), + transport_tasks, + transport, + reader_task, + }, + event_rx, + ) + } + + pub(crate) fn admit_inbound_request( + &self, + request_id: &RequestId, + call_slots: &Arc, + ) -> Result { + let request_id = match request_id { + RequestId::Integer(request_id) if *request_id >= 0 => request_id, + RequestId::Integer(_) | RequestId::String(_) => { + return Err(RpcInboundRequestAdmissionError::InvalidRequestId); + } + }; + let request_id = RequestId::Integer(*request_id); + { + let mut request_ids = self + .inbound_request_ids + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if !request_ids.insert(request_id.clone()) { + return Err(RpcInboundRequestAdmissionError::DuplicateRequestId); + } + } + let call_slot = match Arc::clone(call_slots).try_acquire_owned() { + Ok(call_slot) => call_slot, + Err(_) => { + self.inbound_request_ids + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .remove(&request_id); + return Err(RpcInboundRequestAdmissionError::AtCapacity); + } + }; + Ok(RpcInboundRequestGuard { + request_id, + request_ids: Arc::clone(&self.inbound_request_ids), + _call_slot: call_slot, + }) + } + + pub(crate) async fn notify( + &self, + method: &str, + params: &P, + ) -> Result<(), RpcCallError> { + let params = serde_json::to_value(params).map_err(RpcCallError::Json)?; + if self.closed.load(Ordering::Acquire) || *self.disconnected_rx.borrow() { + return Err(RpcCallError::Closed); + } + self.write_tx + .send(JSONRPCMessage::Notification(JSONRPCNotification { + method: method.to_string(), + params: Some(params), + })) + .await + .map_err(|_| RpcCallError::Closed) + } + + pub(crate) async fn respond( + &self, + request_id: RequestId, + result: &T, + ) -> Result<(), RpcCallError> { + let result = serde_json::to_value(result).map_err(RpcCallError::Json)?; + if self.closed.load(Ordering::Acquire) || *self.disconnected_rx.borrow() { + return Err(RpcCallError::Closed); + } + self.write_tx + .send(JSONRPCMessage::Response(JSONRPCResponse { + id: request_id, + result, + })) + .await + .map_err(|_| RpcCallError::Closed) + } + + pub(crate) async fn respond_error( + &self, + request_id: RequestId, + error: JSONRPCErrorError, + ) -> Result<(), RpcCallError> { + if self.closed.load(Ordering::Acquire) || *self.disconnected_rx.borrow() { + return Err(RpcCallError::Closed); + } + self.write_tx + .send(JSONRPCMessage::Error(JSONRPCError { + id: request_id, + error, + })) + .await + .map_err(|_| RpcCallError::Closed) + } + + pub(crate) fn is_disconnected(&self) -> bool { + self.closed.load(Ordering::Acquire) || *self.disconnected_rx.borrow() + } + + pub(crate) async fn close_transport(&self) { + self.closed.store(true, Ordering::Release); + self.transport.terminate(); + for task in &self.transport_tasks { + task.abort(); + } + drain_pending(&self.pending).await; + } + + // Callers keep this permit until `call_inner` returns, so an executor + // cannot free admission early by guessing a request id and replying before + // the request leaves the outbound queue. + fn acquire_regular_call_slot(&self) -> Result, RpcCallError> { + self.shared_call_slots.try_acquire().map_err(|_| { + RpcCallError::PendingRequestLimitExceeded { + limit: MAX_IN_FLIGHT_REGULAR_CALLS, + } + }) + } + + #[tracing::instrument( + name = "codex.exec_server.request", + level = "info", + skip_all, + fields( + otel.kind = "client", + otel.name = method, + method, + ) + )] + pub(crate) async fn call(&self, method: &str, params: &P) -> Result + where + P: Serialize, + T: DeserializeOwned, + { + self.call_untraced(method, params).await + } + + /// Send one request without creating the standard request span. + /// + /// Callers use this only when they install a more precise request span + /// around the same wire operation. + pub(crate) async fn call_untraced( + &self, + method: &str, + params: &P, + ) -> Result + where + P: Serialize, + T: DeserializeOwned, + { + record_client_request(self.metrics.as_ref(), method); + let _call_slot = self.acquire_regular_call_slot()?; + self.call_inner(method, params, RpcCallTimeout::None).await + } + + pub(crate) async fn call_with_timeout( + &self, + method: &str, + params: &P, + call_timeout: Duration, + ) -> Result + where + P: Serialize, + T: DeserializeOwned, + { + record_client_request(self.metrics.as_ref(), method); + let _call_slot = self.acquire_regular_call_slot()?; + self.call_inner(method, params, RpcCallTimeout::After(call_timeout)) + .await + } + + #[tracing::instrument( + name = "codex.exec_server.request", + level = "info", + skip_all, + fields( + otel.kind = "client", + otel.name = method, + method, + ) + )] + pub(crate) async fn call_for_cleanup( + &self, + method: &str, + params: &P, + ) -> Result + where + P: Serialize, + T: DeserializeOwned, + { + record_client_request(self.metrics.as_ref(), method); + let _call_slot = match self.shared_call_slots.try_acquire() { + Ok(call_slot) => call_slot, + Err(_) => match self.cleanup_call_slots.try_acquire() { + Ok(call_slot) => call_slot, + Err(_) => { + self.close_transport().await; + return Err(RpcCallError::Closed); + } + }, + }; + self.call_inner(method, params, RpcCallTimeout::None).await + } + + async fn call_inner( + &self, + method: &str, + params: &P, + call_timeout: RpcCallTimeout, + ) -> Result + where + P: Serialize, + T: DeserializeOwned, + { + let request_id = RequestId::Integer(self.next_request_id.fetch_add(1, Ordering::SeqCst)); + let (response_tx, response_rx) = oneshot::channel(); + { + let mut pending = self.pending.lock().await; + // Registering the pending request and checking disconnect must be + // atomic with the reader's drain_pending path. Otherwise a call + // can sneak in after the drain and wait forever. + if self.closed.load(Ordering::Acquire) || *self.disconnected_rx.borrow() { + return Err(RpcCallError::Closed); + } + pending.retain(|_, response_tx| !response_tx.is_closed()); + pending.insert(request_id.clone(), response_tx); + } + + let params = match serde_json::to_value(params) { + Ok(params) => params, + Err(err) => { + self.pending.lock().await.remove(&request_id); + return Err(RpcCallError::Json(err)); + } + }; + if self + .write_tx + .send(JSONRPCMessage::Request(JSONRPCRequest { + id: request_id.clone(), + method: method.to_string(), + params: Some(params), + trace: codex_otel::current_span_w3c_trace_context(), + })) + .await + .is_err() + { + self.pending.lock().await.remove(&request_id); + return Err(RpcCallError::Closed); + } + + // Do not race in-flight requests directly against the transport-close + // watch value. The connection reader receives JSON-RPC messages and + // the terminal disconnect event on one ordered queue, then drains any + // still-pending requests. Awaiting this receiver preserves that order: + // responses already read before EOF still win, and truly pending calls + // are failed once the reader observes the disconnect. + let response = match call_timeout { + RpcCallTimeout::None => response_rx.await, + RpcCallTimeout::After(call_timeout) => match timeout(call_timeout, response_rx).await { + Ok(response) => response, + Err(_) => { + self.pending.lock().await.remove(&request_id); + return Err(RpcCallError::TimedOut { + method: method.to_string(), + timeout: call_timeout, + }); + } + }, + }; + let result: Result = response.map_err(|_| RpcCallError::Closed)?; + let response = match result { + Ok(response) => response, + Err(error) => return Err(error), + }; + serde_json::from_value(response).map_err(RpcCallError::Json) + } + + #[cfg(test)] + pub(crate) async fn pending_request_count(&self) -> usize { + self.pending.lock().await.len() + } +} + +impl Drop for RpcClient { + fn drop(&mut self) { + self.transport.terminate(); + for task in &self.transport_tasks { + task.abort(); + } + self.reader_task.abort(); + } +} + +pub(crate) fn encode_server_message( + message: RpcServerOutboundMessage, +) -> Result { + match message { + RpcServerOutboundMessage::Request(request) => Ok(JSONRPCMessage::Request(request)), + RpcServerOutboundMessage::Response { request_id, result } => { + Ok(JSONRPCMessage::Response(JSONRPCResponse { + id: request_id, + result, + })) + } + RpcServerOutboundMessage::Error { request_id, error } => { + Ok(JSONRPCMessage::Error(JSONRPCError { + id: request_id, + error, + })) + } + RpcServerOutboundMessage::Notification(notification) => { + Ok(JSONRPCMessage::Notification(notification)) + } + } +} + +pub(crate) fn invalid_request(message: String) -> JSONRPCErrorError { + JSONRPCErrorError { + code: -32600, + data: None, + message, + } +} + +pub(crate) fn session_already_attached(message: String) -> JSONRPCErrorError { + JSONRPCErrorError { + code: SESSION_ALREADY_ATTACHED_ERROR_CODE, + data: None, + message, + } +} + +pub(crate) fn method_not_found(message: String) -> JSONRPCErrorError { + JSONRPCErrorError { + code: -32601, + data: None, + message, + } +} + +pub(crate) fn invalid_params(message: String) -> JSONRPCErrorError { + JSONRPCErrorError { + code: -32602, + data: None, + message, + } +} + +pub(crate) fn not_found(message: String) -> JSONRPCErrorError { + JSONRPCErrorError { + code: -32004, + data: None, + message, + } +} + +pub(crate) fn internal_error(message: String) -> JSONRPCErrorError { + JSONRPCErrorError { + code: -32603, + data: None, + message, + } +} + +fn decode_request_params

(params: Option) -> Result +where + P: DeserializeOwned, +{ + decode_params(params).map_err(|err| invalid_params(err.to_string())) +} + +fn decode_notification_params

(params: Option) -> Result +where + P: DeserializeOwned, +{ + decode_params(params).map_err(|err| err.to_string()) +} + +fn decode_params

(params: Option) -> Result +where + P: DeserializeOwned, +{ + let params = params.unwrap_or(Value::Null); + let retry_as_null = matches!(¶ms, Value::Object(map) if map.is_empty()); + match serde_json::from_value(params) { + Ok(params) => Ok(params), + Err(err) => { + if retry_as_null { + serde_json::from_value(Value::Null).map_err(|_| err) + } else { + Err(err) + } + } + } +} + +async fn handle_server_message( + pending: &Mutex>, + event_tx: &mpsc::Sender, + message: JSONRPCMessage, +) -> Result<(), String> { + match message { + JSONRPCMessage::Response(JSONRPCResponse { id, result }) => { + if let Some(pending) = pending.lock().await.remove(&id) { + let _ = pending.send(Ok(result)); + } + } + JSONRPCMessage::Error(JSONRPCError { id, error }) => { + if let Some(pending) = pending.lock().await.remove(&id) { + let _ = pending.send(Err(RpcCallError::Server(error))); + } + } + JSONRPCMessage::Notification(notification) => { + let _ = event_tx + .send(RpcClientEvent::Notification(notification)) + .await; + } + JSONRPCMessage::Request(request) => { + event_tx + .send(RpcClientEvent::Request { + request, + request_span: tracing::Span::none(), + }) + .await + .map_err(|_| "RPC client event receiver closed".to_string())?; + } + } + + Ok(()) +} + +async fn drain_pending(pending: &Mutex>) { + let pending = { + let mut pending = pending.lock().await; + pending + .drain() + .map(|(_, pending)| pending) + .collect::>() + }; + for pending in pending { + let _ = pending.send(Err(RpcCallError::Closed)); + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + use std::time::Duration; + + use codex_exec_server_protocol::JSONRPCMessage; + use codex_exec_server_protocol::JSONRPCNotification; + use codex_exec_server_protocol::JSONRPCRequest; + use codex_exec_server_protocol::JSONRPCResponse; + use codex_exec_server_protocol::RequestId; + use opentelemetry::trace::TracerProvider as _; + use opentelemetry_sdk::trace::InMemorySpanExporter; + use opentelemetry_sdk::trace::SdkTracerProvider; + use pretty_assertions::assert_eq; + use tokio::io::AsyncBufReadExt; + use tokio::io::AsyncWriteExt; + use tokio::io::BufReader; + use tokio::sync::mpsc; + use tokio::task::JoinSet; + use tokio::time::timeout; + use tracing::Instrument; + use tracing_subscriber::filter::filter_fn; + use tracing_subscriber::prelude::*; + + use super::MAX_IN_FLIGHT_REGULAR_CALLS; + use super::RESERVED_OUTBOUND_CONTROL_MESSAGES; + use super::RpcCallError; + use super::RpcClient; + use super::RpcClientEvent; + use super::RpcNotificationSender; + use crate::connection::JsonRpcConnection; + use crate::connection::JsonRpcConnectionEvent; + use crate::connection::JsonRpcTransport; + + #[tokio::test] + async fn best_effort_notifications_preserve_outbound_control_capacity() { + let (outgoing_tx, _outgoing_rx) = mpsc::channel(RESERVED_OUTBOUND_CONTROL_MESSAGES + 2); + let notifications = RpcNotificationSender::new(outgoing_tx); + + assert!(notifications.try_notify("network/policyDecision", &serde_json::json!({"n": 1}))); + assert!(notifications.try_notify("network/policyDecision", &serde_json::json!({"n": 2}))); + assert!(!notifications.try_notify("network/policyDecision", &serde_json::json!({"n": 3}))); + notifications + .response(RequestId::Integer(7), serde_json::json!({"ok": true})) + .await + .expect("reserved capacity must remain available for controller responses"); + } + + async fn read_jsonrpc_line(lines: &mut tokio::io::Lines>) -> JSONRPCMessage + where + R: tokio::io::AsyncRead + Unpin, + { + let next_line = timeout(Duration::from_secs(1), lines.next_line()).await; + let line_result = match next_line { + Ok(line_result) => line_result, + Err(err) => panic!("timed out waiting for JSON-RPC line: {err}"), + }; + let maybe_line = match line_result { + Ok(maybe_line) => maybe_line, + Err(err) => panic!("failed to read JSON-RPC line: {err}"), + }; + let line = match maybe_line { + Some(line) => line, + None => panic!("server connection closed before JSON-RPC line arrived"), + }; + match serde_json::from_str::(&line) { + Ok(message) => message, + Err(err) => panic!("failed to parse JSON-RPC line: {err}"), + } + } + + async fn write_jsonrpc_line(writer: &mut W, message: JSONRPCMessage) + where + W: tokio::io::AsyncWrite + Unpin, + { + let encoded = match serde_json::to_string(&message) { + Ok(encoded) => encoded, + Err(err) => panic!("failed to encode JSON-RPC message: {err}"), + }; + if let Err(err) = writer.write_all(format!("{encoded}\n").as_bytes()).await { + panic!("failed to write JSON-RPC line: {err}"); + } + } + + #[tokio::test] + async fn inbound_request_span_stays_open_until_event_consumption() { + let span_exporter = InMemorySpanExporter::default(); + let tracer_provider = SdkTracerProvider::builder() + .with_simple_exporter(span_exporter.clone()) + .build(); + let subscriber = tracing_subscriber::registry().with( + tracing_opentelemetry::layer() + .with_tracer(tracer_provider.tracer("exec-server-test")) + .with_filter(filter_fn(codex_otel::OtelProvider::trace_export_filter)), + ); + let _subscriber = tracing::subscriber::set_default(subscriber); + tracing::callsite::rebuild_interest_cache(); + + let (outgoing_tx, _outgoing_rx) = tokio::sync::mpsc::channel(/*buffer*/ 1); + let (incoming_tx, incoming_rx) = tokio::sync::mpsc::channel(/*buffer*/ 1); + let (_disconnected_tx, disconnected_rx) = tokio::sync::watch::channel(/*init*/ false); + let connection = JsonRpcConnection { + outgoing_tx, + incoming_rx, + disconnected_rx, + task_handles: Vec::new(), + transport: JsonRpcTransport::Plain, + }; + let (_client, mut events_rx) = RpcClient::new(connection); + + incoming_tx + .send(JsonRpcConnectionEvent::message(JSONRPCMessage::Request( + JSONRPCRequest { + id: RequestId::Integer(1), + method: "test/callback".to_string(), + params: None, + trace: None, + }, + ))) + .await + .expect("queue inbound client request"); + timeout(Duration::from_secs(1), async { + while events_rx.is_empty() { + tokio::task::yield_now().await; + } + }) + .await + .expect("inbound request should enter the client event queue"); + assert!( + span_exporter + .get_finished_spans() + .expect("request span export") + .is_empty(), + "the request span must remain open until the client consumes the event" + ); + + let Some(RpcClientEvent::Request { + request, + request_span, + }) = events_rx.recv().await + else { + panic!("expected an inbound client request"); + }; + assert_eq!(request.method, "test/callback"); + request_span.record("otel.name", "test/callback"); + drop(request_span); + + tracer_provider.force_flush().expect("flush traces"); + let spans = span_exporter.get_finished_spans().expect("span export"); + assert!( + spans + .iter() + .any(|span| span.name.as_ref() == "test/callback"), + "the request span should cover the complete event queue wait" + ); + } + + #[tokio::test] + async fn rpc_client_matches_out_of_order_responses_by_request_id() { + let (client_stdin, server_reader) = tokio::io::duplex(4096); + let (mut server_writer, client_stdout) = tokio::io::duplex(4096); + let connection = + JsonRpcConnection::from_stdio(client_stdout, client_stdin, "test-rpc".to_string()); + let (client, _events_rx) = RpcClient::new(connection); + + let server = tokio::spawn(async move { + let mut lines = BufReader::new(server_reader).lines(); + + let first = read_jsonrpc_line(&mut lines).await; + let second = read_jsonrpc_line(&mut lines).await; + let (slow_request, fast_request) = match (first, second) { + ( + JSONRPCMessage::Request(first_request), + JSONRPCMessage::Request(second_request), + ) if first_request.method == "slow" && second_request.method == "fast" => { + (first_request, second_request) + } + ( + JSONRPCMessage::Request(first_request), + JSONRPCMessage::Request(second_request), + ) if first_request.method == "fast" && second_request.method == "slow" => { + (second_request, first_request) + } + _ => panic!("expected slow and fast requests"), + }; + + write_jsonrpc_line( + &mut server_writer, + JSONRPCMessage::Response(JSONRPCResponse { + id: fast_request.id, + result: serde_json::json!({ "value": "fast" }), + }), + ) + .await; + write_jsonrpc_line( + &mut server_writer, + JSONRPCMessage::Response(JSONRPCResponse { + id: slow_request.id, + result: serde_json::json!({ "value": "slow" }), + }), + ) + .await; + }); + + let slow_params = serde_json::json!({ "n": 1 }); + let fast_params = serde_json::json!({ "n": 2 }); + let (slow, fast) = tokio::join!( + client.call::<_, serde_json::Value>("slow", &slow_params), + client.call::<_, serde_json::Value>("fast", &fast_params), + ); + + let slow = slow.unwrap_or_else(|err| panic!("slow request failed: {err:?}")); + let fast = fast.unwrap_or_else(|err| panic!("fast request failed: {err:?}")); + assert_eq!(slow, serde_json::json!({ "value": "slow" })); + assert_eq!(fast, serde_json::json!({ "value": "fast" })); + + assert_eq!(client.pending_request_count().await, 0); + + if let Err(err) = server.await { + panic!("server task failed: {err}"); + } + } + + #[tokio::test(start_paused = true)] + async fn rpc_client_call_has_no_implicit_deadline() { + let (client_stdin, server_reader) = tokio::io::duplex(4096); + let (mut server_writer, client_stdout) = tokio::io::duplex(4096); + let connection = + JsonRpcConnection::from_stdio(client_stdout, client_stdin, "test-rpc".to_string()); + let (client, _events_rx) = RpcClient::new(connection); + let mut lines = BufReader::new(server_reader).lines(); + + let params = serde_json::json!({}); + let call = client.call::<_, serde_json::Value>("slow", ¶ms); + tokio::pin!(call); + assert!(futures::poll!(call.as_mut()).is_pending()); + let request = match read_jsonrpc_line(&mut lines).await { + JSONRPCMessage::Request(request) => request, + other => panic!("expected JSON-RPC request, got {other:?}"), + }; + + tokio::time::advance(Duration::from_secs(61)).await; + assert!(futures::poll!(call.as_mut()).is_pending()); + + let expected = serde_json::json!({ "value": "done" }); + write_jsonrpc_line( + &mut server_writer, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: expected.clone(), + }), + ) + .await; + assert_eq!(call.await.expect("RPC response"), expected); + } + + #[tokio::test] + async fn rpc_client_timeout_removes_pending_request() { + let (client_stdin, server_reader) = tokio::io::duplex(4096); + let (server_writer, client_stdout) = tokio::io::duplex(4096); + let (release_server_tx, release_server_rx) = tokio::sync::oneshot::channel(); + let connection = + JsonRpcConnection::from_stdio(client_stdout, client_stdin, "test-rpc".to_string()); + let (client, _events_rx) = RpcClient::new(connection); + + let server = tokio::spawn(async move { + let mut lines = BufReader::new(server_reader).lines(); + let request = read_jsonrpc_line(&mut lines).await; + assert!(matches!(request, JSONRPCMessage::Request(_))); + let _server_writer = server_writer; + let _ = release_server_rx.await; + }); + + let call_timeout = Duration::from_millis(10); + let result = client + .call_with_timeout::<_, serde_json::Value>("slow", &serde_json::json!({}), call_timeout) + .await; + assert!(matches!( + result, + Err(super::RpcCallError::TimedOut { method, timeout }) + if method == "slow" && timeout == call_timeout + )); + assert_eq!(client.pending_request_count().await, 0); + + let _ = release_server_tx.send(()); + if let Err(err) = server.await { + panic!("server task failed: {err}"); + } + } + + #[tokio::test] + async fn rpc_client_bounds_in_flight_calls_and_preserves_cleanup() { + let (outgoing_tx, outgoing_rx) = tokio::sync::mpsc::channel(/*buffer*/ 1); + outgoing_tx + .send(JSONRPCMessage::Notification(JSONRPCNotification { + method: "blocker".to_string(), + params: None, + })) + .await + .expect("outbound queue should accept the blocker"); + let (incoming_tx, incoming_rx) = tokio::sync::mpsc::channel(MAX_IN_FLIGHT_REGULAR_CALLS); + let (_disconnected_tx, disconnected_rx) = tokio::sync::watch::channel(/*init*/ false); + let connection = JsonRpcConnection { + outgoing_tx, + incoming_rx, + disconnected_rx, + task_handles: Vec::new(), + transport: JsonRpcTransport::Plain, + }; + let (client, _events_rx) = RpcClient::new(connection); + let client = Arc::new(client); + let mut calls = JoinSet::new(); + + for index in 0..MAX_IN_FLIGHT_REGULAR_CALLS { + let client = Arc::clone(&client); + calls.spawn(async move { + client + .call::<_, serde_json::Value>("pending", &serde_json::json!({ "index": index })) + .await + }); + } + timeout(Duration::from_secs(1), async { + while client.pending_request_count().await < MAX_IN_FLIGHT_REGULAR_CALLS { + tokio::task::yield_now().await; + } + }) + .await + .expect("pending requests should reach the regular limit"); + + for request_id in 1..=MAX_IN_FLIGHT_REGULAR_CALLS { + incoming_tx + .send(JsonRpcConnectionEvent::Message(JSONRPCMessage::Response( + JSONRPCResponse { + id: RequestId::Integer( + i64::try_from(request_id).expect("request id should fit in i64"), + ), + result: serde_json::json!({}), + }, + ))) + .await + .expect("reader should accept the spoofed response"); + } + timeout(Duration::from_secs(1), async { + while client.pending_request_count().await != 0 { + tokio::task::yield_now().await; + } + }) + .await + .expect("spoofed responses should drain response routing"); + + let params = serde_json::json!({}); + let overflow = client.call::<_, serde_json::Value>("overflow", ¶ms); + tokio::pin!(overflow); + assert!(matches!( + futures::poll!(overflow.as_mut()), + std::task::Poll::Ready(Err(RpcCallError::PendingRequestLimitExceeded { limit })) + if limit == MAX_IN_FLIGHT_REGULAR_CALLS + )); + + let cleanup_client = Arc::clone(&client); + calls.spawn(async move { + let params = serde_json::json!({}); + cleanup_client + .call_for_cleanup::<_, serde_json::Value>("cleanup", ¶ms) + .await + }); + timeout(Duration::from_secs(1), async { + while client.pending_request_count().await != 1 { + tokio::task::yield_now().await; + } + }) + .await + .expect("cleanup request should use the reserved capacity"); + + let cleanup_params = serde_json::json!({}); + let cleanup_overflow = timeout( + Duration::from_secs(1), + client.call_for_cleanup::<_, serde_json::Value>("cleanup-overflow", &cleanup_params), + ) + .await + .expect("cleanup circuit breaker should not block"); + assert!(matches!(cleanup_overflow, Err(RpcCallError::Closed))); + assert!(client.is_disconnected()); + assert_eq!(client.pending_request_count().await, 0); + + drop(outgoing_rx); + while let Some(call) = calls.join_next().await { + assert!(matches!( + call.expect("pending call task should join"), + Err(RpcCallError::Closed) + )); + } + } + + #[tokio::test(flavor = "current_thread")] + async fn rpc_client_propagates_current_trace_context() { + let span_exporter = InMemorySpanExporter::default(); + let tracer_provider = SdkTracerProvider::builder() + .with_simple_exporter(span_exporter) + .build(); + let tracer = tracer_provider.tracer("exec-server-test"); + let subscriber = tracing_subscriber::registry().with( + tracing_opentelemetry::layer() + .with_tracer(tracer) + .with_filter(filter_fn(codex_otel::OtelProvider::trace_export_filter)), + ); + let _subscriber_guard = tracing::subscriber::set_default(subscriber); + tracing::callsite::rebuild_interest_cache(); + let parent_span = tracing::info_span!("outbound-parent"); + let expected_trace = codex_otel::span_w3c_trace_context(&parent_span) + .expect("parent span should have trace context"); + + let (client_stdin, server_reader) = tokio::io::duplex(4096); + let (mut server_writer, client_stdout) = tokio::io::duplex(4096); + let connection = + JsonRpcConnection::from_stdio(client_stdout, client_stdin, "test-rpc".to_string()); + let (client, _events_rx) = RpcClient::new(connection); + + let server = tokio::spawn(async move { + let mut lines = BufReader::new(server_reader).lines(); + let request = match read_jsonrpc_line(&mut lines).await { + JSONRPCMessage::Request(request) => request, + other => panic!("expected JSON-RPC request, got {other:?}"), + }; + write_jsonrpc_line( + &mut server_writer, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id.clone(), + result: serde_json::json!({}), + }), + ) + .await; + request.trace + }); + + let response = client + .call::<_, serde_json::Value>("traced", &serde_json::json!({})) + .instrument(parent_span) + .await + .expect("RPC response"); + assert_eq!(response, serde_json::json!({})); + let trace = server.await.expect("server task").expect("trace context"); + let expected_traceparent = expected_trace + .traceparent + .as_deref() + .expect("parent traceparent"); + let traceparent = trace.traceparent.as_deref().expect("request traceparent"); + let expected_parts = expected_traceparent.split('-').collect::>(); + let parts = traceparent.split('-').collect::>(); + assert_eq!(parts[1], expected_parts[1]); + assert_ne!(parts[2], expected_parts[2]); + assert_eq!(trace.tracestate, expected_trace.tracestate); + } +} diff --git a/codex-rs/exec-server/src/rpc_client_metrics_tests.rs b/codex-rs/exec-server/src/rpc_client_metrics_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..e58d504c4382a3f758832e00ea8f7f77cca2e996 --- /dev/null +++ b/codex-rs/exec-server/src/rpc_client_metrics_tests.rs @@ -0,0 +1,252 @@ +//! Exercise caller-side counters through the RPC transport, including failures and cancellation. + +use std::collections::BTreeMap; +use std::time::Duration; + +use codex_otel::MetricsClient; +use codex_otel::MetricsConfig; +use opentelemetry_sdk::metrics::InMemoryMetricExporter; +use opentelemetry_sdk::metrics::data::AggregatedMetrics; +use opentelemetry_sdk::metrics::data::MetricData; +use pretty_assertions::assert_eq; +use serde_json::Value; +use tokio::sync::mpsc; +use tokio::sync::watch; + +use super::RpcCallError; +use super::RpcClient; +use super::RpcClientEvent; +use crate::connection::JsonRpcConnection; +use crate::connection::JsonRpcConnectionEvent; +use crate::connection::JsonRpcTransport; +use crate::protocol::FS_READ_FILE_METHOD; +use crate::protocol::INITIALIZED_METHOD; +use crate::protocol::JSONRPCMessage; +use crate::protocol::JSONRPCResponse; +use crate::protocol::RequestId; + +struct Harness { + client: RpcClient, + metrics: MetricsClient, + outgoing: mpsc::Receiver, + incoming: mpsc::Sender, + _events: mpsc::Receiver, + _disconnected: watch::Sender, +} + +impl Harness { + fn new() -> Self { + let metrics = MetricsClient::new( + MetricsConfig::in_memory( + "test", + "exec-server-client-test", + env!("CARGO_PKG_VERSION"), + InMemoryMetricExporter::default(), + ) + .with_runtime_reader(), + ) + .expect("metrics client"); + let (outgoing_tx, outgoing) = mpsc::channel(/*buffer*/ 8); + let (incoming, incoming_rx) = mpsc::channel(/*buffer*/ 8); + let (disconnected, disconnected_rx) = watch::channel(/*init*/ false); + let (mut client, events) = RpcClient::new(JsonRpcConnection { + outgoing_tx, + incoming_rx, + disconnected_rx, + task_handles: Vec::new(), + transport: JsonRpcTransport::Plain, + }); + client.metrics = Some(metrics.clone()); + Self { + client, + metrics, + outgoing, + incoming, + _events: events, + _disconnected: disconnected, + } + } + + fn counts(&self) -> BTreeMap { + let snapshot = self.metrics.snapshot().expect("metrics snapshot"); + let mut counts = BTreeMap::new(); + for metric in snapshot + .scope_metrics() + .flat_map(opentelemetry_sdk::metrics::data::ScopeMetrics::metrics) + .filter(|metric| metric.name() == "exec_server_client_requests_total") + { + let AggregatedMetrics::U64(MetricData::Sum(sum)) = metric.data() else { + panic!("client request count should be a u64 sum"); + }; + for point in sum.data_points() { + let attributes = point + .attributes() + .map(|attribute| { + ( + attribute.key.as_str(), + attribute.value.as_str().into_owned(), + ) + }) + .collect::>(); + assert_eq!(attributes.len(), 1, "only the method is labeled"); + assert_eq!(attributes[0].0, "method"); + *counts.entry(attributes[0].1.clone()).or_default() += point.value(); + } + } + counts + } +} + +#[derive(Clone, Copy)] +enum CallKind { + Regular, + Untraced, + WithTimeout, + Cleanup, +} + +#[tokio::test] +async fn each_request_entry_point_counts_once_and_preserves_response() { + for kind in [ + CallKind::Regular, + CallKind::Untraced, + CallKind::WithTimeout, + CallKind::Cleanup, + ] { + let mut harness = Harness::new(); + let params = serde_json::json!({"path": "/sensitive-test-path"}); + let request = async { + match kind { + CallKind::Regular => { + harness + .client + .call::<_, Value>(FS_READ_FILE_METHOD, ¶ms) + .await + } + CallKind::Untraced => { + harness + .client + .call_untraced::<_, Value>(FS_READ_FILE_METHOD, ¶ms) + .await + } + CallKind::WithTimeout => { + harness + .client + .call_with_timeout::<_, Value>( + FS_READ_FILE_METHOD, + ¶ms, + Duration::from_secs(1), + ) + .await + } + CallKind::Cleanup => { + harness + .client + .call_for_cleanup::<_, Value>(FS_READ_FILE_METHOD, ¶ms) + .await + } + } + }; + let server = async { + let Some(JSONRPCMessage::Request(request)) = harness.outgoing.recv().await else { + panic!("expected request"); + }; + assert_eq!(request.method, FS_READ_FILE_METHOD); + assert_eq!(request.params, Some(params.clone())); + harness + .incoming + .send(JsonRpcConnectionEvent::Message(JSONRPCMessage::Response( + JSONRPCResponse { + id: request.id, + result: serde_json::json!({"ok": true}), + }, + ))) + .await + .expect("response accepted"); + }; + let (response, ()) = tokio::join!(request, server); + assert_eq!( + response.expect("RPC response"), + serde_json::json!({"ok": true}) + ); + assert_eq!( + harness.counts(), + BTreeMap::from([(FS_READ_FILE_METHOD.to_string(), 1)]) + ); + } +} + +#[tokio::test] +async fn local_rejections_and_closed_transport_count_as_attempts() { + let harness = Harness::new(); + let slots = harness + .client + .shared_call_slots + .acquire_many(super::MAX_IN_FLIGHT_REGULAR_CALLS as u32) + .await + .expect("occupy slots"); + let rejected = harness + .client + .call::<_, Value>(FS_READ_FILE_METHOD, &()) + .await; + assert!(matches!( + rejected, + Err(RpcCallError::PendingRequestLimitExceeded { .. }) + )); + drop(slots); + harness.client.close_transport().await; + let closed = harness + .client + .call::<_, Value>(FS_READ_FILE_METHOD, &()) + .await; + assert!(matches!(closed, Err(RpcCallError::Closed))); + assert_eq!( + harness.counts(), + BTreeMap::from([(FS_READ_FILE_METHOD.to_string(), 2)]) + ); +} + +#[tokio::test(start_paused = true)] +async fn timeout_and_cancellation_count_as_attempts() { + let mut harness = Harness::new(); + let timed_out = harness + .client + .call_with_timeout::<_, Value>(FS_READ_FILE_METHOD, &(), Duration::from_secs(1)) + .await; + assert!(matches!(timed_out, Err(RpcCallError::TimedOut { .. }))); + harness.outgoing.recv().await.expect("timed out request"); + let mut cancelled = Box::pin(harness.client.call::<_, Value>(FS_READ_FILE_METHOD, &())); + assert!(futures::poll!(cancelled.as_mut()).is_pending()); + harness.outgoing.recv().await.expect("cancelled request"); + drop(cancelled); + assert_eq!( + harness.counts(), + BTreeMap::from([(FS_READ_FILE_METHOD.to_string(), 2)]) + ); +} + +#[tokio::test] +async fn notifications_responses_and_disabled_metrics_do_not_record_attempts() { + let mut harness = Harness::new(); + harness + .client + .notify(INITIALIZED_METHOD, &()) + .await + .expect("notification"); + harness + .client + .respond(RequestId::Integer(1), &()) + .await + .expect("response"); + assert_eq!(harness.counts(), BTreeMap::new()); + harness.client.metrics = None; + harness.client.close_transport().await; + assert!(matches!( + harness + .client + .call::<_, Value>(FS_READ_FILE_METHOD, &()) + .await, + Err(RpcCallError::Closed) + )); + assert_eq!(harness.counts(), BTreeMap::new()); +} diff --git a/codex-rs/exec-server/src/rpc_server_requests.rs b/codex-rs/exec-server/src/rpc_server_requests.rs new file mode 100644 index 0000000000000000000000000000000000000000..b7d158802b8ecaa49a9055e29848b99e4d194e62 --- /dev/null +++ b/codex-rs/exec-server/src/rpc_server_requests.rs @@ -0,0 +1,170 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::atomic::AtomicI64; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::RequestId; +use serde::Serialize; +use serde::de::DeserializeOwned; +use serde_json::Value; +use tokio::sync::Semaphore; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::time::timeout; +use tokio_util::sync::CancellationToken; + +use crate::rpc::RpcCallError; +use crate::rpc::RpcServerOutboundMessage; + +pub(crate) const MAX_IN_FLIGHT_SERVER_CALLS: usize = 256; + +type PendingRequest = oneshot::Sender>; + +#[derive(Clone)] +pub(crate) struct RpcServerRequestSender { + inner: Arc, +} + +struct RpcServerRequestSenderInner { + outgoing_tx: mpsc::Sender, + pending: Mutex>, + call_slots: Semaphore, + next_request_id: AtomicI64, + closed: CancellationToken, +} + +struct PendingServerRequestGuard { + inner: Arc, + request_id: RequestId, +} + +impl Drop for PendingServerRequestGuard { + fn drop(&mut self) { + self.inner + .pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .remove(&self.request_id); + } +} + +impl RpcServerRequestSender { + pub(crate) fn new(outgoing_tx: mpsc::Sender) -> Self { + Self { + inner: Arc::new(RpcServerRequestSenderInner { + outgoing_tx, + pending: Mutex::new(HashMap::new()), + call_slots: Semaphore::new(MAX_IN_FLIGHT_SERVER_CALLS), + next_request_id: AtomicI64::new(1), + closed: CancellationToken::new(), + }), + } + } + + pub(crate) async fn call_with_timeout( + &self, + method: &str, + params: &P, + call_timeout: Duration, + ) -> Result + where + P: Serialize, + T: DeserializeOwned, + { + let _call_slot = self.inner.call_slots.try_acquire().map_err(|_| { + RpcCallError::PendingRequestLimitExceeded { + limit: MAX_IN_FLIGHT_SERVER_CALLS, + } + })?; + let params = serde_json::to_value(params).map_err(RpcCallError::Json)?; + let request_id = + RequestId::Integer(self.inner.next_request_id.fetch_add(1, Ordering::SeqCst)); + let (response_tx, response_rx) = oneshot::channel(); + { + let mut pending = self + .inner + .pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if self.inner.closed.is_cancelled() { + return Err(RpcCallError::Closed); + } + pending.insert(request_id.clone(), response_tx); + } + let _pending = PendingServerRequestGuard { + inner: Arc::clone(&self.inner), + request_id: request_id.clone(), + }; + let request = RpcServerOutboundMessage::Request(JSONRPCRequest { + id: request_id, + method: method.to_string(), + params: Some(params), + trace: codex_otel::current_span_w3c_trace_context(), + }); + + let response = timeout(call_timeout, async { + tokio::select! { + biased; + _ = self.inner.closed.cancelled() => return Err(RpcCallError::Closed), + result = self.inner.outgoing_tx.send(request) => { + result.map_err(|_| RpcCallError::Closed)?; + } + } + response_rx.await.map_err(|_| RpcCallError::Closed)? + }) + .await + .map_err(|_| RpcCallError::TimedOut { + method: method.to_string(), + timeout: call_timeout, + })??; + serde_json::from_value(response).map_err(RpcCallError::Json) + } + + pub(crate) fn complete( + &self, + request_id: RequestId, + result: Result, + ) -> bool { + if let Some(pending) = self + .inner + .pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .remove(&request_id) + { + let _ = pending.send(result); + true + } else { + matches!( + request_id, + RequestId::Integer(id) + if id > 0 && id < self.inner.next_request_id.load(Ordering::Acquire) + ) + } + } + + pub(crate) fn close(&self) { + self.inner.closed.cancel(); + let pending = { + let mut pending = self + .inner + .pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + pending + .drain() + .map(|(_, pending)| pending) + .collect::>() + }; + for pending in pending { + let _ = pending.send(Err(RpcCallError::Closed)); + } + } +} + +#[cfg(test)] +#[path = "rpc_server_requests_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/rpc_server_requests_tests.rs b/codex-rs/exec-server/src/rpc_server_requests_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..05692c32aaaecb87ba6505de570ce13a83a5843a --- /dev/null +++ b/codex-rs/exec-server/src/rpc_server_requests_tests.rs @@ -0,0 +1,178 @@ +use std::sync::Arc; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::RequestId; +use pretty_assertions::assert_eq; +use tokio::sync::mpsc; +use tokio::task::JoinSet; +use tokio::time::timeout; + +use super::MAX_IN_FLIGHT_SERVER_CALLS; +use super::RpcServerRequestSender; +use crate::rpc::RpcCallError; +use crate::rpc::RpcServerOutboundMessage; + +async fn receive_server_request( + outgoing_rx: &mut mpsc::Receiver, +) -> JSONRPCRequest { + let message = timeout(Duration::from_secs(1), outgoing_rx.recv()) + .await + .expect("server request should arrive") + .expect("server request"); + match message { + RpcServerOutboundMessage::Request(request) => request, + other => panic!("expected server request, got {other:?}"), + } +} + +impl RpcServerRequestSender { + pub(crate) fn pending_request_count(&self) -> usize { + self.inner + .pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .len() + } +} + +#[tokio::test] +async fn rpc_server_sender_matches_out_of_order_responses_by_request_id() { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(/*buffer*/ 8); + let requests = RpcServerRequestSender::new(outgoing_tx); + let slow_requests = requests.clone(); + let slow = tokio::spawn(async move { + slow_requests + .call_with_timeout::<_, serde_json::Value>( + "slow", + &serde_json::json!({ "n": 1 }), + Duration::from_secs(1), + ) + .await + }); + let fast_requests = requests.clone(); + let fast = tokio::spawn(async move { + fast_requests + .call_with_timeout::<_, serde_json::Value>( + "fast", + &serde_json::json!({ "n": 2 }), + Duration::from_secs(1), + ) + .await + }); + + let first = receive_server_request(&mut outgoing_rx).await; + let second = receive_server_request(&mut outgoing_rx).await; + let (slow_request, fast_request) = if first.method == "slow" { + (first, second) + } else { + (second, first) + }; + requests.complete(fast_request.id, Ok(serde_json::json!({ "value": "fast" }))); + requests.complete(slow_request.id, Ok(serde_json::json!({ "value": "slow" }))); + + assert_eq!( + slow.await.expect("slow task").expect("slow server request"), + serde_json::json!({ "value": "slow" }) + ); + assert_eq!( + fast.await.expect("fast task").expect("fast server request"), + serde_json::json!({ "value": "fast" }) + ); + assert_eq!(requests.pending_request_count(), 0); +} + +#[tokio::test] +async fn rpc_server_sender_preserves_response_received_before_close() { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(/*buffer*/ 1); + let requests = RpcServerRequestSender::new(outgoing_tx); + let caller = requests.clone(); + let call = tokio::spawn(async move { + caller + .call_with_timeout::<_, serde_json::Value>( + "ordered", + &serde_json::json!({}), + Duration::from_secs(1), + ) + .await + }); + + let request = receive_server_request(&mut outgoing_rx).await; + requests.complete(request.id, Ok(serde_json::json!({ "value": "accepted" }))); + requests.close(); + + assert_eq!( + call.await + .expect("server request task should join") + .expect("server request should preserve its response"), + serde_json::json!({ "value": "accepted" }) + ); + assert_eq!(requests.pending_request_count(), 0); +} + +#[tokio::test(start_paused = true)] +async fn rpc_server_sender_timeout_removes_pending_request() { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(/*buffer*/ 1); + let requests = RpcServerRequestSender::new(outgoing_tx); + let call_timeout = Duration::from_secs(1); + let params = serde_json::json!({}); + let call = requests.call_with_timeout::<_, serde_json::Value>("slow", ¶ms, call_timeout); + tokio::pin!(call); + assert!(futures::poll!(call.as_mut()).is_pending()); + let request = receive_server_request(&mut outgoing_rx).await; + + tokio::time::advance(call_timeout).await; + assert!(matches!( + call.await, + Err(RpcCallError::TimedOut { method, timeout }) + if method == "slow" && timeout == call_timeout + )); + assert_eq!(requests.pending_request_count(), 0); + assert!(requests.complete(request.id, Ok(serde_json::Value::Null))); + assert!(!requests.complete(RequestId::Integer(2), Ok(serde_json::Value::Null))); +} + +#[tokio::test] +async fn rpc_server_sender_bounds_and_drains_pending_requests_on_close() { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(MAX_IN_FLIGHT_SERVER_CALLS); + let requests = Arc::new(RpcServerRequestSender::new(outgoing_tx)); + let mut calls = JoinSet::new(); + for index in 0..MAX_IN_FLIGHT_SERVER_CALLS { + let requests = Arc::clone(&requests); + calls.spawn(async move { + requests + .call_with_timeout::<_, serde_json::Value>( + "pending", + &serde_json::json!({ "index": index }), + Duration::from_secs(30), + ) + .await + }); + } + for _ in 0..MAX_IN_FLIGHT_SERVER_CALLS { + receive_server_request(&mut outgoing_rx).await; + } + assert_eq!(requests.pending_request_count(), MAX_IN_FLIGHT_SERVER_CALLS); + + let overflow = requests + .call_with_timeout::<_, serde_json::Value>( + "overflow", + &serde_json::json!({}), + Duration::from_secs(1), + ) + .await; + assert!(matches!( + overflow, + Err(RpcCallError::PendingRequestLimitExceeded { limit }) + if limit == MAX_IN_FLIGHT_SERVER_CALLS + )); + + requests.close(); + assert_eq!(requests.pending_request_count(), 0); + while let Some(call) = calls.join_next().await { + assert!(matches!( + call.expect("pending server request task should join"), + Err(RpcCallError::Closed) + )); + } +} diff --git a/codex-rs/exec-server/src/runtime_paths.rs b/codex-rs/exec-server/src/runtime_paths.rs new file mode 100644 index 0000000000000000000000000000000000000000..0947ed87728493bab97f71b92ba01620276d322b --- /dev/null +++ b/codex-rs/exec-server/src/runtime_paths.rs @@ -0,0 +1,58 @@ +use std::path::PathBuf; + +use codex_utils_absolute_path::AbsolutePathBuf; + +/// Paths and sandbox settings initialized when creating an executor. +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct ExecServerRuntimePaths { + /// Stable path to the Codex executable used to launch hidden helper modes. + pub codex_self_exe: AbsolutePathBuf, + /// Path to the Linux sandbox helper alias used when the platform sandbox + /// needs to re-enter Codex by argv0. + pub codex_linux_sandbox_exe: Option, + /// User-config opt-out of writable-root symlink checks beneath this host's home. + #[cfg(target_os = "macos")] + pub allowed_symlinked_codex_home: Option, +} + +impl ExecServerRuntimePaths { + pub fn from_optional_paths( + codex_self_exe: Option, + codex_linux_sandbox_exe: Option, + ) -> std::io::Result { + let codex_self_exe = codex_self_exe.ok_or_else(|| { + std::io::Error::new( + std::io::ErrorKind::InvalidInput, + "Codex executable path is not configured", + ) + })?; + Self::new(codex_self_exe, codex_linux_sandbox_exe) + } + + pub fn new( + codex_self_exe: PathBuf, + codex_linux_sandbox_exe: Option, + ) -> std::io::Result { + Ok(Self { + codex_self_exe: absolute_path(codex_self_exe)?, + codex_linux_sandbox_exe: codex_linux_sandbox_exe.map(absolute_path).transpose()?, + #[cfg(target_os = "macos")] + allowed_symlinked_codex_home: None, + }) + } + + /// Applies the symlink opt-in resolved by the execution host's config loader. + #[cfg(target_os = "macos")] + pub fn with_allowed_symlinked_codex_home( + mut self, + allowed_symlinked_codex_home: Option, + ) -> Self { + self.allowed_symlinked_codex_home = allowed_symlinked_codex_home; + self + } +} + +fn absolute_path(path: PathBuf) -> std::io::Result { + AbsolutePathBuf::from_absolute_path(path.as_path()) + .map_err(|err| std::io::Error::new(std::io::ErrorKind::InvalidInput, err)) +} diff --git a/codex-rs/exec-server/src/sandbox_selection.rs b/codex-rs/exec-server/src/sandbox_selection.rs new file mode 100644 index 0000000000000000000000000000000000000000..0815067c4917f35133230f4a5feca8bf5cbe5bbd --- /dev/null +++ b/codex-rs/exec-server/src/sandbox_selection.rs @@ -0,0 +1,30 @@ +//! Resolves an executor sandbox context to a concrete local sandbox implementation. + +use codex_file_system::FileSystemSandboxContext; +use codex_file_system::WindowsSandboxSelection; +use codex_protocol::config_types::WindowsSandboxLevel; +use codex_protocol::models::PermissionProfile; +use codex_sandboxing::SandboxManager; +use codex_sandboxing::SandboxType; +use codex_sandboxing::SandboxablePreference; + +pub(crate) fn select_sandbox( + manager: &SandboxManager, + permission_profile: &PermissionProfile, + sandbox_context: &FileSystemSandboxContext, + has_managed_network_requirements: bool, +) -> (SandboxType, Option) { + let windows_sandbox_level = match sandbox_context.windows_sandbox_selection { + WindowsSandboxSelection::Disabled => WindowsSandboxLevel::Disabled, + WindowsSandboxSelection::RestrictedToken => WindowsSandboxLevel::RestrictedToken, + WindowsSandboxSelection::Elevated => WindowsSandboxLevel::Elevated, + WindowsSandboxSelection::Mxc => return (SandboxType::WindowsMxc, None), + }; + let sandbox_type = manager.select_initial( + permission_profile, + SandboxablePreference::Require, + windows_sandbox_level, + has_managed_network_requirements, + ); + (sandbox_type, Some(windows_sandbox_level)) +} diff --git a/codex-rs/exec-server/src/sandboxed_file_open.rs b/codex-rs/exec-server/src/sandboxed_file_open.rs new file mode 100644 index 0000000000000000000000000000000000000000..9cb4c75537e28f90b7ad80575093c5fc14eb0149 --- /dev/null +++ b/codex-rs/exec-server/src/sandboxed_file_open.rs @@ -0,0 +1,212 @@ +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_sandboxing::SandboxExecRequest; +use codex_utils_path_uri::PathUri; +use tokio::io; + +use crate::fs_helper::FsHelperOpenResponse; +use crate::fs_helper::FsHelperPayload; +use crate::fs_helper::FsHelperRequest; +use crate::fs_helper::FsHelperResponse; +#[cfg(windows)] +use crate::fs_sandbox::drain_helper_stderr; +use crate::fs_sandbox::io_error; +#[cfg(windows)] +use crate::fs_sandbox::read_helper_response; +#[cfg(windows)] +use crate::fs_sandbox::reap_helper_after_response; +use crate::fs_sandbox::spawn_command; +#[cfg(unix)] +use crate::fs_sandbox::wait_for_helper_output; +use crate::protocol::FsReadFileParams; +use crate::rpc::internal_error; +use crate::rpc::invalid_request; + +pub(crate) async fn open( + command: SandboxExecRequest, + path: PathUri, +) -> Result { + let request = serde_json::to_vec(&FsHelperRequest::Open(FsReadFileParams { + path, + follow_symlinks: None, + sandbox: None, + })) + .map_err(|error| internal_error(format!("invalid fs sandbox helper request: {error}")))?; + open_platform(command, request).await +} + +fn open_response(response: &[u8]) -> Result { + match serde_json::from_slice(response).map_err(|error| { + internal_error(format!("invalid fs sandbox helper open response: {error}")) + })? { + FsHelperResponse::Ok(FsHelperPayload::Open(response)) => Ok(response), + FsHelperResponse::Ok(_) => Err(invalid_request( + "invalid fs sandbox helper open response".to_string(), + )), + FsHelperResponse::Error(error) => Err(error), + } +} + +// Unix passes the opened fd over the helper's stdin socket. +#[cfg(unix)] +async fn open_platform( + command: SandboxExecRequest, + request: Vec, +) -> Result { + use std::io::Write; + use std::os::fd::OwnedFd; + use std::os::unix::net::UnixStream; + + let (mut receiver, sender) = UnixStream::pair().map_err(io_error)?; + let sender: OwnedFd = sender.into(); + let child = spawn_command(command, std::process::Stdio::from(sender))?; + receiver.write_all(&request).map_err(io_error)?; + receiver + .shutdown(std::net::Shutdown::Write) + .map_err(io_error)?; + + let output = wait_for_helper_output(child).await?; + open_response(&output.stdout)?; + let descriptor = receive_file_descriptor(&receiver).map_err(io_error)?; + Ok(tokio::fs::File::from_std(std::fs::File::from(descriptor))) +} + +// Windows duplicates the helper's handle before letting it exit. +#[cfg(windows)] +async fn open_platform( + command: SandboxExecRequest, + mut request: Vec, +) -> Result { + use tokio::io::AsyncWriteExt; + + let mut child = spawn_command(command, std::process::Stdio::piped())?; + let mut stdin = child + .stdin + .take() + .ok_or_else(|| internal_error("missing fs sandbox helper stdin".to_string()))?; + let stdout = child + .stdout + .take() + .ok_or_else(|| internal_error("missing fs sandbox helper stdout".to_string()))?; + request.push(b'\n'); + stdin.write_all(&request).await.map_err(io_error)?; + stdin.flush().await.map_err(io_error)?; + let stderr = drain_helper_stderr(&mut child); + + let result = async { + let response = read_helper_response(stdout).await?; + let response = open_response(&response)?; + duplicate_file_handle(response.process_id, response.file_handle).map_err(io_error) + } + .await; + drop(stdin); + reap_helper_after_response(child, stderr).await?; + result.map(tokio::fs::File::from_std) +} + +// SCM_RIGHTS is Unix-only. +#[cfg(unix)] +pub(crate) fn transfer_file(file: &tokio::fs::File) -> io::Result<()> { + use rustix::net::SendAncillaryBuffer; + use rustix::net::SendAncillaryMessage; + use rustix::net::SendFlags; + use std::io::IoSlice; + use std::os::fd::AsFd; + + let descriptors = [file.as_fd()]; + let mut space = [std::mem::MaybeUninit::uninit(); rustix::cmsg_space!(ScmRights(1))]; + let mut control = SendAncillaryBuffer::new(&mut space); + if !control.push(SendAncillaryMessage::ScmRights(&descriptors)) { + return Err(io::Error::other("missing file-descriptor control header")); + } + if rustix::net::sendmsg( + std::io::stdin(), + &[IoSlice::new(&[0])], + &mut control, + SendFlags::empty(), + )? != 1 + { + return Err(io::Error::other( + "fs sandbox helper did not transfer its opened file descriptor", + )); + } + Ok(()) +} + +// File-descriptor passing is only available on Unix. +#[cfg(unix)] +fn receive_file_descriptor( + socket: &std::os::unix::net::UnixStream, +) -> io::Result { + use rustix::net::RecvAncillaryBuffer; + use rustix::net::RecvAncillaryMessage; + use rustix::net::RecvFlags; + use rustix::net::ReturnFlags; + use std::io::IoSliceMut; + + let mut byte = [0_u8]; + let mut buffers = [IoSliceMut::new(&mut byte)]; + let mut space = [std::mem::MaybeUninit::uninit(); rustix::cmsg_space!(ScmRights(1))]; + let mut control = RecvAncillaryBuffer::new(&mut space); + // Linux can set close-on-exec while receiving the fd. + #[cfg(target_os = "linux")] + let flags = RecvFlags::CMSG_CLOEXEC; + // Other Unix platforms need the non-atomic fcntl call below. + #[cfg(not(target_os = "linux"))] + let flags = RecvFlags::empty(); + let message = rustix::net::recvmsg(socket, &mut buffers, &mut control, flags)?; + if message.bytes != 1 || message.flags.contains(ReturnFlags::CTRUNC) { + return Err(io::Error::other("invalid file-descriptor control message")); + } + let descriptor = control + .drain() + .find_map(|message| match message { + RecvAncillaryMessage::ScmRights(mut descriptors) => descriptors.next(), + _ => None, + }) + .ok_or_else(|| io::Error::other("missing transferred file descriptor"))?; + // macOS cannot set this atomically, so the fd is briefly inheritable. + // Shell and filesystem helper launches close inherited fds to limit that race. + #[cfg(not(target_os = "linux"))] + rustix::io::fcntl_setfd(&descriptor, rustix::io::FdFlags::CLOEXEC)?; + Ok(descriptor) +} + +// Windows file handles must be duplicated across processes. +#[cfg(windows)] +fn duplicate_file_handle(process_id: u32, file_handle: u64) -> io::Result { + use std::os::windows::io::AsRawHandle; + use std::os::windows::io::FromRawHandle; + use std::os::windows::io::OwnedHandle; + use windows_sys::Win32::Foundation::DUPLICATE_SAME_ACCESS; + use windows_sys::Win32::Foundation::DuplicateHandle; + use windows_sys::Win32::Foundation::HANDLE; + use windows_sys::Win32::System::Threading::GetCurrentProcess; + use windows_sys::Win32::System::Threading::OpenProcess; + use windows_sys::Win32::System::Threading::PROCESS_DUP_HANDLE; + + // SAFETY: OpenProcess returns an owned handle or null on failure. + let process = unsafe { OpenProcess(PROCESS_DUP_HANDLE, 0, process_id) }; + if process == 0 { + return Err(io::Error::last_os_error()); + } + // SAFETY: The successful OpenProcess result is owned by this scope. + let process = unsafe { OwnedHandle::from_raw_handle(process as _) }; + let mut duplicated: HANDLE = 0; + // SAFETY: Both process handles remain valid and duplicated receives an owned file handle. + if unsafe { + DuplicateHandle( + process.as_raw_handle() as HANDLE, + file_handle as HANDLE, + GetCurrentProcess(), + &raw mut duplicated, + 0, + 0, + DUPLICATE_SAME_ACCESS, + ) + } == 0 + { + return Err(io::Error::last_os_error()); + } + // SAFETY: DuplicateHandle transferred ownership of the new file handle. + Ok(unsafe { std::fs::File::from_raw_handle(duplicated as _) }) +} diff --git a/codex-rs/exec-server/src/sandboxed_file_system.rs b/codex-rs/exec-server/src/sandboxed_file_system.rs new file mode 100644 index 0000000000000000000000000000000000000000..61c3c061267bb79b8fc2eef490501190dece8d0a --- /dev/null +++ b/codex-rs/exec-server/src/sandboxed_file_system.rs @@ -0,0 +1,466 @@ +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD; +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_utils_path_uri::PathUri; +use tokio::io; +use tokio_util::io::ReaderStream; + +use crate::CapabilityRootsDiscoverParams; +use crate::CapabilityRootsDiscoverResponse; +use crate::CopyOptions; +use crate::CreateDirectoryOptions; +use crate::ExecServerRuntimePaths; +use crate::ExecutorFileSystem; +use crate::ExecutorFileSystemFuture; +use crate::FILE_READ_CHUNK_SIZE; +use crate::FileMetadata; +use crate::FileSystemReadStream; +use crate::FileSystemResult; +use crate::FileSystemSandboxContext; +use crate::GetMetadataOptions; +use crate::ReadDirectoryEntry; +use crate::ReadFileOptions; +use crate::RemoveOptions; +use crate::WalkOptions; +use crate::WalkOutcome; +use crate::WriteFileOptions; +use crate::fs_helper::FsHelperPayload; +use crate::fs_helper::FsHelperRequest; +use crate::fs_sandbox::FileSystemSandboxRunner; +use crate::protocol::FsCanonicalizeParams; +use crate::protocol::FsCopyParams; +use crate::protocol::FsCreateDirectoryParams; +use crate::protocol::FsGetMetadataParams; +use crate::protocol::FsReadDirectoryParams; +use crate::protocol::FsReadFileParams; +use crate::protocol::FsRemoveParams; +use crate::protocol::FsWalkParams; +use crate::protocol::FsWriteFileParams; + +#[derive(Clone)] +pub struct SandboxedFileSystem { + sandbox_runner: FileSystemSandboxRunner, +} + +impl SandboxedFileSystem { + #[tracing::instrument( + name = "capability_roots.discover_v1", + skip_all, + fields(root_count = params.roots.len()) + )] + pub(crate) async fn discover_capability_roots( + &self, + params: CapabilityRootsDiscoverParams, + sandbox: &FileSystemSandboxContext, + ) -> FileSystemResult { + self.run_sandboxed(sandbox, FsHelperRequest::DiscoverCapabilityRoots(params)) + .await? + .expect_capability_roots_discover() + .map_err(map_sandbox_error) + } + + pub(crate) async fn open_file_for_read( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + let sandbox = require_platform_sandbox(sandbox)?; + validate_native_path(path)?; + let command = self + .sandbox_runner + .sandbox_command(sandbox) + .map_err(map_sandbox_error)?; + crate::sandboxed_file_open::open(command, path.clone()) + .await + .map_err(map_sandbox_error) + } + + pub fn new(runtime_paths: ExecServerRuntimePaths) -> Self { + Self { + sandbox_runner: FileSystemSandboxRunner::new(runtime_paths), + } + } + + async fn run_sandboxed( + &self, + sandbox: &FileSystemSandboxContext, + request: FsHelperRequest, + ) -> FileSystemResult { + self.sandbox_runner + .run(sandbox, request) + .await + .map_err(map_sandbox_error) + } +} + +impl SandboxedFileSystem { + async fn canonicalize( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + let sandbox = require_platform_sandbox(sandbox)?; + validate_native_path(path)?; + let response = self + .run_sandboxed( + sandbox, + FsHelperRequest::Canonicalize(FsCanonicalizeParams { + path: path.clone(), + sandbox: None, + }), + ) + .await? + .expect_canonicalize() + .map_err(map_sandbox_error)?; + Ok(response.path) + } + + async fn read_file( + &self, + path: &PathUri, + options: ReadFileOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult> { + let sandbox = require_platform_sandbox(sandbox)?; + validate_native_path(path)?; + let response = self + .run_sandboxed( + sandbox, + FsHelperRequest::ReadFile(FsReadFileParams { + path: path.clone(), + follow_symlinks: (!options.follow_symlinks).then_some(false), + sandbox: None, + }), + ) + .await? + .expect_read_file() + .map_err(map_sandbox_error)?; + STANDARD.decode(response.data_base64).map_err(|err| { + io::Error::new( + io::ErrorKind::InvalidData, + format!("fs/readFile returned invalid base64 dataBase64: {err}"), + ) + }) + } + + async fn write_file( + &self, + path: &PathUri, + contents: Vec, + options: WriteFileOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + let sandbox = require_platform_sandbox(sandbox)?; + validate_native_path(path)?; + self.run_sandboxed( + sandbox, + FsHelperRequest::WriteFile(FsWriteFileParams { + path: path.clone(), + data_base64: STANDARD.encode(contents), + follow_symlinks: (!options.follow_symlinks).then_some(false), + sandbox: None, + }), + ) + .await? + .expect_write_file() + .map_err(map_sandbox_error)?; + Ok(()) + } + + async fn create_directory( + &self, + path: &PathUri, + options: CreateDirectoryOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + let sandbox = require_platform_sandbox(sandbox)?; + validate_native_path(path)?; + self.run_sandboxed( + sandbox, + FsHelperRequest::CreateDirectory(FsCreateDirectoryParams { + path: path.clone(), + recursive: Some(options.recursive), + follow_symlinks: (!options.follow_symlinks).then_some(false), + sandbox: None, + }), + ) + .await? + .expect_create_directory() + .map_err(map_sandbox_error)?; + Ok(()) + } + + async fn get_metadata( + &self, + path: &PathUri, + options: GetMetadataOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + let sandbox = require_platform_sandbox(sandbox)?; + validate_native_path(path)?; + let response = self + .run_sandboxed( + sandbox, + FsHelperRequest::GetMetadata(FsGetMetadataParams { + path: path.clone(), + follow_symlinks: (!options.follow_symlinks).then_some(false), + sandbox: None, + }), + ) + .await? + .expect_get_metadata() + .map_err(map_sandbox_error)?; + Ok(FileMetadata { + is_directory: response.is_directory, + is_file: response.is_file, + is_symlink: response.is_symlink, + size: response.size, + created_at_ms: response.created_at_ms, + modified_at_ms: response.modified_at_ms, + }) + } + + async fn read_directory( + &self, + path: &PathUri, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult> { + let sandbox = require_platform_sandbox(sandbox)?; + validate_native_path(path)?; + let response = self + .run_sandboxed( + sandbox, + FsHelperRequest::ReadDirectory(FsReadDirectoryParams { + path: path.clone(), + sandbox: None, + }), + ) + .await? + .expect_read_directory() + .map_err(map_sandbox_error)?; + Ok(response + .entries + .into_iter() + .map(|entry| ReadDirectoryEntry { + file_name: entry.file_name, + is_directory: entry.is_directory, + is_file: entry.is_file, + }) + .collect()) + } + + async fn walk( + &self, + path: &PathUri, + options: WalkOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult { + let sandbox = require_platform_sandbox(sandbox)?; + validate_native_path(path)?; + let response = self + .run_sandboxed( + sandbox, + FsHelperRequest::Walk(FsWalkParams { + path: path.clone(), + options, + sandbox: None, + }), + ) + .await? + .expect_walk() + .map_err(map_sandbox_error)?; + Ok(response) + } + + async fn remove( + &self, + path: &PathUri, + remove_options: RemoveOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + let sandbox = require_platform_sandbox(sandbox)?; + validate_native_path(path)?; + self.run_sandboxed( + sandbox, + FsHelperRequest::Remove(FsRemoveParams { + path: path.clone(), + recursive: Some(remove_options.recursive), + force: Some(remove_options.force), + follow_symlinks: (!remove_options.follow_symlinks).then_some(false), + sandbox: None, + }), + ) + .await? + .expect_remove() + .map_err(map_sandbox_error)?; + Ok(()) + } + + async fn copy( + &self, + source_path: &PathUri, + destination_path: &PathUri, + options: CopyOptions, + sandbox: Option<&FileSystemSandboxContext>, + ) -> FileSystemResult<()> { + let sandbox = require_platform_sandbox(sandbox)?; + validate_native_path(source_path)?; + validate_native_path(destination_path)?; + self.run_sandboxed( + sandbox, + FsHelperRequest::Copy(FsCopyParams { + source_path: source_path.clone(), + destination_path: destination_path.clone(), + recursive: options.recursive, + sandbox: None, + }), + ) + .await? + .expect_copy() + .map_err(map_sandbox_error)?; + Ok(()) + } +} + +impl ExecutorFileSystem for SandboxedFileSystem { + fn canonicalize<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, PathUri> { + Box::pin(SandboxedFileSystem::canonicalize(self, path, sandbox)) + } + + fn read_file<'a>( + &'a self, + path: &'a PathUri, + options: ReadFileOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, Vec> { + Box::pin(SandboxedFileSystem::read_file(self, path, options, sandbox)) + } + + fn read_file_stream<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, FileSystemReadStream> { + Box::pin(async move { + let file = self.open_file_for_read(path, sandbox).await?; + Ok(FileSystemReadStream::new(ReaderStream::with_capacity( + file, + FILE_READ_CHUNK_SIZE, + ))) + }) + } + + fn write_file<'a>( + &'a self, + path: &'a PathUri, + contents: Vec, + options: WriteFileOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(SandboxedFileSystem::write_file( + self, path, contents, options, sandbox, + )) + } + + fn create_directory<'a>( + &'a self, + path: &'a PathUri, + options: CreateDirectoryOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(SandboxedFileSystem::create_directory( + self, path, options, sandbox, + )) + } + + fn get_metadata<'a>( + &'a self, + path: &'a PathUri, + options: GetMetadataOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, FileMetadata> { + Box::pin(SandboxedFileSystem::get_metadata( + self, path, options, sandbox, + )) + } + + fn read_directory<'a>( + &'a self, + path: &'a PathUri, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, Vec> { + Box::pin(SandboxedFileSystem::read_directory(self, path, sandbox)) + } + + fn walk<'a>( + &'a self, + path: &'a PathUri, + options: WalkOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, WalkOutcome> { + Box::pin(SandboxedFileSystem::walk(self, path, options, sandbox)) + } + + fn remove<'a>( + &'a self, + path: &'a PathUri, + remove_options: RemoveOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(SandboxedFileSystem::remove( + self, + path, + remove_options, + sandbox, + )) + } + + fn copy<'a>( + &'a self, + source_path: &'a PathUri, + destination_path: &'a PathUri, + options: CopyOptions, + sandbox: Option<&'a FileSystemSandboxContext>, + ) -> ExecutorFileSystemFuture<'a, ()> { + Box::pin(SandboxedFileSystem::copy( + self, + source_path, + destination_path, + options, + sandbox, + )) + } +} + +fn validate_native_path(path: &PathUri) -> FileSystemResult<()> { + path.to_abs_path().map(drop) +} + +fn require_platform_sandbox( + sandbox: Option<&FileSystemSandboxContext>, +) -> FileSystemResult<&FileSystemSandboxContext> { + sandbox + .filter(|sandbox| sandbox.should_run_in_sandbox()) + .ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidInput, + "sandboxed filesystem operations require ReadOnly or WorkspaceWrite sandbox policy", + ) + }) +} + +fn map_sandbox_error(error: JSONRPCErrorError) -> io::Error { + match error.code { + -32004 => io::Error::new(io::ErrorKind::NotFound, error.message), + -32600 => io::Error::new(io::ErrorKind::InvalidInput, error.message), + _ => io::Error::other(error.message), + } +} + +#[cfg(all(test, any(unix, windows)))] +#[path = "sandboxed_file_system_path_uri_tests.rs"] +mod path_uri_tests; diff --git a/codex-rs/exec-server/src/sandboxed_file_system_path_uri_tests.rs b/codex-rs/exec-server/src/sandboxed_file_system_path_uri_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..031d89080f21a369010161dec5a5971bdb26aa27 --- /dev/null +++ b/codex-rs/exec-server/src/sandboxed_file_system_path_uri_tests.rs @@ -0,0 +1,43 @@ +use codex_protocol::models::PermissionProfile; +use codex_protocol::permissions::FileSystemSandboxPolicy; +use codex_protocol::permissions::NetworkSandboxPolicy; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +use tokio::io; + +use super::*; + +#[tokio::test] +async fn sandboxed_file_system_rejects_non_native_uri_as_invalid_input() { + let runtime_paths = ExecServerRuntimePaths::new( + std::env::current_exe().expect("current exe"), + /*codex_linux_sandbox_exe*/ None, + ) + .expect("runtime paths"); + let file_system = SandboxedFileSystem::new(runtime_paths); + let sandbox = FileSystemSandboxContext::from_permission_profile( + PermissionProfile::from_runtime_permissions( + &FileSystemSandboxPolicy::restricted(Vec::new()), + NetworkSandboxPolicy::Restricted, + ), + ); + + let error = file_system + .read_file(&non_native_uri(), Default::default(), Some(&sandbox)) + .await + .expect_err("non-native URI should be rejected"); + + assert_eq!(error.kind(), io::ErrorKind::InvalidInput); +} + +fn non_native_uri() -> PathUri { + #[cfg(unix)] + let uri = "file://server/share/file.txt"; + #[cfg(windows)] + let uri = "file:///usr/local/file.txt"; + + match PathUri::parse(uri) { + Ok(uri) => uri, + Err(err) => panic!("valid non-native URI should parse: {err}"), + } +} diff --git a/codex-rs/exec-server/src/server.rs b/codex-rs/exec-server/src/server.rs new file mode 100644 index 0000000000000000000000000000000000000000..c52cf3222e200381531c10daba8f873baf5faf74 --- /dev/null +++ b/codex-rs/exec-server/src/server.rs @@ -0,0 +1,114 @@ +mod build_identity; +mod file_system_handler; +mod handler; +mod process_handler; +mod processor; +mod registry; +mod release_version; +mod request_dispatcher; +mod session_registry; +mod transport; + +#[cfg(all(test, unix))] +#[path = "server/process_otel_tests.rs"] +mod process_otel_tests; + +pub(crate) use handler::ExecServerHandler; +pub(crate) use processor::ConnectionProcessor; +pub use request_dispatcher::ConcurrentRequestLimit; +pub use request_dispatcher::RequestDispatchMode; +pub use transport::DEFAULT_LISTEN_URL; +pub use transport::ExecServerListenUrlParseError; + +use crate::ExecServerRuntimePaths; +use crate::ExecServerTelemetry; +use codex_http_client::HttpClientFactory; + +pub async fn run_main( + listen_url: &str, + runtime_paths: ExecServerRuntimePaths, + http_client_factory: HttpClientFactory, +) -> Result<(), Box> { + run_main_with_telemetry( + listen_url, + runtime_paths, + ExecServerTelemetry::default(), + http_client_factory, + RequestDispatchMode::Inline, + ) + .await +} + +#[tracing::instrument( + name = "codex.exec_server", + skip_all, + fields(otel.kind = "internal") +)] +pub async fn run_main_with_telemetry( + listen_url: &str, + runtime_paths: ExecServerRuntimePaths, + telemetry: ExecServerTelemetry, + http_client_factory: HttpClientFactory, + request_dispatch_mode: RequestDispatchMode, +) -> Result<(), Box> { + std::sync::LazyLock::force(&build_identity::PROVIDER_ID); + transport::run_transport( + listen_url, + runtime_paths, + telemetry, + http_client_factory, + request_dispatch_mode, + ) + .await +} + +#[cfg(test)] +mod tests { + use codex_http_client::HttpClientFactory; + use codex_http_client::OutboundProxyPolicy; + use opentelemetry::trace::TracerProvider as _; + use opentelemetry_sdk::trace::InMemorySpanExporter; + use opentelemetry_sdk::trace::SdkTracerProvider; + use tracing::instrument::WithSubscriber; + use tracing_subscriber::prelude::*; + + use super::run_main_with_telemetry; + use crate::ExecServerRuntimePaths; + use crate::ExecServerTelemetry; + + #[tokio::test] + async fn telemetry_entrypoint_emits_root_span() { + let exporter = InMemorySpanExporter::default(); + let provider = SdkTracerProvider::builder() + .with_simple_exporter(exporter.clone()) + .build(); + let subscriber = tracing_subscriber::registry() + .with(tracing_opentelemetry::layer().with_tracer(provider.tracer("exec-server-test"))); + + async { + tracing::callsite::rebuild_interest_cache(); + run_main_with_telemetry( + "invalid", + ExecServerRuntimePaths::new( + std::env::current_exe().expect("current executable"), + /*codex_linux_sandbox_exe*/ None, + ) + .expect("runtime paths"), + ExecServerTelemetry::default(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + super::RequestDispatchMode::Inline, + ) + .await + .expect_err("invalid listen URL should fail"); + } + .with_subscriber(subscriber) + .await; + + provider.force_flush().expect("flush traces"); + let spans = exporter.get_finished_spans().expect("span export"); + assert!( + spans.iter().any(|span| span.name == "codex.exec_server"), + "root exec-server span missing: {spans:?}" + ); + } +} diff --git a/codex-rs/exec-server/src/server/build_identity.rs b/codex-rs/exec-server/src/server/build_identity.rs new file mode 100644 index 0000000000000000000000000000000000000000..691b592535e6362da9219398b8b6b8ea0ba16303 --- /dev/null +++ b/codex-rs/exec-server/src/server/build_identity.rs @@ -0,0 +1,21 @@ +//! Cache the running executor's build identity before accepting connections. + +use std::sync::LazyLock; + +use codex_build_info::BuildInfo; +use codex_build_info::build_id; + +use crate::protocol::EnvironmentInfo; + +pub(super) static PROVIDER_ID: LazyLock> = LazyLock::new(|| { + let info = BuildInfo::get(); + info.target() + .and_then(|target| build_id(info.build_commit(), target)) +}); + +pub(super) fn local_environment_info() -> EnvironmentInfo { + EnvironmentInfo { + provider_id: PROVIDER_ID.clone(), + ..super::release_version::local_environment_info() + } +} diff --git a/codex-rs/exec-server/src/server/file_system_handler.rs b/codex-rs/exec-server/src/server/file_system_handler.rs new file mode 100644 index 0000000000000000000000000000000000000000..8acc2a39c644ff441eba59ab4f9ac5121210455e --- /dev/null +++ b/codex-rs/exec-server/src/server/file_system_handler.rs @@ -0,0 +1,433 @@ +use std::io; + +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD; +use codex_exec_server_protocol::JSONRPCErrorError; + +use crate::CapabilityRootsDiscoverParams; +use crate::CapabilityRootsDiscoverResponse; +use crate::CopyOptions; +use crate::CreateDirectoryOptions; +use crate::ExecServerRuntimePaths; +use crate::ExecutorFileSystem; +use crate::GetMetadataOptions; +use crate::ReadFileOptions; +use crate::RemoveOptions; +use crate::WriteFileOptions; +use crate::file_read::FileReadHandleManager; +use crate::local_file_system::LocalFileSystem; +use crate::protocol::FS_READ_DIRECTORY_METHOD; +use crate::protocol::FS_WRITE_FILE_METHOD; +use crate::protocol::FsCanonicalizeParams; +use crate::protocol::FsCanonicalizeResponse; +use crate::protocol::FsCloseParams; +use crate::protocol::FsCloseResponse; +use crate::protocol::FsCopyParams; +use crate::protocol::FsCopyResponse; +use crate::protocol::FsCreateDirectoryParams; +use crate::protocol::FsCreateDirectoryResponse; +use crate::protocol::FsGetMetadataParams; +use crate::protocol::FsGetMetadataResponse; +use crate::protocol::FsOpenParams; +use crate::protocol::FsOpenResponse; +use crate::protocol::FsReadBlockParams; +use crate::protocol::FsReadBlockResponse; +use crate::protocol::FsReadDirectoryEntry; +use crate::protocol::FsReadDirectoryParams; +use crate::protocol::FsReadDirectoryResponse; +use crate::protocol::FsReadFileParams; +use crate::protocol::FsReadFileResponse; +use crate::protocol::FsRemoveParams; +use crate::protocol::FsRemoveResponse; +use crate::protocol::FsWalkParams; +use crate::protocol::FsWalkResponse; +use crate::protocol::FsWriteFileParams; +use crate::protocol::FsWriteFileResponse; +use crate::rpc::internal_error; +use crate::rpc::invalid_request; +use crate::rpc::not_found; + +const MAX_FILE_READ_HANDLE_ID_BYTES: usize = 32; +// Each read-directory entry needs four JSON values. Keep same-version +// producers comfortably below the shared 256K-value decoder budget. +const MAX_READ_DIRECTORY_ENTRIES: usize = 50_000; + +#[derive(Clone)] +pub(crate) struct FileSystemHandler { + file_system: LocalFileSystem, + file_reads: FileReadHandleManager, +} + +impl FileSystemHandler { + pub(crate) fn new(runtime_paths: ExecServerRuntimePaths) -> Self { + Self { + file_system: LocalFileSystem::with_runtime_paths(runtime_paths), + file_reads: FileReadHandleManager::default(), + } + } + + pub(crate) async fn shutdown(&self) { + self.file_reads.close_all().await; + } + + pub(crate) async fn discover_capability_roots( + &self, + params: CapabilityRootsDiscoverParams, + ) -> Result { + let sandbox = params + .roots + .first() + .and_then(|root| root.sandbox.as_ref()) + .filter(|sandbox| { + sandbox.should_run_in_sandbox() + && (!cfg!(target_os = "windows") || sandbox.windows_sandbox_is_requested()) + && params + .roots + .iter() + .all(|root| root.sandbox.as_ref() == Some(*sandbox)) + }) + .cloned(); + + if let Some(sandbox) = sandbox { + let mut batched_params = params.clone(); + for root in &mut batched_params.roots { + root.sandbox = None; + } + let result = match self.file_system.sandboxed() { + Ok(file_system) => { + file_system + .discover_capability_roots(batched_params, &sandbox) + .await + } + Err(error) => Err(error), + }; + match result { + Ok(response) => return Ok(response), + Err(error) => { + tracing::warn!(%error, "batched capability discovery failed; retrying roots separately"); + } + } + } + + crate::discover_capability_roots(&self.file_system, params) + .await + .map_err(|error| invalid_request(error.to_string())) + } + + pub(crate) async fn open( + &self, + params: FsOpenParams, + ) -> Result { + validate_file_read_handle_id(¶ms.handle_id)?; + let file = self + .file_system + .open_file_for_read(¶ms.path, params.sandbox.as_ref()) + .await + .map_err(map_fs_error)?; + let handle_id = self + .file_reads + .open(params.handle_id, file) + .await + .map_err(map_fs_error)?; + Ok(FsOpenResponse { handle_id }) + } + + pub(crate) async fn read_block( + &self, + params: FsReadBlockParams, + ) -> Result { + validate_file_read_handle_id(¶ms.handle_id)?; + let block = self + .file_reads + .read_block(¶ms.handle_id, params.offset, params.len) + .await + .map_err(map_fs_error)?; + Ok(FsReadBlockResponse { + chunk: block.bytes.into(), + eof: block.eof, + }) + } + + pub(crate) async fn close( + &self, + params: FsCloseParams, + ) -> Result { + validate_file_read_handle_id(¶ms.handle_id)?; + self.file_reads.close(¶ms.handle_id).await; + Ok(FsCloseResponse {}) + } + + pub(crate) async fn read_file( + &self, + params: FsReadFileParams, + ) -> Result { + let bytes = self + .file_system + .read_file( + ¶ms.path, + ReadFileOptions { + follow_symlinks: params.follow_symlinks.unwrap_or(true), + }, + params.sandbox.as_ref(), + ) + .await + .map_err(map_fs_error)?; + Ok(FsReadFileResponse { + data_base64: STANDARD.encode(bytes), + }) + } + + pub(crate) async fn write_file( + &self, + params: FsWriteFileParams, + ) -> Result { + let bytes = STANDARD.decode(params.data_base64).map_err(|err| { + invalid_request(format!( + "{FS_WRITE_FILE_METHOD} requires valid base64 dataBase64: {err}" + )) + })?; + self.file_system + .write_file( + ¶ms.path, + bytes, + WriteFileOptions { + follow_symlinks: params.follow_symlinks.unwrap_or(true), + }, + params.sandbox.as_ref(), + ) + .await + .map_err(map_fs_error)?; + Ok(FsWriteFileResponse {}) + } + + pub(crate) async fn create_directory( + &self, + params: FsCreateDirectoryParams, + ) -> Result { + let recursive = params.recursive.unwrap_or(true); + self.file_system + .create_directory( + ¶ms.path, + CreateDirectoryOptions { + recursive, + follow_symlinks: params.follow_symlinks.unwrap_or(true), + }, + params.sandbox.as_ref(), + ) + .await + .map_err(map_fs_error)?; + Ok(FsCreateDirectoryResponse {}) + } + + pub(crate) async fn get_metadata( + &self, + params: FsGetMetadataParams, + ) -> Result { + let metadata = self + .file_system + .get_metadata( + ¶ms.path, + GetMetadataOptions { + follow_symlinks: params.follow_symlinks.unwrap_or(true), + }, + params.sandbox.as_ref(), + ) + .await + .map_err(map_fs_error)?; + Ok(FsGetMetadataResponse { + is_directory: metadata.is_directory, + is_file: metadata.is_file, + is_symlink: metadata.is_symlink, + size: metadata.size, + created_at_ms: metadata.created_at_ms, + modified_at_ms: metadata.modified_at_ms, + }) + } + + pub(crate) async fn canonicalize( + &self, + params: FsCanonicalizeParams, + ) -> Result { + let path = self + .file_system + .canonicalize(¶ms.path, params.sandbox.as_ref()) + .await + .map_err(map_fs_error)?; + Ok(FsCanonicalizeResponse { path }) + } + + pub(crate) async fn read_directory( + &self, + params: FsReadDirectoryParams, + ) -> Result { + let entries = self + .file_system + .read_directory(¶ms.path, params.sandbox.as_ref()) + .await + .map_err(map_fs_error)?; + let entry_count = entries.len(); + if entry_count > MAX_READ_DIRECTORY_ENTRIES { + return Err(internal_error(format!( + "{FS_READ_DIRECTORY_METHOD} returned {entry_count} entries; limit is {MAX_READ_DIRECTORY_ENTRIES}" + ))); + } + let entries = entries + .into_iter() + .map(|entry| FsReadDirectoryEntry { + file_name: entry.file_name, + is_directory: entry.is_directory, + is_file: entry.is_file, + }) + .collect(); + Ok(FsReadDirectoryResponse { entries }) + } + + pub(crate) async fn walk( + &self, + params: FsWalkParams, + ) -> Result { + self.file_system + .walk(¶ms.path, params.options, params.sandbox.as_ref()) + .await + .map_err(map_fs_error) + } + + pub(crate) async fn remove( + &self, + params: FsRemoveParams, + ) -> Result { + let recursive = params.recursive.unwrap_or(true); + let force = params.force.unwrap_or(true); + self.file_system + .remove( + ¶ms.path, + RemoveOptions { + recursive, + force, + follow_symlinks: params.follow_symlinks.unwrap_or(true), + }, + params.sandbox.as_ref(), + ) + .await + .map_err(map_fs_error)?; + Ok(FsRemoveResponse {}) + } + + pub(crate) async fn copy( + &self, + params: FsCopyParams, + ) -> Result { + self.file_system + .copy( + ¶ms.source_path, + ¶ms.destination_path, + CopyOptions { + recursive: params.recursive, + }, + params.sandbox.as_ref(), + ) + .await + .map_err(map_fs_error)?; + Ok(FsCopyResponse {}) + } +} + +fn validate_file_read_handle_id(handle_id: &str) -> Result<(), JSONRPCErrorError> { + if handle_id.len() > MAX_FILE_READ_HANDLE_ID_BYTES { + return Err(invalid_request(format!( + "file read handle ID must not exceed {MAX_FILE_READ_HANDLE_ID_BYTES} bytes" + ))); + } + Ok(()) +} + +fn map_fs_error(err: io::Error) -> JSONRPCErrorError { + match err.kind() { + io::ErrorKind::NotFound => not_found(err.to_string()), + io::ErrorKind::InvalidInput | io::ErrorKind::PermissionDenied => { + invalid_request(err.to_string()) + } + _ => internal_error(err.to_string()), + } +} + +#[cfg(test)] +mod tests { + use codex_protocol::protocol::NetworkAccess; + use codex_protocol::protocol::SandboxPolicy; + use codex_utils_path_uri::PathUri; + use pretty_assertions::assert_eq; + + use super::*; + use crate::FileSystemSandboxContext; + use crate::protocol::FsReadFileParams; + use crate::protocol::FsWriteFileParams; + + #[tokio::test] + async fn no_platform_sandbox_policies_do_not_require_configured_sandbox_helper() { + let temp_dir = tempfile::tempdir().expect("tempdir"); + let runtime_paths = ExecServerRuntimePaths::new( + std::env::current_exe().expect("current exe"), + /*codex_linux_sandbox_exe*/ None, + ) + .expect("runtime paths"); + let handler = FileSystemHandler::new(runtime_paths); + let sandbox_cwd = PathUri::from_host_native_path(temp_dir.path()).expect("tempdir URI"); + let sandbox_context = |sandbox_policy| { + FileSystemSandboxContext::from_legacy_sandbox_policy( + sandbox_policy, + sandbox_cwd.clone(), + ) + .expect("sandbox context") + }; + + for (file_name, sandbox_policy) in [ + ("danger.txt", SandboxPolicy::DangerFullAccess), + ( + "external.txt", + SandboxPolicy::ExternalSandbox { + network_access: NetworkAccess::Restricted, + }, + ), + ] { + let path = + PathUri::from_host_native_path(temp_dir.path().join(file_name)).expect("path URI"); + + handler + .write_file(FsWriteFileParams { + path: path.clone(), + follow_symlinks: None, + data_base64: STANDARD.encode("ok"), + sandbox: Some(sandbox_context(sandbox_policy.clone())), + }) + .await + .expect("write file"); + + let canonicalized = handler + .canonicalize(FsCanonicalizeParams { + path: path.clone(), + sandbox: Some(sandbox_context(sandbox_policy.clone())), + }) + .await + .expect("canonicalize file"); + assert_eq!( + canonicalized.path, + PathUri::from_host_native_path( + std::fs::canonicalize(temp_dir.path().join(file_name)).expect("canonical path"), + ) + .expect("canonical path URI"), + ); + + let response = handler + .read_file(FsReadFileParams { + path, + follow_symlinks: None, + sandbox: Some(sandbox_context(sandbox_policy)), + }) + .await + .expect("read file"); + + assert_eq!(response.data_base64, STANDARD.encode("ok")); + } + } +} diff --git a/codex-rs/exec-server/src/server/handler.rs b/codex-rs/exec-server/src/server/handler.rs new file mode 100644 index 0000000000000000000000000000000000000000..3041a102a7e926bb39ac611e07e4a131c9b4a616 --- /dev/null +++ b/codex-rs/exec-server/src/server/handler.rs @@ -0,0 +1,485 @@ +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_exec_server_protocol::RequestId; +use codex_http_client::HttpClientFactory; +use opentelemetry::trace::SpanContext; +use serde_json::to_value; +use std::collections::HashSet; +use tokio::sync::Mutex; +use tokio_util::sync::CancellationToken; +use tokio_util::task::TaskTracker; + +use crate::ExecServerRuntimePaths; +use crate::client::http_client::PendingRouteAwareHttpBodyStream; +use crate::client::http_client::RouteAwareHttpClient; +use crate::client::http_client::RouteAwareHttpRequestRunner; +use crate::environment_config::ReadEnvironmentConfigError; +use crate::environment_config::read_environment_config; +use crate::protocol::CapabilityRootsDiscoverParams; +use crate::protocol::CapabilityRootsDiscoverResponse; +use crate::protocol::EnvironmentConfigReadParams; +use crate::protocol::EnvironmentConfigReadResponse; +use crate::protocol::EnvironmentInfo; +use crate::protocol::EnvironmentStatus; +use crate::protocol::EnvironmentStatusKind; +use crate::protocol::ExecParams; +use crate::protocol::ExecResponse; +use crate::protocol::FsCanonicalizeParams; +use crate::protocol::FsCanonicalizeResponse; +use crate::protocol::FsCloseParams; +use crate::protocol::FsCloseResponse; +use crate::protocol::FsCopyParams; +use crate::protocol::FsCopyResponse; +use crate::protocol::FsCreateDirectoryParams; +use crate::protocol::FsCreateDirectoryResponse; +use crate::protocol::FsGetMetadataParams; +use crate::protocol::FsGetMetadataResponse; +use crate::protocol::FsOpenParams; +use crate::protocol::FsOpenResponse; +use crate::protocol::FsReadBlockParams; +use crate::protocol::FsReadBlockResponse; +use crate::protocol::FsReadDirectoryParams; +use crate::protocol::FsReadDirectoryResponse; +use crate::protocol::FsReadFileParams; +use crate::protocol::FsReadFileResponse; +use crate::protocol::FsRemoveParams; +use crate::protocol::FsRemoveResponse; +use crate::protocol::FsWalkParams; +use crate::protocol::FsWalkResponse; +use crate::protocol::FsWriteFileParams; +use crate::protocol::FsWriteFileResponse; +use crate::protocol::HttpRequestParams; +use crate::protocol::InitializeParams; +use crate::protocol::InitializeResponse; +use crate::protocol::ReadParams; +use crate::protocol::ReadResponse; +use crate::protocol::SignalParams; +use crate::protocol::SignalResponse; +use crate::protocol::TerminateParams; +use crate::protocol::TerminateResponse; +use crate::protocol::WriteParams; +use crate::protocol::WriteResponse; +use crate::rpc::RpcNotificationSender; +use crate::rpc::internal_error; +use crate::rpc::invalid_params; +use crate::rpc::invalid_request; +use crate::server::build_identity::local_environment_info; +use crate::server::file_system_handler::FileSystemHandler; +use crate::server::session_registry::SessionHandle; +use crate::server::session_registry::SessionRegistry; +use crate::telemetry::ExecutorRegistration; + +pub(crate) struct ExecServerHandler { + pub(super) executor_registration: Option>, + session_registry: Arc, + notifications: RpcNotificationSender, + session: StdMutex>, + active_body_stream_ids: Mutex>, + background_task_shutdown: CancellationToken, + background_tasks: TaskTracker, + file_system: FileSystemHandler, + runtime_paths: ExecServerRuntimePaths, + http_client: RouteAwareHttpClient, + initialize_requested: AtomicBool, + initialized: AtomicBool, +} + +impl ExecServerHandler { + pub(crate) fn new( + session_registry: Arc, + notifications: RpcNotificationSender, + runtime_paths: ExecServerRuntimePaths, + http_client_factory: HttpClientFactory, + ) -> Self { + Self { + executor_registration: None, + session_registry, + notifications, + session: StdMutex::new(None), + active_body_stream_ids: Mutex::new(HashSet::new()), + background_task_shutdown: CancellationToken::new(), + background_tasks: TaskTracker::new(), + file_system: FileSystemHandler::new(runtime_paths.clone()), + runtime_paths, + http_client: RouteAwareHttpClient::new(http_client_factory), + initialize_requested: AtomicBool::new(false), + initialized: AtomicBool::new(false), + } + } + + pub(crate) async fn shutdown(&self) { + self.background_task_shutdown.cancel(); + self.background_tasks.close(); + self.background_tasks.wait().await; + self.file_system.shutdown().await; + if let Some(session) = self.session() { + session.detach().await; + } + } + + pub(crate) fn is_session_attached(&self) -> bool { + self.session() + .is_none_or(|session| session.is_session_attached()) + } + + pub(crate) async fn initialize( + &self, + params: InitializeParams, + ) -> Result { + if self.initialize_requested.swap(true, Ordering::SeqCst) { + return Err(invalid_request( + "initialize may only be sent once per connection".to_string(), + )); + } + + let session = match self + .session_registry + .attach( + params.resume_session_id.clone(), + self.notifications.clone(), + self.runtime_paths.clone(), + ) + .await + { + Ok(session) => session, + Err(error) => { + self.initialize_requested.store(false, Ordering::SeqCst); + return Err(error); + } + }; + let session_id = session.session_id().to_string(); + tracing::debug!( + session_id, + connection_id = %session.connection_id(), + "exec-server session attached" + ); + *self + .session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(session); + Ok(InitializeResponse { + session_id, + environment_info: Some(local_environment_info()), + }) + } + + pub(crate) fn initialized(&self) -> Result<(), String> { + if !self.initialize_requested.load(Ordering::SeqCst) { + return Err("received `initialized` notification before `initialize`".into()); + } + self.require_session_attached() + .map_err(|error| error.message)?; + self.initialized.store(true, Ordering::SeqCst); + Ok(()) + } + + pub(crate) async fn exec( + &self, + params: ExecParams, + launch_context: Option, + ) -> Result { + let session = self.require_initialized_for("exec")?; + session + .process() + .exec( + params, + crate::process_telemetry::ProcessTelemetry { + launch_context, + executor_registration: self.executor_registration.clone(), + ..Default::default() + }, + ) + .await + } + + pub(crate) fn environment_info(&self) -> Result { + self.require_initialized_for("environment info")?; + Ok(local_environment_info()) + } + + pub(crate) async fn environment_config_read( + &self, + params: EnvironmentConfigReadParams, + ) -> Result { + self.require_initialized_for("environment config")?; + read_environment_config(crate::LOCAL_FS.as_ref(), params) + .await + .map_err(|error| match error { + ReadEnvironmentConfigError::InvalidParams(message) => invalid_params(message), + ReadEnvironmentConfigError::Internal(message) => internal_error(message), + }) + } + + pub(crate) fn environment_status(&self) -> Result { + self.require_initialized_for("environment status")?; + Ok(EnvironmentStatus { + status: EnvironmentStatusKind::Ready, + }) + } + + pub(crate) async fn exec_read( + &self, + params: ReadParams, + ) -> Result { + let session = self.require_initialized_for("exec")?; + let response = session.process().exec_read(params).await?; + self.require_session_attached()?; + Ok(response) + } + + pub(crate) async fn exec_write( + &self, + params: WriteParams, + ) -> Result { + let session = self.require_initialized_for("exec")?; + session.process().exec_write(params).await + } + + pub(crate) async fn signal( + &self, + params: SignalParams, + ) -> Result { + let session = self.require_initialized_for("exec")?; + session.process().signal(params).await + } + + pub(crate) async fn terminate( + &self, + params: TerminateParams, + ) -> Result { + let session = self.require_initialized_for("exec")?; + session.process().terminate(params).await + } + + pub(crate) async fn http_request( + self: &Arc, + request_id: RequestId, + params: HttpRequestParams, + ) -> Result<(), JSONRPCErrorError> { + self.require_initialized_for("http")?; + let stream_response = params.stream_response; + let http_request_id = params.request_id.clone(); + if stream_response { + self.reserve_http_body_stream(&http_request_id).await?; + } + let response = self + .http_client + .runner(params.redirect_policy) + .run(params) + .await; + if response.is_err() && stream_response { + self.release_http_body_stream(&http_request_id).await; + } + let (response, mut pending_stream) = response?; + let result = match to_value(response) { + Ok(result) => result, + Err(err) => { + if let Some(pending_stream) = pending_stream.take() { + self.release_http_body_stream(&pending_stream.request_id) + .await; + } + return Err(internal_error(err.to_string())); + } + }; + if let Err(error) = self.notifications.response(request_id, result).await { + if let Some(pending_stream) = pending_stream.take() { + self.release_http_body_stream(&pending_stream.request_id) + .await; + } + return Err(error); + } + if let Some(pending_stream) = pending_stream { + self.start_http_body_stream(pending_stream).await; + } + Ok(()) + } + + pub(crate) async fn fs_read_file( + &self, + params: FsReadFileParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.read_file(params).await + } + + pub(crate) async fn discover_capability_roots( + &self, + params: CapabilityRootsDiscoverParams, + ) -> Result { + self.require_initialized_for("capability discovery")?; + self.file_system.discover_capability_roots(params).await + } + + pub(crate) async fn fs_open( + &self, + params: FsOpenParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.open(params).await + } + + pub(crate) async fn fs_read_block( + &self, + params: FsReadBlockParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.read_block(params).await + } + + pub(crate) async fn fs_close( + &self, + params: FsCloseParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.close(params).await + } + + pub(crate) async fn fs_write_file( + &self, + params: FsWriteFileParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.write_file(params).await + } + + pub(crate) async fn fs_create_directory( + &self, + params: FsCreateDirectoryParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.create_directory(params).await + } + + pub(crate) async fn fs_get_metadata( + &self, + params: FsGetMetadataParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.get_metadata(params).await + } + + pub(crate) async fn fs_canonicalize( + &self, + params: FsCanonicalizeParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.canonicalize(params).await + } + + pub(crate) async fn fs_read_directory( + &self, + params: FsReadDirectoryParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.read_directory(params).await + } + + pub(crate) async fn fs_walk( + &self, + params: FsWalkParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.walk(params).await + } + + pub(crate) async fn fs_remove( + &self, + params: FsRemoveParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.remove(params).await + } + + pub(crate) async fn fs_copy( + &self, + params: FsCopyParams, + ) -> Result { + self.require_initialized_for("filesystem")?; + self.file_system.copy(params).await + } + + fn require_initialized_for( + &self, + method_family: &str, + ) -> Result { + if !self.initialize_requested.load(Ordering::SeqCst) { + return Err(invalid_request(format!( + "client must call initialize before using {method_family} methods" + ))); + } + let session = self.require_session_attached()?; + if !self.initialized.load(Ordering::SeqCst) { + return Err(invalid_request(format!( + "client must send initialized before using {method_family} methods" + ))); + } + Ok(session) + } + + fn require_session_attached(&self) -> Result { + let Some(session) = self.session() else { + return Err(invalid_request( + "client must call initialize before using methods".to_string(), + )); + }; + if session.is_session_attached() { + return Ok(session); + } + + Err(invalid_request( + "session has been resumed by another connection".to_string(), + )) + } + + fn session(&self) -> Option { + self.session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + } + + async fn start_http_body_stream( + self: &Arc, + pending_stream: PendingRouteAwareHttpBodyStream, + ) { + let request_id = pending_stream.request_id.clone(); + if self.background_task_shutdown.is_cancelled() { + self.release_http_body_stream(&request_id).await; + return; + } + let finished_request_id = request_id.clone(); + let handler = Arc::clone(self); + let notifications = self.notifications.clone(); + let shutdown = self.background_task_shutdown.clone(); + self.background_tasks.spawn(async move { + tokio::select! { + _ = shutdown.cancelled() => {} + _ = RouteAwareHttpRequestRunner::stream_body(pending_stream, notifications) => {} + } + handler.release_http_body_stream(&finished_request_id).await; + }); + } + + async fn release_http_body_stream(&self, request_id: &str) { + let mut active_body_stream_ids = self.active_body_stream_ids.lock().await; + active_body_stream_ids.remove(request_id); + } + + async fn reserve_http_body_stream(&self, request_id: &str) -> Result<(), JSONRPCErrorError> { + let mut active_body_stream_ids = self.active_body_stream_ids.lock().await; + if active_body_stream_ids.contains(request_id) { + return Err(invalid_params(format!( + "http/request streamResponse requestId `{request_id}` is already active" + ))); + } + active_body_stream_ids.insert(request_id.to_string()); + Ok(()) + } +} + +#[cfg(test)] +mod tests; diff --git a/codex-rs/exec-server/src/server/handler/tests.rs b/codex-rs/exec-server/src/server/handler/tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..5d2f43d11aa44319035c8d19dd165597403a2e86 --- /dev/null +++ b/codex-rs/exec-server/src/server/handler/tests.rs @@ -0,0 +1,380 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +use tokio::sync::mpsc; +use uuid::Uuid; + +use super::ExecServerHandler; +use crate::ExecServerRuntimePaths; +use crate::ProcessId; +use crate::protocol::ExecParams; +use crate::protocol::InitializeParams; +use crate::protocol::ReadParams; +use crate::protocol::ReadResponse; +use crate::protocol::TerminateParams; +use crate::protocol::TerminateResponse; +use crate::rpc::RpcNotificationSender; +use crate::server::session_registry::SessionRegistry; + +fn exec_params(process_id: &str) -> ExecParams { + exec_params_with_argv(process_id, sleep_argv()) +} + +fn exec_params_with_argv(process_id: &str, argv: Vec) -> ExecParams { + ExecParams { + metadata: Default::default(), + process_id: ProcessId::from(process_id), + argv, + cwd: PathUri::from_host_native_path(std::env::current_dir().expect("cwd")) + .expect("cwd URI"), + shell_snapshot: None, + env_policy: None, + env: inherited_path_env(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + } +} + +fn inherited_path_env() -> HashMap { + let mut env = HashMap::new(); + if let Some(path) = std::env::var_os("PATH") { + env.insert("PATH".to_string(), path.to_string_lossy().into_owned()); + } + env +} + +fn sleep_argv() -> Vec { + shell_argv("sleep 0.1", "ping -n 2 127.0.0.1 >NUL") +} + +fn shell_argv(unix_script: &str, windows_script: &str) -> Vec { + if cfg!(windows) { + vec![ + windows_command_processor(), + "/C".to_string(), + windows_script.to_string(), + ] + } else { + vec![ + "/bin/sh".to_string(), + "-c".to_string(), + unix_script.to_string(), + ] + } +} + +fn windows_command_processor() -> String { + std::env::var("COMSPEC").unwrap_or_else(|_| "cmd.exe".to_string()) +} + +fn test_runtime_paths() -> ExecServerRuntimePaths { + ExecServerRuntimePaths::new( + std::env::current_exe().expect("current exe"), + /*codex_linux_sandbox_exe*/ None, + ) + .expect("runtime paths") +} + +fn test_http_client_factory() -> HttpClientFactory { + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault) +} + +async fn initialized_handler() -> Arc { + let (outgoing_tx, _outgoing_rx) = mpsc::channel(16); + let registry = SessionRegistry::new(crate::ExecServerTelemetry::default()); + let handler = Arc::new(ExecServerHandler::new( + registry, + RpcNotificationSender::new(outgoing_tx), + test_runtime_paths(), + test_http_client_factory(), + )); + let initialize_response = handler + .initialize(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + }) + .await + .expect("initialize"); + Uuid::parse_str(&initialize_response.session_id).expect("session id should be a UUID"); + handler.initialized().expect("initialized"); + handler +} + +#[tokio::test] +async fn duplicate_process_ids_allow_only_one_successful_start() { + let handler = initialized_handler().await; + let first_handler = Arc::clone(&handler); + let second_handler = Arc::clone(&handler); + + let (first, second) = tokio::join!( + first_handler.exec(exec_params("proc-1"), /*launch_context*/ None), + second_handler.exec(exec_params("proc-1"), /*launch_context*/ None), + ); + + let (successes, failures): (Vec<_>, Vec<_>) = + [first, second].into_iter().partition(Result::is_ok); + assert_eq!(successes.len(), 1); + assert_eq!(failures.len(), 1); + + let error = failures + .into_iter() + .next() + .expect("one failed request") + .expect_err("expected duplicate process error"); + assert_eq!(error.code, -32600); + assert_eq!(error.message, "process proc-1 already exists"); + + tokio::time::sleep(Duration::from_millis(150)).await; + handler.shutdown().await; +} + +#[tokio::test] +async fn terminate_reports_false_after_process_exit() { + let handler = initialized_handler().await; + handler + .exec(exec_params("proc-1"), /*launch_context*/ None) + .await + .expect("start process"); + + let deadline = tokio::time::Instant::now() + Duration::from_secs(1); + loop { + let response = handler + .terminate(TerminateParams { + process_id: ProcessId::from("proc-1"), + }) + .await + .expect("terminate response"); + if response == (TerminateResponse { running: false }) { + break; + } + assert!( + tokio::time::Instant::now() < deadline, + "process should have exited within 1s" + ); + tokio::time::sleep(Duration::from_millis(25)).await; + } + + handler.shutdown().await; +} + +#[tokio::test] +async fn long_poll_read_fails_after_session_resume() { + let (first_tx, _first_rx) = mpsc::channel(16); + let registry = SessionRegistry::new(crate::ExecServerTelemetry::default()); + let first_handler = Arc::new(ExecServerHandler::new( + Arc::clone(®istry), + RpcNotificationSender::new(first_tx), + test_runtime_paths(), + test_http_client_factory(), + )); + let initialize_response = first_handler + .initialize(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + }) + .await + .expect("initialize"); + first_handler.initialized().expect("initialized"); + + // Keep the process quiet and alive so the pending read can only complete + // after session resume, not because the process produced output or exited. + first_handler + .exec( + exec_params_with_argv( + "proc-long-poll", + shell_argv("sleep 5", "ping -n 6 127.0.0.1 >NUL"), + ), + /*launch_context*/ None, + ) + .await + .expect("start process"); + + let first_read_handler = Arc::clone(&first_handler); + let read_task = tokio::spawn(async move { + first_read_handler + .exec_read(ReadParams { + process_id: ProcessId::from("proc-long-poll"), + after_seq: None, + max_bytes: None, + wait_ms: Some(500), + }) + .await + }); + + tokio::time::sleep(Duration::from_millis(50)).await; + first_handler.shutdown().await; + + let (second_tx, _second_rx) = mpsc::channel(16); + let second_handler = Arc::new(ExecServerHandler::new( + registry, + RpcNotificationSender::new(second_tx), + test_runtime_paths(), + test_http_client_factory(), + )); + second_handler + .initialize(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: Some(initialize_response.session_id), + }) + .await + .expect("initialize second connection"); + second_handler + .initialized() + .expect("initialized second connection"); + + let err = read_task + .await + .expect("read task should join") + .expect_err("evicted long-poll read should fail"); + assert_eq!(err.code, -32600); + assert_eq!( + err.message, + "session has been resumed by another connection" + ); + + second_handler.shutdown().await; +} + +#[tokio::test] +async fn active_session_resume_is_rejected() { + let (first_tx, _first_rx) = mpsc::channel(16); + let registry = SessionRegistry::new(crate::ExecServerTelemetry::default()); + let first_handler = Arc::new(ExecServerHandler::new( + Arc::clone(®istry), + RpcNotificationSender::new(first_tx), + test_runtime_paths(), + test_http_client_factory(), + )); + let initialize_response = first_handler + .initialize(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + }) + .await + .expect("initialize"); + + let (second_tx, _second_rx) = mpsc::channel(16); + let second_handler = Arc::new(ExecServerHandler::new( + registry, + RpcNotificationSender::new(second_tx), + test_runtime_paths(), + test_http_client_factory(), + )); + let err = second_handler + .initialize(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: Some(initialize_response.session_id.clone()), + }) + .await + .expect_err("active session resume should fail"); + + assert_eq!(err.code, crate::rpc::SESSION_ALREADY_ATTACHED_ERROR_CODE); + assert_eq!( + err.message, + format!( + "session {} is already attached to another connection", + initialize_response.session_id + ) + ); + + first_handler.shutdown().await; +} + +#[tokio::test] +async fn output_and_exit_are_retained_after_notification_receiver_closes() { + let (outgoing_tx, outgoing_rx) = mpsc::channel(16); + let handler = Arc::new(ExecServerHandler::new( + SessionRegistry::new(crate::ExecServerTelemetry::default()), + RpcNotificationSender::new(outgoing_tx), + test_runtime_paths(), + test_http_client_factory(), + )); + handler + .initialize(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + }) + .await + .expect("initialize"); + handler.initialized().expect("initialized"); + + let process_id = ProcessId::from("proc-notification-fail"); + handler + .exec( + exec_params_with_argv( + process_id.as_str(), + shell_argv( + "sleep 0.05; printf 'first\\n'; sleep 0.05; printf 'second\\n'", + "echo first&& ping -n 2 127.0.0.1 >NUL&& echo second", + ), + ), + /*launch_context*/ None, + ) + .await + .expect("start process"); + + drop(outgoing_rx); + + let (output, exit_code) = read_process_until_closed(&handler, process_id.clone()).await; + assert_eq!(output.replace("\r\n", "\n"), "first\nsecond\n"); + assert_eq!(exit_code, Some(0)); + + tokio::time::sleep(Duration::from_millis(100)).await; + handler + .exec( + exec_params(process_id.as_str()), + /*launch_context*/ None, + ) + .await + .expect("process id should be reusable after exit retention"); + + handler.shutdown().await; +} + +async fn read_process_until_closed( + handler: &ExecServerHandler, + process_id: ProcessId, +) -> (String, Option) { + let deadline = tokio::time::Instant::now() + Duration::from_secs(5); + let mut output = String::new(); + let mut exit_code = None; + let mut after_seq = None; + + loop { + let response: ReadResponse = handler + .exec_read(ReadParams { + process_id: process_id.clone(), + after_seq, + max_bytes: None, + wait_ms: Some(500), + }) + .await + .expect("read process"); + + for chunk in response.chunks { + output.push_str(&String::from_utf8_lossy(&chunk.chunk.into_inner())); + after_seq = Some(chunk.seq); + } + if response.exited { + exit_code = response.exit_code; + } + if response.closed { + return (output, exit_code); + } + after_seq = response.next_seq.checked_sub(1).or(after_seq); + assert!( + tokio::time::Instant::now() < deadline, + "process should close within 5s" + ); + } +} diff --git a/codex-rs/exec-server/src/server/process_handler.rs b/codex-rs/exec-server/src/server/process_handler.rs new file mode 100644 index 0000000000000000000000000000000000000000..a4aca9d72c2b94da9b59054d95b93ffd578aea49 --- /dev/null +++ b/codex-rs/exec-server/src/server/process_handler.rs @@ -0,0 +1,78 @@ +use crate::process_telemetry::ProcessTelemetry; +use codex_exec_server_protocol::JSONRPCErrorError; + +use crate::ExecServerRuntimePaths; +use crate::local_process::LocalProcess; +use crate::protocol::ExecParams; +use crate::protocol::ExecResponse; +use crate::protocol::ReadParams; +use crate::protocol::ReadResponse; +use crate::protocol::SignalParams; +use crate::protocol::SignalResponse; +use crate::protocol::TerminateParams; +use crate::protocol::TerminateResponse; +use crate::protocol::WriteParams; +use crate::protocol::WriteResponse; +use crate::rpc::RpcNotificationSender; +use crate::telemetry::ExecServerTelemetry; + +#[derive(Clone)] +pub(crate) struct ProcessHandler { + process: LocalProcess, +} + +impl ProcessHandler { + pub(crate) fn new( + notifications: RpcNotificationSender, + telemetry: ExecServerTelemetry, + runtime_paths: ExecServerRuntimePaths, + ) -> Self { + Self { + process: LocalProcess::new(notifications, telemetry, runtime_paths), + } + } + + pub(crate) async fn shutdown(&self) { + self.process.shutdown().await; + } + + pub(crate) fn set_notification_sender(&self, notifications: Option) { + self.process.set_notification_sender(notifications); + } + + pub(crate) async fn exec( + &self, + params: ExecParams, + telemetry: ProcessTelemetry, + ) -> Result { + self.process.exec(params, telemetry).await + } + + pub(crate) async fn exec_read( + &self, + params: ReadParams, + ) -> Result { + self.process.exec_read(params).await + } + + pub(crate) async fn exec_write( + &self, + params: WriteParams, + ) -> Result { + self.process.exec_write(params).await + } + + pub(crate) async fn signal( + &self, + params: SignalParams, + ) -> Result { + self.process.signal_process(params).await + } + + pub(crate) async fn terminate( + &self, + params: TerminateParams, + ) -> Result { + self.process.terminate_process(params).await + } +} diff --git a/codex-rs/exec-server/src/server/process_otel_tests.rs b/codex-rs/exec-server/src/server/process_otel_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..3460408ea563d100720ce2688911307e96d9decb --- /dev/null +++ b/codex-rs/exec-server/src/server/process_otel_tests.rs @@ -0,0 +1,752 @@ +//! Verifies exported lifecycle and network logs retain launch identity without private payloads. + +#![allow(clippy::expect_used)] + +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::RequestId; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_network_proxy::NetworkProxyConfig; +use codex_network_proxy::RemoteNetworkProxyConfig; +use codex_network_proxy::RemoteNetworkProxyLaunchConfig; +use codex_otel::OtelExporter; +use codex_otel::OtelHttpProtocol; +use codex_otel::OtelProvider; +use codex_otel::OtelSettings; +use codex_protocol::protocol::W3cTraceContext; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +use serde_json::Value; +use serde_json::json; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::sync::mpsc; +use tracing::Instrument; +use tracing_subscriber::prelude::*; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; + +use super::ExecServerHandler; +use super::registry::build_router; +use super::session_registry::SessionRegistry; +use crate::ExecServerRuntimePaths; +use crate::ExecServerTelemetry; +use crate::connection::JsonRpcConnectionEvent; +use crate::protocol::EXEC_METHOD; +use crate::protocol::EXEC_READ_METHOD; +use crate::protocol::EXEC_TERMINATE_METHOD; +use crate::protocol::InitializeParams; +use crate::protocol::ReadResponse; +use crate::rpc::RpcNotificationSender; +use crate::rpc::RpcRouter; +use crate::rpc::RpcServerOutboundMessage; +use crate::telemetry::ExecutorRegistration; + +const PRIVATE_PAYLOAD: &str = "private-process-payload"; +const TRACE_ID: &str = "11111111111111111111111111111111"; +const LAUNCH_SPANS: [&str; 2] = ["2222222222222222", "3333333333333333"]; +const THREADS: [&str; 2] = [ + "11111111-1111-4111-8111-111111111111", + "22222222-2222-4222-8222-222222222222", +]; +const CALLS: [&str; 2] = ["call-first", "call-second"]; +const LATER_TRACE: &str = "00-44444444444444444444444444444444-5555555555555555-01"; + +#[test] +fn exported_process_and_network_logs_keep_the_validated_launch_reference() { + // Keep the subscriber on the runtime's only thread, including proxy background tasks. + for (export_traces, registered) in [(false, true), (true, true), (false, false)] { + let records = exported_logs(export_traces, |runtime, _, _| { + runtime.block_on(async { + let server = TestServer::new(); + let mut handlers = Vec::new(); + let mut session_ids = Vec::new(); + for client_name in ["first-orchestrator", "second-orchestrator"] { + let (handler, session_id) = server.new_orchestrator( + client_name, registered.then_some("original"), + ).await; + handlers.push(handler); + session_ids.push(session_id); + } + let proxy = RemoteNetworkProxyLaunchConfig::new( + RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig { + enabled: true, + ..NetworkProxyConfig::default() + }).expect("proxy config"), + ); + // Unsampled incoming context must still correlate logs-only export. + let flags = if export_traces { "01" } else { "00" }; + let launches = LAUNCH_SPANS.map(|span| format!("00-{TRACE_ID}-{span}-{flags}")); + for (batch, traces) in [ + [Some(launches[0].as_str()), Some(launches[1].as_str())], + [None, Some(PRIVATE_PAYLOAD)], + ].into_iter().enumerate() { + // Process IDs are session-scoped and deliberately reused across clients. + let process_id = format!("{PRIVATE_PAYLOAD}-{batch}"); + for (index, (handler, trace)) in handlers.iter().zip(traces).enumerate() { + let mut params = process_start_params( + &process_id, + json!(["/bin/sh", "-c", "printf '%s\\n' \"$HTTP_PROXY\"; read ignored", PRIVATE_PAYLOAD]), + PathUri::from_host_native_path(std::env::current_dir().expect("cwd")).expect("cwd URI"), + ); + params["metadata"] = if batch == 0 { + process_metadata(index) + } else { + json!({"toolCallId": if index == 0 { None } else { Some("private-process-payload\n") }}) + }; + params["metadata"]["executorRegistrationId"] = json!(PRIVATE_PAYLOAD); + params["metadata"]["environmentId"] = json!(PRIVATE_PAYLOAD); + params["executorRegistrationId"] = json!(PRIVATE_PAYLOAD); + params["environmentId"] = json!(PRIVATE_PAYLOAD); + params["pipeStdin"] = json!(true); + params["networkProxy"] = json!(proxy); + request(&server.router, handler, EXEC_METHOD, params, trace).await; + } + if batch == 0 { + handlers[0].shutdown().await; + let resumed = new_handler( + &server.sessions, server.outgoing.clone(), registered.then_some("resumed"), + ); + resumed.initialize(InitializeParams { + client_name: "resumed-orchestrator".to_string(), + resume_session_id: Some(session_ids[0].clone()), + }).await.expect("resume session"); + resumed.initialized().expect("initialized resumed orchestrator"); + handlers[0] = resumed; + } + for handler in &handlers { + let output: ReadResponse = serde_json::from_value(request( + &server.router, handler, EXEC_READ_METHOD, + json!({"processId": process_id, "waitMs": 1000}), Some(LATER_TRACE), + ).await).expect("read proxy address"); + let proxy_address = String::from_utf8(output.chunks.into_iter() + .flat_map(|chunk| chunk.chunk.into_inner()).collect()).expect("UTF-8 output"); + let mut connection = tokio::net::TcpStream::connect( + proxy_address.trim().strip_prefix("http://").expect("HTTP proxy address"), + ).await.expect("connect to process proxy"); + connection.write_all(b"CONNECT 8.8.8.8:443 HTTP/1.1\r\nHost: 8.8.8.8:443\r\n\r\n") + .await.expect("request denied destination"); + let mut response = [0_u8; 256]; + let length = tokio::time::timeout(Duration::from_secs(5), connection.read(&mut response)) + .await.expect("proxy response timeout").expect("proxy response"); + assert!(String::from_utf8_lossy(&response[..length]).starts_with("HTTP/1.1 403")); + request(&server.router, handler, EXEC_TERMINATE_METHOD, + json!({"processId": process_id}), Some(LATER_TRACE)).await; + read_until_closed(&server.router, handler, &process_id).await; + } + } + server.sessions.shutdown().await; + }); + }); + let mut events = Vec::new(); + for record in &records { + let name = attribute(record, "event.name").expect("event name"); + let trace = attribute(record, "launch.trace_id"); + let span = attribute(record, "launch.span_id"); + if export_traces + && name.starts_with("codex.exec_server.process_") + && let Some(span) = span + { + assert_eq!(record["traceId"].as_str(), trace); + assert_ne!(record["spanId"].as_str().expect("native child span"), span); + } + events.push(( + name.to_string(), + trace.map(str::to_string), + span.map(str::to_string), + attribute(record, "conversation.id").map(str::to_string), + attribute(record, "tool.call_id").map(str::to_string), + attribute(record, "executor.environment_id").map(str::to_string), + attribute(record, "executor.registration_id").map(str::to_string), + )); + } + let mut expected = Vec::new(); + for (index, span) in [Some(LAUNCH_SPANS[0]), Some(LAUNCH_SPANS[1]), None, None] + .into_iter() + .enumerate() + { + for name in [ + "codex.exec_server.process_start", + "codex.exec_server.process_exit", + "codex.network_proxy.policy_decision", + ] { + expected.push(( + name.to_string(), + span.map(|_| TRACE_ID.to_string()), + span.map(str::to_string), + (index < 2).then(|| THREADS[index].to_string()), + (index < 2).then(|| CALLS[index].to_string()), + registered.then(|| "environment".to_string()), + registered.then(|| if index == 2 { "resumed" } else { "original" }.to_string()), + )); + } + } + events.sort(); + expected.sort(); + assert_eq!(events, expected); + } +} + +#[test] +fn launch_rpc_span_finishes_before_its_process_exits() { + let mut observed = None; + let records = exported_logs(/*export_traces*/ true, |runtime, otel, collector| { + runtime.block_on(async { + let server = TestServer::new(); + let (handler, _) = server + .new_orchestrator("process-span-orchestrator", Some("original")) + .await; + let launch = format!("00-{TRACE_ID}-{}-01", LAUNCH_SPANS[0]); + let mut params = process_start_params( + PRIVATE_PAYLOAD, + json!(["/bin/sh", "-c", "read ignored", PRIVATE_PAYLOAD]), + PathUri::from_host_native_path(std::env::current_dir().expect("cwd")) + .expect("cwd URI"), + ); + params["metadata"] = process_metadata(/*index*/ 0); + params["pipeStdin"] = json!(true); + request(&server.router, &handler, EXEC_METHOD, params, Some(&launch)).await; + + // Flush completed spans while the process is still waiting on its open stdin. + // The old task's in_current_span() retains the RPC span, so this snapshot lacks it. + let spans_before_exit = flushed_spans(otel, collector).await; + let running: ReadResponse = serde_json::from_value( + request( + &server.router, + &handler, + EXEC_READ_METHOD, + json!({"processId": PRIVATE_PAYLOAD, "waitMs": 0}), + Some(LATER_TRACE), + ) + .await, + ) + .expect("read running process"); + request( + &server.router, + &handler, + EXEC_TERMINATE_METHOD, + json!({"processId": PRIVATE_PAYLOAD}), + Some(LATER_TRACE), + ) + .await; + let exited = read_until_closed(&server.router, &handler, PRIVATE_PAYLOAD).await; + server.sessions.shutdown().await; + let spans_after_exit = flushed_spans(otel, collector).await; + // Finish cleanup before asserting the regression so a failure leaves no child behind. + observed = Some((running, exited, spans_before_exit, spans_after_exit)); + }); + }); + let (running, exited, spans_before_exit, spans_after_exit) = + observed.expect("process observations"); + assert_eq!( + (running.exited, running.closed, running.exit_code), + (false, false, None) + ); + let launch_spans: Vec<_> = spans_before_exit + .iter() + .filter(|span| span["name"] == EXEC_METHOD) + .collect(); + assert_eq!( + launch_spans.len(), + 1, + "process/start span must export before child exit" + ); + let launch_span = launch_spans[0]; + assert_eq!( + ( + launch_span["traceId"].as_str(), + launch_span["parentSpanId"].as_str() + ), + (Some(TRACE_ID), Some(LAUNCH_SPANS[0])) + ); + let exit_record = records + .iter() + .find(|record| attribute(record, "event.name") == Some("codex.exec_server.process_exit")) + .expect("process exit log"); + let exit_span_id = exit_record["spanId"] + .as_str() + .expect("process exit span ID"); + // ProcessMetricGuard has the same span name; select the span that emitted the exit log. + let process_spans: Vec<_> = spans_after_exit + .iter() + .filter(|span| span["spanId"].as_str() == Some(exit_span_id)) + .collect(); + assert_eq!( + process_spans.len(), + 1, + "process completion has its own span" + ); + let process_span = process_spans[0]; + assert_eq!(process_span["name"], "codex.exec_server.process"); + assert_eq!( + ( + process_span["traceId"].as_str(), + process_span["parentSpanId"].as_str() + ), + (Some(TRACE_ID), Some(LAUNCH_SPANS[0])) + ); + assert_ne!(process_span["spanId"], launch_span["spanId"]); + let mut expected_exit = + expected_process_attributes("codex.exec_server.process_exit", /*index*/ 0, "None"); + expected_exit["process.exit_code"] = + json!({"intValue": exited.exit_code.expect("exit code").to_string()}); + expected_exit["process.termination_requested"] = json!({"boolValue": true}); + assert_exported_attributes( + &records, + vec![ + expected_process_attributes( + "codex.exec_server.process_start", + /*index*/ 0, + "None", + ), + expected_exit, + ], + ); + for record in &records { + assert_eq!(record["traceId"].as_str(), Some(TRACE_ID)); + let span = if attribute(record, "event.name") == Some("codex.exec_server.process_start") { + launch_span + } else { + process_span + }; + assert_eq!(record["spanId"], span["spanId"]); + } +} + +async fn flushed_spans(otel: &OtelProvider, collector: &MockServer) -> Vec { + let provider = otel + .tracer_provider + .as_ref() + .expect("trace exporter") + .clone(); + // The collector uses this runtime, so the synchronous export flush must not block it. + tokio::time::timeout( + Duration::from_secs(5), + tokio::task::spawn_blocking(move || provider.force_flush()), + ) + .await + .expect("trace flush timeout") + .expect("trace flush task") + .expect("trace flush"); + let mut spans = Vec::new(); + for request in collector + .received_requests() + .await + .expect("exported requests") + { + let body: Value = serde_json::from_slice(&request.body).expect("OTLP JSON"); + let Some(resources) = body["resourceSpans"].as_array() else { + continue; + }; + for resource in resources { + for scope in resource["scopeSpans"].as_array().expect("scope spans") { + spans.extend(scope["spans"].as_array().expect("spans").iter().cloned()); + } + } + } + spans +} + +#[test] +fn exported_spawn_failure_keeps_launch_identity_without_outcome_or_error_text() { + let records = exported_logs(/*export_traces*/ false, |runtime, _, _| { + runtime.block_on(async { + let server = TestServer::new(); + let (handler, _) = server + .new_orchestrator("spawn-failure-orchestrator", Some("original")) + .await; + let directory = tempfile::tempdir().expect("process directory"); + let missing_executable = directory.path().join(PRIVATE_PAYLOAD); + let launch = format!("00-{TRACE_ID}-{}-00", LAUNCH_SPANS[0]); + let mut params = process_start_params( + PRIVATE_PAYLOAD, + json!([missing_executable, PRIVATE_PAYLOAD]), + PathUri::from_host_native_path(directory.path()).expect("cwd URI"), + ); + params["metadata"] = process_metadata(/*index*/ 0); + let response = + request_result(&server.router, &handler, EXEC_METHOD, params, Some(&launch)).await; + let Some(RpcServerOutboundMessage::Error { error, .. }) = response else { + panic!("missing executable must fail to spawn: {response:?}"); + }; + assert_eq!(error.code, -32603); + assert!( + !error.message.is_empty(), + "caller still receives the spawn error" + ); + server.sessions.shutdown().await; + }); + }); + assert_exported_attributes( + &records, + vec![expected_process_attributes( + "codex.exec_server.process_spawn_failed", + /*index*/ 0, + "None", + )], + ); +} + +#[cfg(target_os = "macos")] +#[test] +fn exported_sandbox_denial_keeps_launch_identity_and_separate_exit_outcome() { + let records = exported_logs(/*export_traces*/ false, |runtime, _, _| { + runtime.block_on(async { + let server = TestServer::new(); + let (handler, _) = server.new_orchestrator("sandbox-denial-orchestrator", Some("original")).await; + let directory = tempfile::tempdir().expect("process directory"); + let private_file = directory.path().join(PRIVATE_PAYLOAD); + std::fs::write(&private_file, PRIVATE_PAYLOAD).expect("write test file"); + let cwd = PathUri::from_host_native_path(directory.path()).expect("cwd URI"); + let sandbox = crate::FileSystemSandboxContext::from_legacy_sandbox_policy( + codex_protocol::protocol::SandboxPolicy::new_read_only_policy(), cwd.clone(), + ).expect("read-only sandbox"); + for (index, span) in LAUNCH_SPANS.into_iter().enumerate() { + let process_id = format!("{PRIVATE_PAYLOAD}-{index}"); + let launch = format!("00-{TRACE_ID}-{span}-00"); + let argv = if index == 0 { + json!(["/bin/cat", private_file]) + } else { + // Emit private output, then attempt a real denied write. Exit 23 is produced + // by the child shell only after that write fails, not by sandbox startup. + json!(["/bin/sh", "-c", + "printf '%s' \"$PRIVATE_TEST_VALUE\"; if printf changed > \"$1\"; then exit 0; else exit 23; fi", + PRIVATE_PAYLOAD, private_file]) + }; + let mut params = process_start_params(&process_id, argv, cwd.clone()); + params["metadata"] = process_metadata(index); + params["env"]["LC_ALL"] = json!("C"); + params["sandbox"] = json!(sandbox); + let started: crate::protocol::ExecResponse = serde_json::from_value(request( + &server.router, &handler, EXEC_METHOD, params, Some(&launch), + ).await).expect("start sandboxed process"); + assert_eq!(started.sandbox_type, Some(crate::protocol::ProcessSandboxType::MacosSeatbelt)); + let output = read_until_closed(&server.router, &handler, &process_id).await; + let stdout: Vec = output.chunks.iter() + .filter(|chunk| chunk.stream == crate::protocol::ExecOutputStream::Stdout) + .flat_map(|chunk| chunk.chunk.0.iter().copied()).collect(); + assert_eq!(stdout, PRIVATE_PAYLOAD.as_bytes()); + assert_eq!((output.exit_code, output.sandbox_denied), if index == 0 { + (Some(0), false) + } else { + (Some(23), true) + }); + if index == 1 { + let stderr: Vec = output.chunks.iter() + .filter(|chunk| chunk.stream == crate::protocol::ExecOutputStream::Stderr) + .flat_map(|chunk| chunk.chunk.0.iter().copied()).collect(); + assert!(String::from_utf8_lossy(&stderr).contains("Operation not permitted")); + } + } + assert_eq!(std::fs::read(&private_file).expect("read test file"), PRIVATE_PAYLOAD.as_bytes()); + server.sessions.shutdown().await; + }); + }); + let mut expected = Vec::new(); + for index in 0..2 { + expected.push(expected_process_attributes( + "codex.exec_server.process_start", + index, + "MacosSeatbelt", + )); + let mut exited = + expected_process_attributes("codex.exec_server.process_exit", index, "MacosSeatbelt"); + exited["process.exit_code"] = json!({"intValue": if index == 0 { "0" } else { "23" }}); + exited["process.termination_requested"] = json!({"boolValue": false}); + expected.push(exited); + } + let mut denied = expected_process_attributes( + "codex.exec_server.sandbox_denied", + /*index*/ 1, + "MacosSeatbelt", + ); + denied["reason"] = json!({"stringValue": "inferred_denial"}); + expected.push(denied); + assert_exported_attributes(&records, expected); +} + +async fn read_until_closed( + router: &RpcRouter, + handler: &Arc, + process_id: &str, +) -> ReadResponse { + tokio::time::timeout(Duration::from_secs(5), async { + loop { + let output: ReadResponse = serde_json::from_value( + request( + router, + handler, + EXEC_READ_METHOD, + json!({"processId": process_id, "waitMs": 100}), + Some(LATER_TRACE), + ) + .await, + ) + .expect("read process exit"); + if output.closed { + return output; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("process exits") +} + +fn expected_process_attributes(name: &str, index: usize, sandbox: &str) -> Value { + json!({ + "event.name": {"stringValue": name}, + "launch.trace_id": {"stringValue": TRACE_ID}, + "launch.span_id": {"stringValue": LAUNCH_SPANS[index]}, + "conversation.id": {"stringValue": THREADS[index]}, + "tool.call_id": {"stringValue": CALLS[index]}, + "executor.environment_id": {"stringValue": "environment"}, + "executor.registration_id": {"stringValue": "original"}, + "sandbox.type": {"stringValue": sandbox}, + }) +} + +fn assert_exported_attributes(records: &[Value], mut expected: Vec) { + let mut actual: Vec = records + .iter() + .map(|record| { + Value::Object( + record["attributes"] + .as_array() + .expect("log attributes") + .iter() + .map(|attribute| { + ( + attribute["key"] + .as_str() + .expect("attribute key") + .to_string(), + attribute["value"].clone(), + ) + }) + .collect(), + ) + }) + .collect(); + actual.sort_by_key(Value::to_string); + expected.sort_by_key(Value::to_string); + assert_eq!(actual, expected); +} + +async fn request( + router: &RpcRouter, + handler: &Arc, + method: &str, + params: Value, + traceparent: Option<&str>, +) -> Value { + let response = request_result(router, handler, method, params, traceparent).await; + let Some(RpcServerOutboundMessage::Response { result, .. }) = response else { + panic!("request failed: {response:?}"); + }; + result +} + +async fn request_result( + router: &RpcRouter, + handler: &Arc, + method: &str, + params: Value, + traceparent: Option<&str>, +) -> Option { + let request = JSONRPCRequest { + id: RequestId::Integer(1), + method: method.to_string(), + params: Some(params), + trace: traceparent.map(|traceparent| W3cTraceContext { + traceparent: Some(traceparent.to_string()), + tracestate: None, + }), + }; + let JsonRpcConnectionEvent::QueuedRequest { + request, + request_span, + .. + } = JsonRpcConnectionEvent::message(JSONRPCMessage::Request(request)) + else { + panic!("request span") + }; + let (_, route) = router + .request_route(method) + .expect("registered request route"); + request_span.record("otel.name", method); + route(Arc::clone(handler), request) + .instrument(request_span) + .await +} + +fn attribute<'a>(record: &'a Value, key: &str) -> Option<&'a str> { + record["attributes"] + .as_array() + .expect("log attributes") + .iter() + .find(|attribute| attribute["key"] == key) + .map(|attribute| { + attribute["value"]["stringValue"] + .as_str() + .expect("string attribute") + }) +} + +fn exported_logs( + export_traces: bool, + run: impl FnOnce(&tokio::runtime::Runtime, &OtelProvider, &MockServer), +) -> Vec { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .expect("test runtime"); + let collector = runtime.block_on(MockServer::start()); + runtime.block_on( + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200)) + .mount(&collector), + ); + let exporter = OtelExporter::OtlpHttp { + endpoint: format!("{}/v1/logs", collector.uri()), + headers: HashMap::new(), + protocol: OtelHttpProtocol::Json, + tls: None, + }; + let otel = OtelProvider::try_new(&OtelSettings { + environment: "test".to_string(), + service_name: "codex-exec-server".to_string(), + service_version: env!("CARGO_PKG_VERSION").to_string(), + codex_home: std::env::current_dir().expect("cwd"), + exporter, + trace_exporter: if export_traces { + OtelExporter::OtlpHttp { + endpoint: format!("{}/v1/traces", collector.uri()), + headers: HashMap::new(), + protocol: OtelHttpProtocol::Json, + tls: None, + } + } else { + OtelExporter::None + }, + metrics_exporter: OtelExporter::None, + runtime_metrics: false, + span_attributes: BTreeMap::new(), + tracestate: BTreeMap::new(), + }) + .expect("OTEL settings") + .expect("OTEL provider"); + let subscriber = tracing_subscriber::registry() + .with(otel.tracing_layer()) + .with(otel.logger_layer()); + tracing::subscriber::with_default(subscriber, || { + tracing::callsite::rebuild_interest_cache(); + run(&runtime, &otel, &collector); + }); + runtime + .block_on(otel.shutdown_with_timeout(Duration::from_secs(5))) + .expect("flush OTEL"); + + let mut records = Vec::new(); + for request in runtime + .block_on(collector.received_requests()) + .expect("exported requests") + { + let body: Value = serde_json::from_slice(&request.body).expect("OTLP JSON"); + let Some(resources) = body["resourceLogs"].as_array() else { + continue; + }; + assert!(!String::from_utf8_lossy(&request.body).contains(PRIVATE_PAYLOAD)); + for resource in resources { + for scope in resource["scopeLogs"].as_array().expect("scope logs") { + for record in scope["logRecords"].as_array().expect("log records") { + assert!(record["body"].is_null(), "lifecycle logs have no raw body"); + records.push(record.clone()); + } + } + } + } + records +} + +struct TestServer { + sessions: Arc, + router: RpcRouter, + outgoing: mpsc::Sender, + _notifications: mpsc::Receiver, +} + +impl TestServer { + fn new() -> Self { + let (outgoing, notifications) = mpsc::channel(/*buffer*/ 128); + Self { + sessions: SessionRegistry::new(ExecServerTelemetry::default()), + router: build_router(), + outgoing, + _notifications: notifications, + } + } + + async fn new_orchestrator( + &self, + client_name: &str, + registration: Option<&str>, + ) -> (Arc, String) { + let handler = new_handler(&self.sessions, self.outgoing.clone(), registration); + let initialized = handler + .initialize(InitializeParams { + client_name: client_name.to_string(), + resume_session_id: None, + }) + .await + .expect("initialize orchestrator"); + handler.initialized().expect("initialized orchestrator"); + (handler, initialized.session_id) + } +} + +fn process_start_params(process_id: &str, argv: Value, cwd: PathUri) -> Value { + json!({ + "processId": process_id, + "argv": argv, + "cwd": cwd, + "env": {"PRIVATE_TEST_VALUE": PRIVATE_PAYLOAD}, + "tty": false, + "arg0": null, + }) +} + +fn process_metadata(index: usize) -> Value { + json!({"threadId": THREADS[index], "toolCallId": CALLS[index]}) +} + +fn new_handler( + sessions: &Arc, + outgoing: mpsc::Sender, + registration: Option<&str>, +) -> Arc { + let mut handler = ExecServerHandler::new( + Arc::clone(sessions), + RpcNotificationSender::new(outgoing), + ExecServerRuntimePaths::new( + std::env::current_exe().expect("test executable"), + /*codex_linux_sandbox_exe*/ None, + ) + .expect("runtime paths"), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ); + handler.executor_registration = registration + .and_then(|registration| { + ExecutorRegistration::new("environment".to_string(), registration.to_string()) + }) + .map(Arc::new); + Arc::new(handler) +} diff --git a/codex-rs/exec-server/src/server/processor.rs b/codex-rs/exec-server/src/server/processor.rs new file mode 100644 index 0000000000000000000000000000000000000000..1fd2a1f1d6a4d52d0bf275ff623a59a9aa6f3ba8 --- /dev/null +++ b/codex-rs/exec-server/src/server/processor.rs @@ -0,0 +1,625 @@ +use std::sync::Arc; +use std::time::Instant; + +use codex_build_info::BuildInfo; +use codex_exec_server_protocol::JSONRPCMessage; +use tokio::sync::mpsc; +use tracing::debug; +use tracing::warn; + +use crate::ExecServerRuntimePaths; +use crate::connection::CHANNEL_CAPACITY; +use crate::connection::JsonRpcConnection; +use crate::connection::JsonRpcConnectionEvent; +use crate::rpc::RpcCallError; +use crate::rpc::RpcNotificationSender; +use crate::rpc::RpcServerOutboundMessage; +use crate::rpc::encode_server_message; +use crate::rpc_server_requests::RpcServerRequestSender; +use crate::server::ExecServerHandler; +use crate::server::RequestDispatchMode; +use crate::server::registry::build_router; +use crate::server::request_dispatcher::RequestDispatcher; +use crate::server::request_dispatcher::RequestTaskResult; +use crate::server::session_registry::SessionRegistry; +use crate::telemetry::ConnectionTransport; +use crate::telemetry::ExecServerTelemetry; +use crate::telemetry::ExecutorRegistration; +use codex_http_client::HttpClientFactory; + +#[derive(Clone)] +pub(crate) struct ConnectionProcessor { + session_registry: Arc, + runtime_paths: ExecServerRuntimePaths, + telemetry: ExecServerTelemetry, + http_client_factory: HttpClientFactory, + request_dispatch_mode: RequestDispatchMode, +} + +impl ConnectionProcessor { + #[cfg(test)] + pub(crate) fn new(runtime_paths: ExecServerRuntimePaths) -> Self { + Self::new_with_telemetry( + runtime_paths, + ExecServerTelemetry::default(), + codex_http_client::HttpClientFactory::new( + codex_http_client::OutboundProxyPolicy::ReqwestDefault, + ), + RequestDispatchMode::Inline, + ) + } + + pub(crate) fn new_with_telemetry( + runtime_paths: ExecServerRuntimePaths, + telemetry: ExecServerTelemetry, + http_client_factory: HttpClientFactory, + request_dispatch_mode: RequestDispatchMode, + ) -> Self { + // Library callers may bypass CLI startup. Capture the version before serving clients. + let _ = BuildInfo::get(); + Self { + session_registry: SessionRegistry::new(telemetry.clone()), + runtime_paths, + telemetry, + http_client_factory, + request_dispatch_mode, + } + } + + pub(crate) async fn run_connection( + &self, + connection: JsonRpcConnection, + transport: ConnectionTransport, + ) { + run_connection( + connection, + self.clone(), + transport, + /*executor_registration*/ None, + ) + .await; + } + + pub(crate) async fn run_registered_connection( + &self, + connection: JsonRpcConnection, + executor_registration: Option>, + ) { + run_connection( + connection, + self.clone(), + ConnectionTransport::Relay, + executor_registration, + ) + .await; + } + + pub(crate) async fn shutdown(&self) { + self.session_registry.shutdown().await; + } +} + +async fn run_connection( + connection: JsonRpcConnection, + processor: ConnectionProcessor, + transport: ConnectionTransport, + executor_registration: Option>, +) { + let ConnectionProcessor { + session_registry, + runtime_paths, + telemetry, + http_client_factory, + request_dispatch_mode, + } = processor; + let _connection_metrics = telemetry.connection_started(transport); + let JsonRpcConnection { + outgoing_tx: json_outgoing_tx, + mut incoming_rx, + mut disconnected_rx, + task_handles: connection_tasks, + transport: _transport, + } = connection; + let (outgoing_tx, mut outgoing_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let notifications = RpcNotificationSender::new(outgoing_tx.clone()); + let requests = notifications.request_sender(); + let mut handler = ExecServerHandler::new( + session_registry, + notifications, + runtime_paths, + http_client_factory, + ); + handler.executor_registration = executor_registration; + let handler = Arc::new(handler); + + let outbound_task = tokio::spawn(async move { + while let Some(message) = outgoing_rx.recv().await { + let json_message = match encode_server_message(message) { + Ok(json_message) => json_message, + Err(err) => { + warn!("failed to serialize exec-server outbound message: {err}"); + break; + } + }; + if json_outgoing_tx.send(json_message).await.is_err() { + break; + } + } + }); + + let mut dispatcher = RequestDispatcher::new( + Arc::new(build_router()), + Arc::clone(&handler), + outgoing_tx.clone(), + disconnected_rx.clone(), + requests.clone(), + telemetry, + request_dispatch_mode, + ); + + loop { + let has_request_tasks = dispatcher.has_tasks(); + let event = tokio::select! { + result = dispatcher.join_next(), if has_request_tasks => { + if result == RequestTaskResult::ConnectionClosed { + break; + } + continue; + } + _ = disconnected_rx.changed() => { + debug!("exec-server transport disconnected"); + break; + } + event = incoming_rx.recv() => { + let Some(event) = event else { + break; + }; + event + } + }; + + if !handler.is_session_attached() { + debug!("exec-server connection evicted after session resume"); + break; + } + + let result = match event { + JsonRpcConnectionEvent::MalformedMessage { reason } => { + dispatcher.handle_malformed_message(reason).await + } + JsonRpcConnectionEvent::Message(message) => match message { + JSONRPCMessage::Request(request) => { + dispatcher + .dispatch_request(request, tracing::Span::none(), Instant::now()) + .await + } + JSONRPCMessage::Notification(notification) => { + dispatcher.handle_notification(notification).await + } + JSONRPCMessage::Response(response) => dispatcher.handle_response(response), + JSONRPCMessage::Error(error) => dispatcher.handle_error(error), + }, + JsonRpcConnectionEvent::QueuedRequest { + request, + request_span, + queued_at, + } => { + dispatcher + .dispatch_request(request, request_span, queued_at) + .await + } + JsonRpcConnectionEvent::Disconnected { reason } => { + if let Some(reason) = reason { + debug!("exec-server connection disconnected: {reason}"); + } + break; + } + }; + if result == RequestTaskResult::ConnectionClosed { + break; + } + } + + if *disconnected_rx.borrow() { + complete_queued_client_responses(&requests, &mut incoming_rx); + } + requests.close(); + dispatcher.shutdown().await; + handler.shutdown().await; + drop(handler); + drop(requests); + drop(outgoing_tx); + for task in connection_tasks { + task.abort(); + let _ = task.await; + } + let _ = outbound_task.await; +} + +fn complete_queued_client_responses( + requests: &RpcServerRequestSender, + incoming_rx: &mut mpsc::Receiver, +) { + while let Ok(event) = incoming_rx.try_recv() { + let (request_id, response) = match event { + JsonRpcConnectionEvent::Message(JSONRPCMessage::Response(response)) => { + (response.id, Ok(response.result)) + } + JsonRpcConnectionEvent::Message(JSONRPCMessage::Error(error)) => { + (error.id, Err(RpcCallError::Server(error.error))) + } + JsonRpcConnectionEvent::Message( + JSONRPCMessage::Request(_) | JSONRPCMessage::Notification(_), + ) + | JsonRpcConnectionEvent::QueuedRequest { .. } + | JsonRpcConnectionEvent::MalformedMessage { .. } + | JsonRpcConnectionEvent::Disconnected { .. } => continue, + }; + if !requests.complete(request_id.clone(), response) { + warn!("ignoring unexpected client response while disconnecting: {request_id:?}"); + } + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + use std::sync::Arc; + use std::time::Duration; + + use codex_exec_server_protocol::JSONRPCMessage; + use codex_exec_server_protocol::JSONRPCNotification; + use codex_exec_server_protocol::JSONRPCRequest; + use codex_exec_server_protocol::JSONRPCResponse; + use codex_exec_server_protocol::RequestId; + use codex_utils_path_uri::PathUri; + use pretty_assertions::assert_eq; + use serde::Serialize; + use serde::de::DeserializeOwned; + use tokio::io::AsyncBufReadExt; + use tokio::io::AsyncWriteExt; + use tokio::io::BufReader; + use tokio::io::DuplexStream; + use tokio::io::Lines; + use tokio::io::duplex; + use tokio::sync::mpsc; + use tokio::task::JoinHandle; + use tokio::time::timeout; + + use super::complete_queued_client_responses; + use super::run_connection; + use crate::ExecServerRuntimePaths; + use crate::ProcessId; + use crate::connection::JsonRpcConnection; + use crate::connection::JsonRpcConnectionEvent; + use crate::protocol::ENVIRONMENT_INFO_METHOD; + use crate::protocol::ENVIRONMENT_STATUS_METHOD; + use crate::protocol::EXEC_METHOD; + use crate::protocol::EXEC_READ_METHOD; + use crate::protocol::EXEC_TERMINATE_METHOD; + use crate::protocol::EnvironmentInfo; + use crate::protocol::EnvironmentStatus; + use crate::protocol::EnvironmentStatusKind; + use crate::protocol::ExecParams; + use crate::protocol::ExecResponse; + use crate::protocol::ExecServerNetworkPolicyDecision; + use crate::protocol::INITIALIZE_METHOD; + use crate::protocol::INITIALIZED_METHOD; + use crate::protocol::InitializeParams; + use crate::protocol::InitializeResponse; + use crate::protocol::NETWORK_POLICY_REQUEST_METHOD; + use crate::protocol::NetworkPolicyRequestResponse; + use crate::protocol::ReadParams; + use crate::protocol::TerminateParams; + use crate::protocol::TerminateResponse; + use crate::rpc::RpcServerOutboundMessage; + use crate::rpc_server_requests::RpcServerRequestSender; + use crate::server::session_registry::SessionRegistry; + + #[tokio::test] + async fn connection_accepts_pipelined_scalar_requests() { + let registry = SessionRegistry::new(crate::ExecServerTelemetry::default()); + let (mut writer, mut lines, task) = spawn_test_connection(registry, "pipelined-scalar"); + + send_request( + &mut writer, + /*id*/ 1, + INITIALIZE_METHOD, + &InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + }, + ) + .await; + let _: InitializeResponse = read_response(&mut lines, /*expected_id*/ 1).await; + send_notification(&mut writer, INITIALIZED_METHOD, &()).await; + + send_request(&mut writer, /*id*/ 2, ENVIRONMENT_INFO_METHOD, &()).await; + send_request(&mut writer, /*id*/ 3, ENVIRONMENT_INFO_METHOD, &()).await; + send_request(&mut writer, /*id*/ 4, ENVIRONMENT_STATUS_METHOD, &()).await; + + let _: EnvironmentInfo = read_response(&mut lines, /*expected_id*/ 2).await; + let _: EnvironmentInfo = read_response(&mut lines, /*expected_id*/ 3).await; + assert_eq!( + read_response::(&mut lines, /*expected_id*/ 4).await, + EnvironmentStatus { + status: EnvironmentStatusKind::Ready, + } + ); + + drop(writer); + drop(lines); + timeout(Duration::from_secs(1), task) + .await + .expect("processor should exit") + .expect("processor should join"); + } + + /// A callback response received before EOF must survive transport shutdown. + #[tokio::test] + async fn disconnect_completes_queued_network_policy_response() { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(/*buffer*/ 1); + let requests = RpcServerRequestSender::new(outgoing_tx); + let pending_request = { + let requests = requests.clone(); + tokio::spawn(async move { + requests + .call_with_timeout::<_, NetworkPolicyRequestResponse>( + NETWORK_POLICY_REQUEST_METHOD, + &(), + Duration::from_secs(1), + ) + .await + }) + }; + let RpcServerOutboundMessage::Request(request) = outgoing_rx + .recv() + .await + .expect("network policy request should be queued") + else { + panic!("expected outbound network policy request"); + }; + + let (incoming_tx, mut incoming_rx) = mpsc::channel(/*buffer*/ 2); + incoming_tx + .try_send(JsonRpcConnectionEvent::Message(JSONRPCMessage::Response( + JSONRPCResponse { + id: request.id, + result: serde_json::to_value(NetworkPolicyRequestResponse { + decision: ExecServerNetworkPolicyDecision::Allow, + }) + .expect("serialize network policy response"), + }, + ))) + .expect("queue network policy response"); + incoming_tx + .try_send(JsonRpcConnectionEvent::Disconnected { reason: None }) + .expect("queue transport disconnect"); + + complete_queued_client_responses(&requests, &mut incoming_rx); + requests.close(); + + assert_eq!( + pending_request + .await + .expect("network policy request should join") + .expect("network policy request should complete"), + NetworkPolicyRequestResponse { + decision: ExecServerNetworkPolicyDecision::Allow, + } + ); + } + + #[tokio::test] + async fn transport_disconnect_detaches_session_during_in_flight_read() { + let registry = SessionRegistry::new(crate::ExecServerTelemetry::default()); + let (mut first_writer, mut first_lines, first_task) = + spawn_test_connection(Arc::clone(®istry), "first"); + + send_request( + &mut first_writer, + /*id*/ 1, + INITIALIZE_METHOD, + &InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + }, + ) + .await; + let initialize_response: InitializeResponse = + read_response(&mut first_lines, /*expected_id*/ 1).await; + send_notification(&mut first_writer, INITIALIZED_METHOD, &()).await; + + let process_id = ProcessId::from("proc-long-poll"); + send_request( + &mut first_writer, + /*id*/ 2, + EXEC_METHOD, + &exec_params(process_id.clone()), + ) + .await; + let _: ExecResponse = read_response(&mut first_lines, /*expected_id*/ 2).await; + + send_request( + &mut first_writer, + /*id*/ 3, + EXEC_READ_METHOD, + &ReadParams { + process_id: process_id.clone(), + after_seq: None, + max_bytes: None, + wait_ms: Some(5_000), + }, + ) + .await; + drop(first_writer); + tokio::time::sleep(Duration::from_millis(25)).await; + + let (mut second_writer, mut second_lines, second_task) = + spawn_test_connection(Arc::clone(®istry), "second"); + send_request( + &mut second_writer, + /*id*/ 1, + INITIALIZE_METHOD, + &InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: Some(initialize_response.session_id.clone()), + }, + ) + .await; + let second_initialize_response = timeout( + Duration::from_secs(1), + read_response::(&mut second_lines, /*expected_id*/ 1), + ) + .await + .expect("resume initialize should not wait for the old read to finish"); + assert_eq!( + second_initialize_response.session_id, + initialize_response.session_id + ); + timeout(Duration::from_secs(1), first_task) + .await + .expect("first processor should exit") + .expect("first processor should join"); + send_notification(&mut second_writer, INITIALIZED_METHOD, &()).await; + + send_request( + &mut second_writer, + /*id*/ 2, + EXEC_TERMINATE_METHOD, + &TerminateParams { process_id }, + ) + .await; + let _: TerminateResponse = read_response(&mut second_lines, /*expected_id*/ 2).await; + + drop(second_writer); + drop(second_lines); + timeout(Duration::from_secs(1), second_task) + .await + .expect("second processor should exit") + .expect("second processor should join"); + } + + fn spawn_test_connection( + registry: Arc, + label: &str, + ) -> (DuplexStream, Lines>, JoinHandle<()>) { + let (client_writer, server_reader) = duplex(1 << 20); + let (server_writer, client_reader) = duplex(1 << 20); + let connection = + JsonRpcConnection::from_stdio(server_reader, server_writer, label.to_string()); + let task = tokio::spawn(run_connection( + connection, + super::ConnectionProcessor { + session_registry: registry, + ..super::ConnectionProcessor::new(test_runtime_paths()) + }, + crate::telemetry::ConnectionTransport::Stdio, + /*executor_registration*/ None, + )); + (client_writer, BufReader::new(client_reader).lines(), task) + } + + fn test_runtime_paths() -> ExecServerRuntimePaths { + ExecServerRuntimePaths::new( + std::env::current_exe().expect("current exe"), + /*codex_linux_sandbox_exe*/ None, + ) + .expect("runtime paths") + } + + async fn send_request( + writer: &mut DuplexStream, + id: i64, + method: &str, + params: &P, + ) { + write_message( + writer, + &JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(id), + method: method.to_string(), + params: Some(serde_json::to_value(params).expect("serialize params")), + trace: None, + }), + ) + .await; + } + + async fn send_notification(writer: &mut DuplexStream, method: &str, params: &P) { + write_message( + writer, + &JSONRPCMessage::Notification(JSONRPCNotification { + method: method.to_string(), + params: Some(serde_json::to_value(params).expect("serialize params")), + }), + ) + .await; + } + + async fn write_message(writer: &mut DuplexStream, message: &JSONRPCMessage) { + let encoded = serde_json::to_vec(message).expect("serialize JSON-RPC message"); + writer.write_all(&encoded).await.expect("write request"); + writer.write_all(b"\n").await.expect("write newline"); + } + + async fn read_response( + lines: &mut Lines>, + expected_id: i64, + ) -> T { + let line = lines + .next_line() + .await + .expect("read response") + .expect("response line"); + match serde_json::from_str::(&line).expect("decode JSON-RPC response") { + JSONRPCMessage::Response(JSONRPCResponse { id, result }) => { + assert_eq!(id, RequestId::Integer(expected_id)); + serde_json::from_value(result).expect("decode response result") + } + JSONRPCMessage::Error(error) => panic!("unexpected JSON-RPC error: {error:?}"), + other => panic!("expected JSON-RPC response, got {other:?}"), + } + } + + fn exec_params(process_id: ProcessId) -> ExecParams { + let mut env = HashMap::new(); + if let Some(path) = std::env::var_os("PATH") { + env.insert("PATH".to_string(), path.to_string_lossy().into_owned()); + } + ExecParams { + metadata: Default::default(), + process_id, + argv: sleep_then_print_argv(), + cwd: PathUri::from_host_native_path(std::env::current_dir().expect("cwd")) + .expect("cwd URI"), + shell_snapshot: None, + env_policy: None, + env, + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + } + } + + fn sleep_then_print_argv() -> Vec { + if cfg!(windows) { + vec![ + std::env::var("COMSPEC").unwrap_or_else(|_| "cmd.exe".to_string()), + "/C".to_string(), + "ping -n 3 127.0.0.1 >NUL && echo late".to_string(), + ] + } else { + vec![ + "/bin/sh".to_string(), + "-c".to_string(), + "sleep 1; printf late".to_string(), + ] + } + } +} diff --git a/codex-rs/exec-server/src/server/registry.rs b/codex-rs/exec-server/src/server/registry.rs new file mode 100644 index 0000000000000000000000000000000000000000..4cd21ad85e55cdc7f871ac918e60ca0b82dee3d5 --- /dev/null +++ b/codex-rs/exec-server/src/server/registry.rs @@ -0,0 +1,199 @@ +use opentelemetry::trace::TraceContextExt; +use std::sync::Arc; + +use crate::protocol::CAPABILITY_ROOTS_DISCOVER_METHOD; +use crate::protocol::CapabilityRootsDiscoverParams; +use crate::protocol::ENVIRONMENT_CONFIG_READ_METHOD; +use crate::protocol::ENVIRONMENT_INFO_METHOD; +use crate::protocol::ENVIRONMENT_STATUS_METHOD; +use crate::protocol::EXEC_METHOD; +use crate::protocol::EXEC_READ_METHOD; +use crate::protocol::EXEC_SIGNAL_METHOD; +use crate::protocol::EXEC_TERMINATE_METHOD; +use crate::protocol::EXEC_WRITE_METHOD; +use crate::protocol::EnvironmentConfigReadParams; +use crate::protocol::ExecParams; +use crate::protocol::FS_CANONICALIZE_METHOD; +use crate::protocol::FS_CLOSE_METHOD; +use crate::protocol::FS_COPY_METHOD; +use crate::protocol::FS_CREATE_DIRECTORY_METHOD; +use crate::protocol::FS_GET_METADATA_METHOD; +use crate::protocol::FS_OPEN_METHOD; +use crate::protocol::FS_READ_BLOCK_METHOD; +use crate::protocol::FS_READ_DIRECTORY_METHOD; +use crate::protocol::FS_READ_FILE_METHOD; +use crate::protocol::FS_REMOVE_METHOD; +use crate::protocol::FS_WALK_METHOD; +use crate::protocol::FS_WRITE_FILE_METHOD; +use crate::protocol::FsCanonicalizeParams; +use crate::protocol::FsCloseParams; +use crate::protocol::FsCopyParams; +use crate::protocol::FsCreateDirectoryParams; +use crate::protocol::FsGetMetadataParams; +use crate::protocol::FsOpenParams; +use crate::protocol::FsReadBlockParams; +use crate::protocol::FsReadDirectoryParams; +use crate::protocol::FsReadFileParams; +use crate::protocol::FsRemoveParams; +use crate::protocol::FsWalkParams; +use crate::protocol::FsWriteFileParams; +use crate::protocol::HTTP_REQUEST_METHOD; +use crate::protocol::HttpRequestParams; +use crate::protocol::INITIALIZE_METHOD; +use crate::protocol::INITIALIZED_METHOD; +use crate::protocol::InitializeParams; +use crate::protocol::ReadParams; +use crate::protocol::SignalParams; +use crate::protocol::TerminateParams; +use crate::protocol::WriteParams; +use crate::rpc::RpcRouter; +use crate::server::ExecServerHandler; + +pub(crate) fn build_router() -> RpcRouter { + let mut router = RpcRouter::new(); + router.notification( + INITIALIZED_METHOD, + |handler: Arc, _params: serde_json::Value| async move { + handler.initialized() + }, + ); + router.request( + INITIALIZE_METHOD, + |handler: Arc, params: InitializeParams| async move { + handler.initialize(params).await + }, + ); + router.request_with_id( + HTTP_REQUEST_METHOD, + |handler: Arc, request_id, params: HttpRequestParams| async move { + handler.http_request(request_id, params).await + }, + ); + router.request_with_trace( + EXEC_METHOD, + |handler: Arc, params: ExecParams, trace| async move { + let launch_context = trace + .as_ref() + .and_then(codex_otel::context_from_w3c_trace_context) + .map(|context| context.span().span_context().clone()); + handler.exec(params, launch_context).await + }, + ); + router.request( + ENVIRONMENT_INFO_METHOD, + |handler: Arc, _params: ()| async move { handler.environment_info() }, + ); + router.request( + ENVIRONMENT_CONFIG_READ_METHOD, + |handler: Arc, params: EnvironmentConfigReadParams| async move { + handler.environment_config_read(params).await + }, + ); + router.request( + ENVIRONMENT_STATUS_METHOD, + |handler: Arc, _params: ()| async move { handler.environment_status() }, + ); + router.request( + CAPABILITY_ROOTS_DISCOVER_METHOD, + |handler: Arc, params: CapabilityRootsDiscoverParams| async move { + handler.discover_capability_roots(params).await + }, + ); + router.request( + EXEC_READ_METHOD, + |handler: Arc, params: ReadParams| async move { + handler.exec_read(params).await + }, + ); + router.request( + EXEC_WRITE_METHOD, + |handler: Arc, params: WriteParams| async move { + handler.exec_write(params).await + }, + ); + router.request( + EXEC_SIGNAL_METHOD, + |handler: Arc, params: SignalParams| async move { + handler.signal(params).await + }, + ); + router.request( + EXEC_TERMINATE_METHOD, + |handler: Arc, params: TerminateParams| async move { + handler.terminate(params).await + }, + ); + router.request( + FS_READ_FILE_METHOD, + |handler: Arc, params: FsReadFileParams| async move { + handler.fs_read_file(params).await + }, + ); + router.request( + FS_OPEN_METHOD, + |handler: Arc, params: FsOpenParams| async move { + handler.fs_open(params).await + }, + ); + router.request( + FS_READ_BLOCK_METHOD, + |handler: Arc, params: FsReadBlockParams| async move { + handler.fs_read_block(params).await + }, + ); + router.request( + FS_CLOSE_METHOD, + |handler: Arc, params: FsCloseParams| async move { + handler.fs_close(params).await + }, + ); + router.request( + FS_WRITE_FILE_METHOD, + |handler: Arc, params: FsWriteFileParams| async move { + handler.fs_write_file(params).await + }, + ); + router.request( + FS_CREATE_DIRECTORY_METHOD, + |handler: Arc, params: FsCreateDirectoryParams| async move { + handler.fs_create_directory(params).await + }, + ); + router.request( + FS_GET_METADATA_METHOD, + |handler: Arc, params: FsGetMetadataParams| async move { + handler.fs_get_metadata(params).await + }, + ); + router.request( + FS_CANONICALIZE_METHOD, + |handler: Arc, params: FsCanonicalizeParams| async move { + handler.fs_canonicalize(params).await + }, + ); + router.request( + FS_READ_DIRECTORY_METHOD, + |handler: Arc, params: FsReadDirectoryParams| async move { + handler.fs_read_directory(params).await + }, + ); + router.request( + FS_WALK_METHOD, + |handler: Arc, params: FsWalkParams| async move { + handler.fs_walk(params).await + }, + ); + router.request( + FS_REMOVE_METHOD, + |handler: Arc, params: FsRemoveParams| async move { + handler.fs_remove(params).await + }, + ); + router.request( + FS_COPY_METHOD, + |handler: Arc, params: FsCopyParams| async move { + handler.fs_copy(params).await + }, + ); + router +} diff --git a/codex-rs/exec-server/src/server/release_version.rs b/codex-rs/exec-server/src/server/release_version.rs new file mode 100644 index 0000000000000000000000000000000000000000..a68fb0e9c2a0ff01e80258af4fd61e1306400a35 --- /dev/null +++ b/codex-rs/exec-server/src/server/release_version.rs @@ -0,0 +1,12 @@ +//! Include the startup-cached package release version in executor metadata. + +use codex_build_info::BuildInfo; + +use crate::protocol::EnvironmentInfo; + +pub(super) fn local_environment_info() -> EnvironmentInfo { + EnvironmentInfo { + executor_version: BuildInfo::get().version().to_string(), + ..EnvironmentInfo::local() + } +} diff --git a/codex-rs/exec-server/src/server/request_dispatcher.rs b/codex-rs/exec-server/src/server/request_dispatcher.rs new file mode 100644 index 0000000000000000000000000000000000000000..c1f16b64fed0c4a88c94bf06070c33bbfe5b26f8 --- /dev/null +++ b/codex-rs/exec-server/src/server/request_dispatcher.rs @@ -0,0 +1,374 @@ +use std::num::NonZeroUsize; +use std::num::ParseIntError; +use std::str::FromStr; +use std::sync::Arc; +use std::time::Instant; + +use codex_exec_server_protocol::JSONRPCError; +use codex_exec_server_protocol::JSONRPCNotification; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::RequestId; +use tokio::sync::Semaphore; +use tokio::sync::mpsc; +use tokio::sync::watch; +use tokio::task::JoinSet; +use tracing::Instrument; +use tracing::debug; +use tracing::warn; + +use crate::protocol::ENVIRONMENT_INFO_METHOD; +use crate::protocol::ENVIRONMENT_STATUS_METHOD; +use crate::protocol::EXEC_SIGNAL_METHOD; +use crate::protocol::EXEC_TERMINATE_METHOD; +use crate::protocol::FS_CLOSE_METHOD; +use crate::protocol::INITIALIZE_METHOD; +use crate::protocol::INITIALIZED_METHOD; +use crate::rpc::RpcCallError; +use crate::rpc::RpcRouter; +use crate::rpc::RpcServerOutboundMessage; +use crate::rpc::invalid_request; +use crate::rpc::method_not_found; +use crate::rpc_server_requests::RpcServerRequestSender; +use crate::server::ExecServerHandler; +use crate::telemetry::ExecServerTelemetry; + +pub(super) struct RequestDispatcher { + router: Arc>, + handler: Arc, + outgoing_tx: mpsc::Sender, + disconnected_rx: watch::Receiver, + requests: RpcServerRequestSender, + telemetry: ExecServerTelemetry, + lanes: Option, + tasks: JoinSet, + initialized: bool, +} + +impl RequestDispatcher { + pub(super) fn new( + router: Arc>, + handler: Arc, + outgoing_tx: mpsc::Sender, + disconnected_rx: watch::Receiver, + requests: RpcServerRequestSender, + telemetry: ExecServerTelemetry, + mode: RequestDispatchMode, + ) -> Self { + let lanes = match mode { + RequestDispatchMode::Inline => None, + RequestDispatchMode::Concurrent { + max_concurrent_requests, + } => Some(RequestLanes { + ordinary: Arc::new(Semaphore::new(max_concurrent_requests.get())), + control: Arc::new(Semaphore::new(max_concurrent_requests.get())), + }), + }; + + Self { + router, + handler, + outgoing_tx, + disconnected_rx, + requests, + telemetry, + lanes, + tasks: JoinSet::new(), + initialized: false, + } + } + + pub(super) fn has_tasks(&self) -> bool { + !self.tasks.is_empty() + } + + pub(super) async fn join_next(&mut self) -> RequestTaskResult { + match self.tasks.join_next().await { + Some(Ok(result)) => result, + Some(Err(error)) => { + warn!("exec-server request task failed: {error}"); + RequestTaskResult::ConnectionClosed + } + None => RequestTaskResult::Completed, + } + } + + pub(super) async fn handle_malformed_message(&self, reason: String) -> RequestTaskResult { + warn!("ignoring malformed exec-server message: {reason}"); + if self + .outgoing_tx + .send(RpcServerOutboundMessage::Error { + request_id: RequestId::Integer(-1), + error: invalid_request(reason), + }) + .await + .is_err() + { + return RequestTaskResult::ConnectionClosed; + } + + RequestTaskResult::Completed + } + + pub(super) async fn handle_notification( + &mut self, + notification: JSONRPCNotification, + ) -> RequestTaskResult { + let is_initialized = notification.method == INITIALIZED_METHOD; + let Some(route) = self.router.notification_route(notification.method.as_str()) else { + warn!( + "closing exec-server connection after unexpected notification: {}", + notification.method + ); + return RequestTaskResult::ConnectionClosed; + }; + let result = tokio::select! { + result = route(Arc::clone(&self.handler), notification) => result, + _ = self.disconnected_rx.changed() => { + debug!("exec-server transport disconnected while handling notification"); + return RequestTaskResult::ConnectionClosed; + } + }; + if let Err(error) = result { + warn!("closing exec-server connection after protocol error: {error}"); + return RequestTaskResult::ConnectionClosed; + } + if is_initialized { + self.initialized = true; + } + + RequestTaskResult::Completed + } + + pub(super) fn handle_response(&self, response: JSONRPCResponse) -> RequestTaskResult { + if !self + .requests + .complete(response.id.clone(), Ok(response.result)) + { + warn!( + "closing exec-server connection after unexpected client response: {:?}", + response.id + ); + return RequestTaskResult::ConnectionClosed; + } + + RequestTaskResult::Completed + } + + pub(super) fn handle_error(&self, error: JSONRPCError) -> RequestTaskResult { + if !self + .requests + .complete(error.id.clone(), Err(RpcCallError::Server(error.error))) + { + warn!( + "closing exec-server connection after unexpected client error: {:?}", + error.id + ); + return RequestTaskResult::ConnectionClosed; + } + + RequestTaskResult::Completed + } + + pub(super) async fn dispatch_request( + &mut self, + request: JSONRPCRequest, + request_span: tracing::Span, + queued_at: Instant, + ) -> RequestTaskResult { + let started_at = Instant::now(); + let Some((method, route)) = self.router.request_route(request.method.as_str()) else { + let method = "unknown"; + self.telemetry + .request_queue_completed(method, queued_at.elapsed()); + request_span.record("otel.name", method); + if self + .outgoing_tx + .send(RpcServerOutboundMessage::Error { + request_id: request.id, + error: method_not_found(format!( + "exec-server stub does not implement `{}` yet", + request.method + )), + }) + .await + .is_err() + { + request_span.record("result", "disconnected"); + self.telemetry.request_completed( + method, + "disconnected", + started_at.elapsed(), + queued_at.elapsed(), + ); + return RequestTaskResult::ConnectionClosed; + } + request_span.record("result", "error"); + self.telemetry.request_completed( + method, + "error", + started_at.elapsed(), + queued_at.elapsed(), + ); + return RequestTaskResult::Completed; + }; + + request_span.record("otel.name", method); + let route_setup_started_at = Instant::now(); + let route = route(Arc::clone(&self.handler), request); + let route_setup_duration = route_setup_started_at.elapsed(); + let outgoing_tx = self.outgoing_tx.clone(); + let mut disconnected_rx = self.disconnected_rx.clone(); + let telemetry = self.telemetry.clone(); + let task = async move { + telemetry.request_queue_completed( + method, + queued_at.elapsed().saturating_sub(route_setup_duration), + ); + let message = tokio::select! { + message = route.instrument(request_span.clone()) => message, + _ = disconnected_rx.changed() => { + request_span.record("result", "disconnected"); + telemetry.request_completed( + method, + "disconnected", + started_at.elapsed(), + queued_at.elapsed(), + ); + return RequestTaskResult::ConnectionClosed; + } + }; + let result = request_result(&message); + let response_sent = match message { + Some(message) => tokio::select! { + result = outgoing_tx.send(message) => result.is_ok(), + _ = disconnected_rx.changed() => false, + }, + None => true, + }; + if !response_sent { + request_span.record("result", "disconnected"); + telemetry.request_completed( + method, + "disconnected", + started_at.elapsed(), + queued_at.elapsed(), + ); + return RequestTaskResult::ConnectionClosed; + } + request_span.record("result", result); + telemetry.request_completed(method, result, started_at.elapsed(), queued_at.elapsed()); + RequestTaskResult::Completed + }; + + let Some(RequestLanes { ordinary, control }) = &self.lanes else { + // Keep requests ordered when concurrent dispatch is not enabled. + return task.await; + }; + // Finish the handshake before concurrent requests can observe session state. + if method == INITIALIZE_METHOD || !self.initialized { + return task.await; + } + + // Reserve capacity for health checks and cleanup while ordinary requests are blocked. + let admission = if matches!( + method, + ENVIRONMENT_INFO_METHOD + | ENVIRONMENT_STATUS_METHOD + | EXEC_SIGNAL_METHOD + | EXEC_TERMINATE_METHOD + | FS_CLOSE_METHOD + ) { + Arc::clone(control) + } else { + Arc::clone(ordinary) + }; + + // TODO(anp) bound queued request bytes without blocking later responses or cleanup. + self.tasks.spawn(async move { + let Ok(_permit) = admission.acquire_owned().await else { + return RequestTaskResult::ConnectionClosed; + }; + task.await + }); + RequestTaskResult::Completed + } + + pub(super) async fn shutdown(mut self) { + self.tasks.abort_all(); + while self.tasks.join_next().await.is_some() {} + } +} + +/// Per-connection request dispatch policy for local and remote exec-servers. +#[derive(Clone, Copy, Debug)] +pub enum RequestDispatchMode { + Inline, + Concurrent { + max_concurrent_requests: ConcurrentRequestLimit, + }, +} + +/// A valid request concurrency limit accepted by Tokio's semaphore. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct ConcurrentRequestLimit(usize); + +impl ConcurrentRequestLimit { + /// Returns a limit when it enables concurrency and fits Tokio's semaphore. + pub fn new(max_concurrent_requests: usize) -> Option { + if !(2..=Semaphore::MAX_PERMITS).contains(&max_concurrent_requests) { + return None; + } + + Some(Self(max_concurrent_requests)) + } + + /// Returns the validated number of concurrent requests. + pub fn get(self) -> usize { + self.0 + } +} + +impl FromStr for RequestDispatchMode { + type Err = ParseIntError; + + fn from_str(value: &str) -> Result { + let max_concurrent_requests = value.parse::()?.get(); + if max_concurrent_requests == 1 { + Ok(Self::Inline) + } else { + Ok(Self::Concurrent { + max_concurrent_requests: ConcurrentRequestLimit( + max_concurrent_requests.min(Semaphore::MAX_PERMITS), + ), + }) + } + } +} + +struct RequestLanes { + ordinary: Arc, + control: Arc, +} + +#[derive(Eq, PartialEq)] +pub(super) enum RequestTaskResult { + Completed, + ConnectionClosed, +} + +fn request_result(message: &Option) -> &'static str { + match message { + Some(RpcServerOutboundMessage::Error { .. }) => "error", + Some( + RpcServerOutboundMessage::Request(_) + | RpcServerOutboundMessage::Response { .. } + | RpcServerOutboundMessage::Notification(_), + ) + | None => "success", + } +} + +#[cfg(test)] +#[path = "request_dispatcher_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/server/request_dispatcher_tests.rs b/codex-rs/exec-server/src/server/request_dispatcher_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..ead6dd07c81498e52344bfa76d8773c7415f89e4 --- /dev/null +++ b/codex-rs/exec-server/src/server/request_dispatcher_tests.rs @@ -0,0 +1,566 @@ +use std::collections::BTreeMap; +use std::sync::Arc; +use std::time::Duration; +use std::time::Instant; + +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::RequestId; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_otel::MetricsClient; +use codex_otel::MetricsConfig; +use opentelemetry::trace::SpanId; +use opentelemetry::trace::TraceId; +use opentelemetry::trace::TracerProvider as _; +use opentelemetry_sdk::metrics::InMemoryMetricExporter; +use opentelemetry_sdk::metrics::data::AggregatedMetrics; +use opentelemetry_sdk::metrics::data::MetricData; +use opentelemetry_sdk::trace::InMemorySpanExporter; +use opentelemetry_sdk::trace::SdkTracerProvider; +use pretty_assertions::assert_eq; +use tokio::sync::Notify; +use tokio::sync::Semaphore; +use tokio::sync::mpsc; +use tokio::sync::watch; +use tokio::time::timeout; +use tracing_subscriber::filter::filter_fn; +use tracing_subscriber::prelude::*; + +use super::ConcurrentRequestLimit; +use super::RequestDispatchMode; +use super::RequestDispatcher; +use super::RequestTaskResult; +use crate::ExecServerRuntimePaths; +use crate::connection::JsonRpcConnectionEvent; +use crate::rpc::RpcNotificationSender; +use crate::rpc::RpcRouter; +use crate::rpc::RpcServerOutboundMessage; +use crate::rpc::invalid_request; +use crate::server::ExecServerHandler; +use crate::server::session_registry::SessionRegistry; +use crate::telemetry::ExecServerTelemetry; + +/// Public limits reject values that cannot safely enable semaphore-backed concurrency. +#[test] +fn concurrent_request_limit_rejects_invalid_values() { + assert_eq!( + ConcurrentRequestLimit::new(/*max_concurrent_requests*/ 0), + None + ); + assert_eq!( + ConcurrentRequestLimit::new(/*max_concurrent_requests*/ 1), + None + ); + assert_eq!( + ConcurrentRequestLimit::new(Semaphore::MAX_PERMITS.saturating_add(1)), + None + ); + assert_eq!( + ConcurrentRequestLimit::new(/*max_concurrent_requests*/ 2).map(ConcurrentRequestLimit::get), + Some(2) + ); +} + +/// CLI parsing keeps one request inline and bounds larger positive concurrency limits. +#[test] +fn request_dispatch_mode_parses_bounded_concurrency() { + assert!(matches!("1".parse(), Ok(RequestDispatchMode::Inline))); + assert!("0".parse::().is_err()); + + let oversized_limit = Semaphore::MAX_PERMITS.saturating_add(1).to_string(); + let mode = oversized_limit + .parse::() + .expect("parse oversized concurrent request limit"); + let RequestDispatchMode::Concurrent { + max_concurrent_requests, + } = mode + else { + panic!("expected concurrent request dispatch"); + }; + assert_eq!(max_concurrent_requests.get(), Semaphore::MAX_PERMITS); +} + +/// End-to-end request spans retain the wire method and inbound trace with bounded names. +#[test] +fn request_span_uses_bounded_name_wire_method_and_inbound_trace_parent() { + let span_exporter = InMemorySpanExporter::default(); + let tracer_provider = SdkTracerProvider::builder() + .with_simple_exporter(span_exporter.clone()) + .build(); + let tracer = tracer_provider.tracer("exec-server-test"); + let subscriber = tracing_subscriber::registry().with( + tracing_opentelemetry::layer() + .with_tracer(tracer) + .with_filter(filter_fn(codex_otel::OtelProvider::trace_export_filter)), + ); + let trace_id = TraceId::from_hex("00000000000000000000000000000001").expect("trace id"); + let parent_span_id = SpanId::from_hex("0000000000000002").expect("span id"); + let trace = codex_protocol::protocol::W3cTraceContext { + traceparent: Some(format!("00-{trace_id}-{parent_span_id}-01")), + tracestate: None, + }; + + let method = "custom/method"; + tracing::subscriber::with_default(subscriber, || { + tracing::callsite::rebuild_interest_cache(); + let request = JSONRPCRequest { + id: RequestId::Integer(1), + method: method.to_string(), + params: None, + trace: Some(trace), + }; + let JsonRpcConnectionEvent::QueuedRequest { request_span, .. } = + JsonRpcConnectionEvent::message(JSONRPCMessage::Request(request)) + else { + panic!("requests should start a server span before dispatch"); + }; + assert!( + span_exporter + .get_finished_spans() + .expect("request span export") + .is_empty(), + "the request span must remain open while the request is waiting" + ); + request_span.record("otel.name", "unknown"); + request_span.in_scope(|| {}); + drop(request_span); + }); + + tracer_provider.force_flush().expect("flush traces"); + let spans = span_exporter.get_finished_spans().expect("span export"); + assert_eq!(spans.len(), 1); + let request_span = spans + .iter() + .find(|span| span.name.as_ref() == "unknown") + .expect("unknown method span"); + assert_eq!( + request_span + .attributes + .iter() + .find(|attribute| attribute.key.as_str() == "method") + .map(|attribute| attribute.value.clone()), + Some(opentelemetry::Value::String(method.into())) + ); + assert_eq!(request_span.span_context.trace_id(), trace_id); + assert_eq!(request_span.parent_span_id, parent_span_id); +} + +/// Total request timing includes queueing without changing dispatch or queue timing. +#[tokio::test] +async fn request_queue_waits_for_dispatcher_admission_before_recording_telemetry() { + let span_exporter = InMemorySpanExporter::default(); + let tracer_provider = SdkTracerProvider::builder() + .with_simple_exporter(span_exporter.clone()) + .build(); + let subscriber = tracing_subscriber::registry().with( + tracing_opentelemetry::layer() + .with_tracer(tracer_provider.tracer("exec-server-test")) + .with_filter(filter_fn(codex_otel::OtelProvider::trace_export_filter)), + ); + let _subscriber = tracing::subscriber::set_default(subscriber); + tracing::callsite::rebuild_interest_cache(); + + let metrics = MetricsClient::new( + MetricsConfig::in_memory( + "test", + "codex-exec-server", + env!("CARGO_PKG_VERSION"), + InMemoryMetricExporter::default(), + ) + .with_runtime_reader(), + ) + .expect("metrics client"); + let telemetry = ExecServerTelemetry::new(metrics.clone()); + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(/*buffer*/ 1); + let notifications = RpcNotificationSender::new(outgoing_tx.clone()); + let requests = notifications.request_sender(); + let handler = Arc::new(ExecServerHandler::new( + SessionRegistry::new(telemetry.clone()), + notifications, + ExecServerRuntimePaths::new( + std::env::current_exe().expect("current executable"), + /*codex_linux_sandbox_exe*/ None, + ) + .expect("runtime paths"), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )); + let mut router = RpcRouter::new(); + let execution_started = Arc::new(Notify::new()); + let release_execution = Arc::new(Notify::new()); + let notify_execution_started = Arc::clone(&execution_started); + let wait_for_execution_release = Arc::clone(&release_execution); + let route_setup_duration = Duration::from_millis(200); + router.request( + "test/queued", + move |_handler: Arc, _params: ()| { + let execution_started = Arc::clone(¬ify_execution_started); + let release_execution = Arc::clone(&wait_for_execution_release); + std::thread::sleep(route_setup_duration); + async move { + execution_started.notify_one(); + release_execution.notified().await; + Ok::<_, codex_exec_server_protocol::JSONRPCErrorError>(()) + } + }, + ); + let (_disconnected_tx, disconnected_rx) = watch::channel(/*init*/ false); + let mut dispatcher = RequestDispatcher::new( + Arc::new(router), + handler, + outgoing_tx, + disconnected_rx, + requests, + telemetry, + RequestDispatchMode::Concurrent { + max_concurrent_requests: ConcurrentRequestLimit::new( + /*max_concurrent_requests*/ 2, + ) + .expect("valid request limit"), + }, + ); + dispatcher.initialized = true; + let admission = Arc::clone( + &dispatcher + .lanes + .as_ref() + .expect("concurrent request lanes") + .ordinary, + ); + let occupied_permits = admission + .acquire_many_owned(/*n*/ 2) + .await + .expect("occupy the request admission lane"); + let JsonRpcConnectionEvent::QueuedRequest { + request, + request_span, + queued_at, + } = JsonRpcConnectionEvent::message(JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(1), + method: "test/queued".to_string(), + params: None, + trace: None, + })) + else { + panic!("requests should start a server span before dispatch"); + }; + + assert!(matches!( + dispatcher + .dispatch_request(request, request_span, queued_at) + .await, + RequestTaskResult::Completed + )); + tokio::task::yield_now().await; + tokio::time::sleep(Duration::from_millis(25)).await; + + assert!( + span_exporter + .get_finished_spans() + .expect("request span export") + .is_empty(), + "the end-to-end request span must remain open before admission" + ); + let queued_snapshot = metrics.snapshot().expect("queued metrics snapshot"); + assert!( + !queued_snapshot + .scope_metrics() + .flat_map(opentelemetry_sdk::metrics::data::ScopeMetrics::metrics) + .any(|metric| metric.name() == "exec_server_request_queue_duration_seconds"), + "queue latency must not be recorded before request admission" + ); + + drop(occupied_permits); + timeout(Duration::from_secs(1), execution_started.notified()) + .await + .expect("queued request should execute after admission"); + assert!( + span_exporter + .get_finished_spans() + .expect("executing span export") + .is_empty(), + "the same request span must remain open during execution" + ); + + let snapshot = metrics.snapshot().expect("metrics snapshot"); + let queue_metric = snapshot + .scope_metrics() + .flat_map(opentelemetry_sdk::metrics::data::ScopeMetrics::metrics) + .find(|metric| metric.name() == "exec_server_request_queue_duration_seconds") + .expect("request queue duration metric"); + let AggregatedMetrics::F64(MetricData::Histogram(histogram)) = queue_metric.data() else { + panic!("request queue duration should be an f64 histogram"); + }; + let data_point = histogram + .data_points() + .next() + .expect("request queue duration data point"); + let queue_duration = Duration::from_secs_f64(data_point.sum()); + + assert_eq!(data_point.count(), 1); + assert!( + queue_duration >= Duration::from_millis(25), + "queue latency should include the occupied admission lane" + ); + assert_eq!( + data_point + .attributes() + .find(|attribute| attribute.key.as_str() == "method") + .map(|attribute| attribute.value.as_str().into_owned()), + Some("test/queued".to_string()) + ); + + release_execution.notify_one(); + let response = timeout(Duration::from_secs(1), outgoing_rx.recv()) + .await + .expect("queued request should send its response") + .expect("queued request response"); + assert!(matches!( + response, + RpcServerOutboundMessage::Response { + request_id: RequestId::Integer(1), + .. + } + )); + assert!(matches!( + dispatcher.join_next().await, + RequestTaskResult::Completed + )); + + tracer_provider.force_flush().expect("flush traces"); + let spans = span_exporter.get_finished_spans().expect("span export"); + assert_eq!(spans.len(), 1); + let request_span = spans.first().expect("end-to-end request span"); + assert_eq!(request_span.name.as_ref(), "test/queued"); + let request_duration = request_span + .end_time + .duration_since(request_span.start_time) + .expect("request span should have a valid interval"); + assert!( + request_duration >= Duration::from_millis(25), + "the end-to-end request span should include the admission wait" + ); + assert!( + queue_duration + <= request_duration.saturating_sub(route_setup_duration) + Duration::from_millis(5), + "queue latency must exclude synchronous request decoding and route setup" + ); + + let [dispatch_duration, total_duration] = + assert_request_completion(&metrics, "test/queued", "success"); + assert!(dispatch_duration >= route_setup_duration + Duration::from_millis(25)); + assert!(total_duration >= dispatch_duration); + assert!( + total_duration <= request_duration + Duration::from_millis(5), + "total request timing must not add admission wait twice" + ); + assert!( + total_duration >= queue_duration + route_setup_duration - Duration::from_millis(5), + "total request timing must include admission and route setup exactly once" + ); + + metrics.shutdown().expect("shutdown metrics"); +} + +struct RequestTelemetryFixture { + metrics: MetricsClient, + dispatcher: RequestDispatcher, + outgoing_rx: mpsc::Receiver, + disconnected_tx: watch::Sender, +} + +fn request_telemetry_fixture(router: RpcRouter) -> RequestTelemetryFixture { + let metrics = MetricsClient::new( + MetricsConfig::in_memory( + "test", + "codex-exec-server", + env!("CARGO_PKG_VERSION"), + InMemoryMetricExporter::default(), + ) + .with_runtime_reader(), + ) + .expect("metrics client"); + let telemetry = ExecServerTelemetry::new(metrics.clone()); + let (outgoing_tx, outgoing_rx) = mpsc::channel(/*buffer*/ 1); + let notifications = RpcNotificationSender::new(outgoing_tx.clone()); + let requests = notifications.request_sender(); + let handler = Arc::new(ExecServerHandler::new( + SessionRegistry::new(telemetry.clone()), + notifications, + ExecServerRuntimePaths::new( + std::env::current_exe().expect("current executable"), + /*codex_linux_sandbox_exe*/ None, + ) + .expect("runtime paths"), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )); + let (disconnected_tx, disconnected_rx) = watch::channel(/*init*/ false); + let dispatcher = RequestDispatcher::new( + Arc::new(router), + handler, + outgoing_tx, + disconnected_rx, + requests, + telemetry, + RequestDispatchMode::Inline, + ); + RequestTelemetryFixture { + metrics, + dispatcher, + outgoing_rx, + disconnected_tx, + } +} + +fn assert_request_completion(metrics: &MetricsClient, method: &str, result: &str) -> [Duration; 2] { + let snapshot = metrics.snapshot().expect("completed request metrics"); + let metric = snapshot + .scope_metrics() + .flat_map(opentelemetry_sdk::metrics::data::ScopeMetrics::metrics) + .find(|metric| metric.name() == "exec_server_requests_total") + .expect("request counter"); + let AggregatedMetrics::U64(MetricData::Sum(sum)) = metric.data() else { + panic!("request counter should be a u64 sum"); + }; + let points: Vec<_> = sum.data_points().collect(); + assert_eq!(points.len(), 1); + assert_eq!(points[0].value(), 1, "record the completion only once"); + + [ + "exec_server_request_duration_seconds", + "exec_server_request_total_duration_seconds", + ] + .map(|name| { + let metric = snapshot + .scope_metrics() + .flat_map(opentelemetry_sdk::metrics::data::ScopeMetrics::metrics) + .find(|metric| metric.name() == name) + .expect("request duration histogram"); + let AggregatedMetrics::F64(MetricData::Histogram(histogram)) = metric.data() else { + panic!("request duration should be an f64 histogram"); + }; + let points: Vec<_> = histogram.data_points().collect(); + assert_eq!(points.len(), 1); + let point = points[0]; + let attributes: BTreeMap<_, _> = point + .attributes() + .map(|attribute| { + ( + attribute.key.as_str().to_string(), + attribute.value.as_str().into_owned(), + ) + }) + .collect(); + assert_eq!( + attributes, + BTreeMap::from([ + ("method".to_string(), method.to_string()), + ("result".to_string(), result.to_string()), + ]) + ); + assert_eq!(point.count(), 1, "record each duration only once"); + Duration::from_secs_f64(point.sum()) + }) +} + +/// A synthetic receipt timestamp makes the pre-dispatch boundary deterministic for every result. +#[tokio::test] +async fn total_duration_preserves_dispatch_duration_and_completion_results() { + for (method, expected_method, expected_result, close_response) in [ + ("test/success", "test/success", "success", false), + ("test/error", "test/error", "error", false), + ("test/unknown", "unknown", "error", false), + ("test/success", "test/success", "disconnected", true), + ("test/unknown", "unknown", "disconnected", true), + ] { + let mut router = RpcRouter::new(); + router.request( + "test/success", + |_handler: Arc, _params: ()| async { + Ok::<_, codex_exec_server_protocol::JSONRPCErrorError>(()) + }, + ); + router.request( + "test/error", + |_handler: Arc, _params: ()| async { + Err::<(), _>(invalid_request("synthetic route error".to_string())) + }, + ); + let mut fixture = request_telemetry_fixture(router); + if close_response { + fixture.outgoing_rx.close(); + } + let pre_dispatch_wait = Duration::from_secs(5); + let received_at = Instant::now() - pre_dispatch_wait; + let result = fixture + .dispatcher + .dispatch_request( + JSONRPCRequest { + id: RequestId::Integer(1), + method: method.to_string(), + params: None, + trace: None, + }, + tracing::Span::none(), + received_at, + ) + .await; + assert_eq!( + matches!(result, RequestTaskResult::ConnectionClosed), + close_response + ); + let [dispatch_duration, total_duration] = + assert_request_completion(&fixture.metrics, expected_method, expected_result); + assert!(total_duration >= dispatch_duration + pre_dispatch_wait); + fixture.metrics.shutdown().expect("shutdown metrics"); + } +} + +/// Disconnecting a running request emits the same result and count for both duration definitions. +#[tokio::test] +async fn total_duration_records_disconnection_during_execution() { + let execution_started = Arc::new(Notify::new()); + let notify_execution_started = Arc::clone(&execution_started); + let mut router = RpcRouter::new(); + router.request( + "test/pending", + move |_handler: Arc, _params: ()| { + let execution_started = Arc::clone(¬ify_execution_started); + async move { + execution_started.notify_one(); + std::future::pending::>() + .await + } + }, + ); + let mut fixture = request_telemetry_fixture(router); + let pre_dispatch_wait = Duration::from_secs(5); + let received_at = Instant::now() - pre_dispatch_wait; + let dispatch = fixture.dispatcher.dispatch_request( + JSONRPCRequest { + id: RequestId::Integer(1), + method: "test/pending".to_string(), + params: None, + trace: None, + }, + tracing::Span::none(), + received_at, + ); + let disconnect = async { + execution_started.notified().await; + fixture + .disconnected_tx + .send(/*value*/ true) + .expect("disconnect request"); + }; + let (result, ()) = timeout(Duration::from_secs(1), async { + tokio::join!(dispatch, disconnect) + }) + .await + .expect("disconnected request should finish"); + assert!(matches!(result, RequestTaskResult::ConnectionClosed)); + let [dispatch_duration, total_duration] = + assert_request_completion(&fixture.metrics, "test/pending", "disconnected"); + assert!(total_duration >= dispatch_duration + pre_dispatch_wait); + fixture.metrics.shutdown().expect("shutdown metrics"); +} diff --git a/codex-rs/exec-server/src/server/session_registry.rs b/codex-rs/exec-server/src/server/session_registry.rs new file mode 100644 index 0000000000000000000000000000000000000000..89a861ab7be7a7f19708cb9ab376bcfdd18960fc --- /dev/null +++ b/codex-rs/exec-server/src/server/session_registry.rs @@ -0,0 +1,270 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCErrorError; +use tokio::sync::Mutex; +use uuid::Uuid; + +use crate::ExecServerRuntimePaths; +use crate::rpc::RpcNotificationSender; +use crate::rpc::invalid_request; +use crate::rpc::session_already_attached; +use crate::server::process_handler::ProcessHandler; +use crate::telemetry::ExecServerTelemetry; + +#[cfg(test)] +const DETACHED_SESSION_TTL: Duration = Duration::from_millis(200); +#[cfg(not(test))] +const DETACHED_SESSION_TTL: Duration = Duration::from_secs(30); + +pub(crate) struct SessionRegistry { + sessions: Mutex>>, + telemetry: ExecServerTelemetry, +} + +struct SessionEntry { + session_id: String, + process: ProcessHandler, + attachment: StdMutex, +} + +struct AttachmentState { + current_connection_id: Option, + detached_connection_id: Option, + detached_expires_at: Option, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +struct ConnectionId(Uuid); + +impl std::fmt::Display for ConnectionId { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + self.0.fmt(f) + } +} + +#[derive(Clone)] +pub(crate) struct SessionHandle { + registry: Arc, + entry: Arc, + connection_id: ConnectionId, +} + +impl SessionRegistry { + pub(crate) fn new(telemetry: ExecServerTelemetry) -> Arc { + Arc::new(Self { + sessions: Mutex::new(HashMap::new()), + telemetry, + }) + } + + pub(crate) async fn attach( + self: &Arc, + resume_session_id: Option, + notifications: RpcNotificationSender, + runtime_paths: ExecServerRuntimePaths, + ) -> Result { + enum AttachOutcome { + Attached(Arc), + Expired { + session_id: String, + entry: Arc, + }, + } + + let connection_id = ConnectionId(Uuid::new_v4()); + let outcome = { + let mut sessions = self.sessions.lock().await; + if let Some(session_id) = resume_session_id { + let entry = sessions + .get(&session_id) + .cloned() + .ok_or_else(|| invalid_request(format!("unknown session id {session_id}")))?; + if entry.is_expired(tokio::time::Instant::now()) { + let entry = sessions.remove(&session_id).ok_or_else(|| { + invalid_request(format!("unknown session id {session_id}")) + })?; + Ok(AttachOutcome::Expired { session_id, entry }) + } else if entry.has_active_connection() { + Err(session_already_attached(format!( + "session {session_id} is already attached to another connection" + ))) + } else { + entry.process.set_notification_sender(Some(notifications)); + entry.attach(connection_id); + Ok(AttachOutcome::Attached(entry)) + } + } else { + let session_id = Uuid::new_v4().to_string(); + let entry = Arc::new(SessionEntry::new( + session_id.clone(), + ProcessHandler::new(notifications, self.telemetry.clone(), runtime_paths), + connection_id, + )); + sessions.insert(session_id, Arc::clone(&entry)); + Ok(AttachOutcome::Attached(entry)) + } + }; + let entry = match outcome? { + AttachOutcome::Attached(entry) => entry, + AttachOutcome::Expired { session_id, entry } => { + entry.process.shutdown().await; + return Err(invalid_request(format!("unknown session id {session_id}"))); + } + }; + + Ok(SessionHandle { + registry: Arc::clone(self), + entry, + connection_id, + }) + } + + pub(crate) async fn shutdown(&self) { + let sessions = std::mem::take(&mut *self.sessions.lock().await); + for entry in sessions.into_values() { + entry.process.shutdown().await; + } + } + + async fn expire_if_detached(&self, session_id: String, connection_id: ConnectionId) { + tokio::time::sleep(DETACHED_SESSION_TTL).await; + + let removed = { + let mut sessions = self.sessions.lock().await; + let Some(entry) = sessions.get(&session_id) else { + return; + }; + if !entry.is_detached_connection_expired(connection_id, tokio::time::Instant::now()) { + return; + } + sessions.remove(&session_id) + }; + + if let Some(entry) = removed { + entry.process.shutdown().await; + } + } +} + +impl Default for SessionRegistry { + fn default() -> Self { + Self { + sessions: Mutex::new(HashMap::new()), + telemetry: ExecServerTelemetry::default(), + } + } +} + +impl SessionEntry { + fn new(session_id: String, process: ProcessHandler, connection_id: ConnectionId) -> Self { + Self { + session_id, + process, + attachment: StdMutex::new(AttachmentState { + current_connection_id: Some(connection_id), + detached_connection_id: None, + detached_expires_at: None, + }), + } + } + + fn attach(&self, connection_id: ConnectionId) { + let mut attachment = self + .attachment + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + attachment.current_connection_id = Some(connection_id); + attachment.detached_connection_id = None; + attachment.detached_expires_at = None; + } + + fn detach(&self, connection_id: ConnectionId) -> bool { + let mut attachment = self + .attachment + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if attachment.current_connection_id != Some(connection_id) { + return false; + } + + self.process.set_notification_sender(/*notifications*/ None); + attachment.current_connection_id = None; + attachment.detached_connection_id = Some(connection_id); + attachment.detached_expires_at = Some(tokio::time::Instant::now() + DETACHED_SESSION_TTL); + true + } + + fn has_active_connection(&self) -> bool { + self.attachment + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .current_connection_id + .is_some() + } + + fn is_attached_to(&self, connection_id: ConnectionId) -> bool { + self.attachment + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .current_connection_id + == Some(connection_id) + } + + fn is_expired(&self, now: tokio::time::Instant) -> bool { + self.attachment + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .detached_expires_at + .is_some_and(|deadline| now >= deadline) + } + + fn is_detached_connection_expired( + &self, + connection_id: ConnectionId, + now: tokio::time::Instant, + ) -> bool { + let attachment = self + .attachment + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + attachment.current_connection_id.is_none() + && attachment.detached_connection_id == Some(connection_id) + && attachment + .detached_expires_at + .is_some_and(|deadline| now >= deadline) + } +} + +impl SessionHandle { + pub(crate) fn session_id(&self) -> &str { + &self.entry.session_id + } + + pub(crate) fn connection_id(&self) -> String { + self.connection_id.to_string() + } + + pub(crate) fn is_session_attached(&self) -> bool { + self.entry.is_attached_to(self.connection_id) + } + + pub(crate) fn process(&self) -> &ProcessHandler { + &self.entry.process + } + + pub(crate) async fn detach(&self) { + if !self.entry.detach(self.connection_id) { + return; + } + + let registry = Arc::clone(&self.registry); + let session_id = self.entry.session_id.clone(); + let connection_id = self.connection_id; + tokio::spawn(async move { + registry.expire_if_detached(session_id, connection_id).await; + }); + } +} diff --git a/codex-rs/exec-server/src/server/transport.rs b/codex-rs/exec-server/src/server/transport.rs new file mode 100644 index 0000000000000000000000000000000000000000..9d556b0dba63d77a70eaec7b0b0f0bf6eb1117a4 --- /dev/null +++ b/codex-rs/exec-server/src/server/transport.rs @@ -0,0 +1,240 @@ +use axum::Router; +use axum::body::Body; +use axum::extract::ConnectInfo; +use axum::extract::State; +use axum::extract::ws::WebSocketUpgrade; +use axum::http::Request; +use axum::http::StatusCode; +use axum::http::header::ORIGIN; +use axum::middleware; +use axum::middleware::Next; +use axum::response::IntoResponse; +use axum::response::Response; +use axum::routing::any; +use axum::routing::get; +use codex_http_client::HttpClientFactory; +use std::io::Write as _; +use std::net::SocketAddr; +use tokio::io; +use tokio::io::AsyncRead; +use tokio::io::AsyncWrite; +use tokio::net::TcpListener; +use tracing::info; +use tracing::warn; + +use crate::ExecServerRuntimePaths; +use crate::ExecServerTelemetry; +use crate::connection::JsonRpcConnection; +use crate::server::RequestDispatchMode; +use crate::server::processor::ConnectionProcessor; +use crate::telemetry::ConnectionTransport; + +pub const DEFAULT_LISTEN_URL: &str = "ws://127.0.0.1:0"; + +#[derive(Debug, Clone, Eq, PartialEq)] +pub(crate) enum ExecServerListenTransport { + WebSocket(SocketAddr), + Stdio, +} + +#[derive(Debug, Clone, Eq, PartialEq)] +pub enum ExecServerListenUrlParseError { + UnsupportedListenUrl(String), + InvalidWebSocketListenUrl(String), +} + +impl std::fmt::Display for ExecServerListenUrlParseError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + ExecServerListenUrlParseError::UnsupportedListenUrl(listen_url) => write!( + f, + "unsupported --listen URL `{listen_url}`; expected `ws://IP:PORT` or `stdio`" + ), + ExecServerListenUrlParseError::InvalidWebSocketListenUrl(listen_url) => write!( + f, + "invalid websocket --listen URL `{listen_url}`; expected `ws://IP:PORT`" + ), + } + } +} + +impl std::error::Error for ExecServerListenUrlParseError {} + +pub(crate) fn parse_listen_url( + listen_url: &str, +) -> Result { + if matches!(listen_url, "stdio" | "stdio://") { + return Ok(ExecServerListenTransport::Stdio); + } + + if let Some(socket_addr) = listen_url.strip_prefix("ws://") { + return socket_addr + .parse::() + .map(ExecServerListenTransport::WebSocket) + .map_err(|_| { + ExecServerListenUrlParseError::InvalidWebSocketListenUrl(listen_url.to_string()) + }); + } + + Err(ExecServerListenUrlParseError::UnsupportedListenUrl( + listen_url.to_string(), + )) +} + +pub(crate) async fn run_transport( + listen_url: &str, + runtime_paths: ExecServerRuntimePaths, + telemetry: ExecServerTelemetry, + http_client_factory: HttpClientFactory, + request_dispatch_mode: RequestDispatchMode, +) -> Result<(), Box> { + match parse_listen_url(listen_url)? { + ExecServerListenTransport::WebSocket(bind_address) => { + run_websocket_listener( + bind_address, + runtime_paths, + telemetry, + http_client_factory, + request_dispatch_mode, + ) + .await + } + ExecServerListenTransport::Stdio => { + run_stdio_connection( + runtime_paths, + telemetry, + http_client_factory, + request_dispatch_mode, + ) + .await + } + } +} + +async fn run_stdio_connection( + runtime_paths: ExecServerRuntimePaths, + telemetry: ExecServerTelemetry, + http_client_factory: HttpClientFactory, + request_dispatch_mode: RequestDispatchMode, +) -> Result<(), Box> { + run_stdio_connection_with_io( + io::stdin(), + io::stdout(), + runtime_paths, + telemetry, + http_client_factory, + request_dispatch_mode, + ) + .await +} + +async fn run_stdio_connection_with_io( + reader: R, + writer: W, + runtime_paths: ExecServerRuntimePaths, + telemetry: ExecServerTelemetry, + http_client_factory: HttpClientFactory, + request_dispatch_mode: RequestDispatchMode, +) -> Result<(), Box> +where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, +{ + let processor = ConnectionProcessor::new_with_telemetry( + runtime_paths, + telemetry, + http_client_factory, + request_dispatch_mode, + ); + tracing::info!("codex-exec-server listening on stdio"); + processor + .run_connection( + JsonRpcConnection::from_stdio(reader, writer, "exec-server stdio".to_string()), + ConnectionTransport::Stdio, + ) + .await; + // Stdio serves exactly one connection, so detached sessions cannot be resumed. + processor.shutdown().await; + Ok(()) +} + +async fn run_websocket_listener( + bind_address: SocketAddr, + runtime_paths: ExecServerRuntimePaths, + telemetry: ExecServerTelemetry, + http_client_factory: HttpClientFactory, + request_dispatch_mode: RequestDispatchMode, +) -> Result<(), Box> { + let listener = TcpListener::bind(bind_address).await?; + let local_addr = listener.local_addr()?; + let processor = ConnectionProcessor::new_with_telemetry( + runtime_paths, + telemetry, + http_client_factory, + request_dispatch_mode, + ); + info!("codex-exec-server listening on ws://{local_addr}"); + println!("ws://{local_addr}"); + std::io::stdout().flush()?; + + let router = Router::new() + .route("/", any(websocket_upgrade_handler)) + .route("/readyz", get(readiness_handler)) + .layer(middleware::from_fn(reject_requests_with_origin_header)) + .with_state(ExecServerWebSocketState { processor }); + axum::serve( + listener, + router.into_make_service_with_connect_info::(), + ) + .await?; + Ok(()) +} + +#[derive(Clone)] +struct ExecServerWebSocketState { + processor: ConnectionProcessor, +} + +async fn readiness_handler() -> StatusCode { + StatusCode::OK +} + +async fn reject_requests_with_origin_header( + request: Request, + next: Next, +) -> Result { + if request.headers().contains_key(ORIGIN) { + warn!( + method = %request.method(), + uri = %request.uri(), + "rejecting exec-server websocket listener request with Origin header" + ); + Err(StatusCode::FORBIDDEN) + } else { + Ok(next.run(request).await) + } +} + +async fn websocket_upgrade_handler( + websocket: WebSocketUpgrade, + ConnectInfo(peer_addr): ConnectInfo, + State(state): State, +) -> impl IntoResponse { + info!(%peer_addr, "exec-server websocket client connected"); + websocket.on_upgrade(move |stream| async move { + state + .processor + .run_connection( + JsonRpcConnection::from_axum_websocket( + stream, + format!("exec-server websocket {peer_addr}"), + ), + ConnectionTransport::WebSocket, + ) + .await; + }) +} + +#[cfg(test)] +#[path = "transport_tests.rs"] +mod transport_tests; diff --git a/codex-rs/exec-server/src/server/transport_tests.rs b/codex-rs/exec-server/src/server/transport_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..c26109a3b8cc373d86070314f842d27b0c116e1b --- /dev/null +++ b/codex-rs/exec-server/src/server/transport_tests.rs @@ -0,0 +1,172 @@ +use std::net::SocketAddr; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCNotification; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::RequestId; +use pretty_assertions::assert_eq; +use tokio::io::AsyncBufReadExt; +use tokio::io::AsyncWriteExt; +use tokio::io::BufReader; +use tokio::io::duplex; +use tokio::time::timeout; + +use super::DEFAULT_LISTEN_URL; +use super::ExecServerListenTransport; +use super::parse_listen_url; +use super::run_stdio_connection_with_io; +use crate::ExecServerRuntimePaths; +use crate::protocol::INITIALIZE_METHOD; +use crate::protocol::INITIALIZED_METHOD; +use crate::protocol::InitializeParams; +use crate::protocol::InitializeResponse; +use crate::server::RequestDispatchMode; + +#[test] +fn parse_listen_url_accepts_default_websocket_url() { + let transport = parse_listen_url(DEFAULT_LISTEN_URL).expect("default listen URL should parse"); + assert_eq!( + transport, + ExecServerListenTransport::WebSocket( + "127.0.0.1:0" + .parse::() + .expect("valid socket address") + ) + ); +} + +#[test] +fn parse_listen_url_accepts_stdio() { + let transport = parse_listen_url("stdio").expect("stdio listen URL should parse"); + assert_eq!(transport, ExecServerListenTransport::Stdio); +} + +#[test] +fn parse_listen_url_accepts_stdio_url() { + let transport = parse_listen_url("stdio://").expect("stdio listen URL should parse"); + assert_eq!(transport, ExecServerListenTransport::Stdio); +} + +#[tokio::test] +async fn stdio_listen_transport_serves_initialize() { + let transport = parse_listen_url("stdio").expect("stdio listen URL should parse"); + let ExecServerListenTransport::Stdio = transport else { + panic!("expected stdio listen transport, got {transport:?}"); + }; + + let (mut client_writer, server_reader) = duplex(1 << 20); + let (server_writer, client_reader) = duplex(1 << 20); + let server_task = tokio::spawn(run_stdio_connection_with_io( + server_reader, + server_writer, + test_runtime_paths(), + crate::ExecServerTelemetry::default(), + codex_http_client::HttpClientFactory::new( + codex_http_client::OutboundProxyPolicy::ReqwestDefault, + ), + RequestDispatchMode::Inline, + )); + let mut client_lines = BufReader::new(client_reader).lines(); + + let initialize = JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(1), + method: INITIALIZE_METHOD.to_string(), + params: Some( + serde_json::to_value(InitializeParams { + client_name: "exec-server-transport-test".to_string(), + resume_session_id: None, + }) + .expect("initialize params should serialize"), + ), + trace: None, + }); + write_jsonrpc_line(&mut client_writer, &initialize).await; + + let response = timeout(Duration::from_secs(1), client_lines.next_line()) + .await + .expect("initialize response should arrive") + .expect("initialize response read should succeed") + .expect("initialize response should be present"); + let response: JSONRPCMessage = + serde_json::from_str(&response).expect("initialize response should parse"); + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = response else { + panic!("expected initialize response, got {response:?}"); + }; + assert_eq!(id, RequestId::Integer(1)); + let initialize_response: InitializeResponse = + serde_json::from_value(result).expect("initialize response should decode"); + assert!( + !initialize_response.session_id.is_empty(), + "initialize should return a session id" + ); + + let initialized = JSONRPCMessage::Notification(JSONRPCNotification { + method: INITIALIZED_METHOD.to_string(), + params: Some(serde_json::to_value(()).expect("initialized params should serialize")), + }); + write_jsonrpc_line(&mut client_writer, &initialized).await; + + drop(client_writer); + drop(client_lines); + timeout(Duration::from_secs(1), server_task) + .await + .expect("stdio transport should finish after client disconnect") + .expect("stdio transport task should join") + .expect("stdio transport should not fail"); +} + +#[test] +fn parse_listen_url_accepts_websocket_url() { + let transport = + parse_listen_url("ws://127.0.0.1:1234").expect("websocket listen URL should parse"); + assert_eq!( + transport, + ExecServerListenTransport::WebSocket( + "127.0.0.1:1234" + .parse::() + .expect("valid socket address") + ) + ); +} + +#[test] +fn parse_listen_url_rejects_invalid_websocket_url() { + let err = parse_listen_url("ws://localhost:1234") + .expect_err("hostname bind address should be rejected"); + assert_eq!( + err.to_string(), + "invalid websocket --listen URL `ws://localhost:1234`; expected `ws://IP:PORT`" + ); +} + +#[test] +fn parse_listen_url_rejects_unsupported_url() { + let err = + parse_listen_url("http://127.0.0.1:1234").expect_err("unsupported scheme should fail"); + assert_eq!( + err.to_string(), + "unsupported --listen URL `http://127.0.0.1:1234`; expected `ws://IP:PORT` or `stdio`" + ); +} + +async fn write_jsonrpc_line(writer: &mut tokio::io::DuplexStream, message: &JSONRPCMessage) { + let encoded = serde_json::to_vec(message).expect("JSON-RPC message should serialize"); + writer + .write_all(&encoded) + .await + .expect("JSON-RPC message should write"); + writer + .write_all(b"\n") + .await + .expect("JSON-RPC newline should write"); +} + +fn test_runtime_paths() -> ExecServerRuntimePaths { + ExecServerRuntimePaths::new( + std::env::current_exe().expect("current exe"), + /*codex_linux_sandbox_exe*/ None, + ) + .expect("runtime paths") +} diff --git a/codex-rs/exec-server/src/shell_snapshot.rs b/codex-rs/exec-server/src/shell_snapshot.rs new file mode 100644 index 0000000000000000000000000000000000000000..e848f6d82c8183248f21dadaa1d0709e86b2eece --- /dev/null +++ b/codex-rs/exec-server/src/shell_snapshot.rs @@ -0,0 +1,406 @@ +use std::collections::HashMap; +use std::collections::VecDeque; +use std::process::Stdio; +use std::sync::Arc; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_network_proxy::PROXY_ACTIVE_ENV_KEY; +use codex_network_proxy::strip_managed_proxy_env; +use codex_protocol::config_types::ShellEnvironmentPolicyInherit; +use codex_protocol::shell_environment; +use codex_shell_command::shell_detect::ShellType; +use codex_shell_command::shell_snapshot::CapturedSnapshot; +use codex_shell_command::shell_snapshot::SnapshotCaptureOptions; +use codex_shell_command::shell_snapshot::SnapshotStartup; +use codex_shell_command::shell_snapshot::snapshot_capture_script; +use codex_utils_path_uri::PathUri; +use tokio::io::AsyncReadExt; +use tokio::process::Command; +use tokio::sync::Mutex; +use tokio::sync::OnceCell; +use tokio::time::Instant; + +use crate::FileSystemSandboxContext; +use crate::local_process::shell_environment_policy; +use crate::process_sandbox::PreparedExecRequest; +use crate::protocol::ExecEnvPolicy; +use crate::protocol::ExecParams; +use crate::protocol::ShellSnapshotRequest; +use crate::rpc::internal_error; +use crate::rpc::invalid_params; +use crate::telemetry::ExecServerTelemetry; + +const MAX_CACHED_SNAPSHOTS: usize = 16; +const MAX_SNAPSHOT_BYTES: usize = 512 * 1024; +// Capture also includes quoted export records and an optional pre-startup environment. +const MAX_SNAPSHOT_CAPTURE_BYTES: usize = 8 * MAX_SNAPSHOT_BYTES; +const MAX_SNAPSHOT_ENV_VALUE_BYTES: usize = 60 * 1024; +const MAX_SNAPSHOT_SCOPE_BYTES: usize = 256; +const SNAPSHOT_TIMEOUT: Duration = Duration::from_secs(10); +const SNAPSHOT_RETRY_BACKOFF: Duration = Duration::from_secs(1); +const MAX_SNAPSHOT_ATTEMPTS: usize = 3; + +#[derive(Default)] +pub(crate) struct ShellSnapshotCache { + entries: Mutex>, +} + +struct CachedShellSnapshot { + request: ShellSnapshotRequest, + cwd: PathUri, + env_policy: Option, + sandbox: Option, + attempts: usize, + // Failed captures store the earliest time another attempt may start. + snapshot: Arc>>, +} + +struct ShellSnapshot { + state: String, + environment: HashMap, +} + +// Keep a bounded metric label alongside the original RPC error. +type CaptureResult = Result; + +#[derive(Clone, Copy, PartialEq, Eq)] +pub(crate) enum CapturePurpose { + Execution, + Prewarm, +} + +impl ShellSnapshotCache { + pub(crate) async fn prepare( + &self, + params: &ExecParams, + prepared: &mut PreparedExecRequest, + telemetry: &ExecServerTelemetry, + purpose: CapturePurpose, + ) -> Result<(), JSONRPCErrorError> { + let Some(request) = params.shell_snapshot.as_ref() else { + return Ok(()); + }; + if request.scope_id.is_empty() || request.scope_id.len() > MAX_SNAPSHOT_SCOPE_BYTES { + return Err(invalid_params(format!( + "shell snapshot scope must be non-empty and at most {MAX_SNAPSHOT_SCOPE_BYTES} bytes" + ))); + } + + if params.argv.len() < 3 + || params.argv[0] != request.shell.path + || params.argv[1] != "-lc" + || !prepared.command.ends_with(¶ms.argv) + { + return Ok(()); + } + + let shell_type = match request.shell.name.as_str() { + "bash" => ShellType::Bash, + "zsh" => ShellType::Zsh, + "sh" => ShellType::Sh, + name => { + return Err(invalid_params(format!( + "shell snapshots are unsupported for shell `{name}`" + ))); + } + }; + + let (snapshot, attempt) = { + let mut entries = self.entries.lock().await; + let position = entries.iter().position(|entry| { + &entry.request == request + && entry.cwd == params.cwd + && entry.env_policy == params.env_policy + && entry.sandbox == params.sandbox + }); + let cached = position.and_then(|position| { + let mut entry = entries.remove(position)?; + // Share each failed attempt during backoff. After the retry + // budget is exhausted, keep falling back until eviction. + if purpose == CapturePurpose::Execution + && entry.attempts < MAX_SNAPSHOT_ATTEMPTS + && let Some(Err(retry_at)) = entry.snapshot.get() + && Instant::now() >= *retry_at + { + entry.attempts += 1; + entry.snapshot = Arc::new(OnceCell::new()); + } + let snapshot = Arc::clone(&entry.snapshot); + let attempt = entry.attempts; + entries.push_back(entry); + Some((snapshot, attempt)) + }); + if let Some(snapshot) = cached { + snapshot + } else { + let snapshot = Arc::new(OnceCell::new()); + let entry = CachedShellSnapshot { + request: request.clone(), + cwd: params.cwd.clone(), + env_policy: params.env_policy.clone(), + sandbox: params.sandbox.clone(), + attempts: 1, + snapshot: Arc::clone(&snapshot), + }; + entries.push_back(entry); + if entries.len() > MAX_CACHED_SNAPSHOTS { + entries.pop_front(); + } + + (snapshot, 1) + } + }; + let capture = async { + let attempt = attempt.to_string(); + let purpose = match purpose { + CapturePurpose::Execution => "execution", + CapturePurpose::Prewarm => "prewarm", + }; + let started_at = std::time::Instant::now(); + let result = capture_snapshot(params, prepared, shell_type).await; + telemetry.shell_snapshot_captured( + started_at.elapsed(), + result.as_ref().map(|_| ()).map_err(|(reason, _)| *reason), + &[ + ("purpose", purpose), + ("attempt", &attempt), + ("shell", request.shell.name.as_str()), + ("sandbox", prepared.sandbox.as_metric_tag()), + ], + ); + result.map_err(|(_, error)| error) + }; + let snapshot = match purpose { + CapturePurpose::Execution => { + snapshot + .get_or_init(|| async { + capture.await.map_err(|err| { + tracing::warn!("failed to capture shell snapshot: {err:?}"); + Instant::now() + SNAPSHOT_RETRY_BACKOFF + }) + }) + .await + } + CapturePurpose::Prewarm => { + // Leave the cell uninitialized on failure: a waiting real command + // can capture immediately, without spending its retry budget. + snapshot + .get_or_try_init(|| async { capture.await.map(Ok) }) + .await? + } + }; + let Ok(snapshot) = snapshot else { + return Ok(()); + }; + if purpose == CapturePurpose::Prewarm { + return Ok(()); + } + + let request_overrides = params + .env + .iter() + .map(|(name, value)| { + ( + name.clone(), + prepared.env.get(name).unwrap_or(value).clone(), + ) + }) + .collect::>(); + prepared.env.extend( + snapshot + .environment + .iter() + .map(|(name, value)| (name.clone(), value.clone())), + ); + prepared.env.extend(request_overrides); + prepared + .env + .retain(|name, _| !shell_environment::is_non_inheritable_env_var(name)); + + let mut state = snapshot.state.as_str(); + let mut state_variables = Vec::new(); + while !state.is_empty() { + let mut end = state.len().min(MAX_SNAPSHOT_ENV_VALUE_BYTES); + while !state.is_char_boundary(end) { + end -= 1; + } + let (chunk, remaining) = state.split_at(end); + let name = format!("__CODEX_SHELL_SNAPSHOT_STATE_{}", state_variables.len()); + prepared.env.insert(name.clone(), chunk.to_string()); + state_variables.push(name); + state = remaining; + } + let state_expansion = state_variables + .iter() + .map(|name| format!("${{{name}}}")) + .collect::(); + let state_variables = state_variables.join(" "); + let shell_start = prepared.command.len() - params.argv.len(); + // Automatic startup files run before the restoration script and could + // reintroduce environment variables that the snapshot already filtered. + let (shell_flag, startup) = match shell_type { + ShellType::Bash => ("-pc", "set +o privileged\n"), + ShellType::Zsh => ("-fc", "setopt RCS\n"), + ShellType::Sh => ("-c", ""), + ShellType::PowerShell | ShellType::Cmd => unreachable!(), + }; + prepared.command[shell_start + 1] = shell_flag.to_string(); + prepared.command[shell_start + 2] = format!( + "{startup}if ! eval \"unset {state_variables}\n{state_expansion}\" >/dev/null; then printf 'failed to restore shell snapshot\\n' >&2; fi\n{}", + params.argv[2] + ); + + Ok(()) + } +} + +async fn capture_snapshot( + params: &ExecParams, + prepared: &PreparedExecRequest, + shell_type: ShellType, +) -> CaptureResult { + let script = snapshot_capture_script( + shell_type, + SnapshotCaptureOptions { + startup: SnapshotStartup::Interactive, + declarations: false, + environment: true, + }, + ) + .ok_or_else(|| { + ( + "unsupported_shell", + invalid_params("unsupported shell snapshot script".to_string()), + ) + })?; + let shell_start = prepared.command.len() - params.argv.len(); + let mut argv = prepared.command.clone(); + argv[shell_start + 2] = script; + let (program, args) = argv.split_first().ok_or_else(|| { + ( + "missing_command", + internal_error("missing shell snapshot command".to_string()), + ) + })?; + + let mut command = Command::new(program); + command + .args(args) + .current_dir(prepared.cwd.as_path()) + .env_clear() + .envs(&prepared.env) + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::null()) + .kill_on_drop(true); + if let Some(arg0) = &prepared.arg0 { + command.arg0(arg0); + } + let mut child = command.spawn().map_err(|err| { + ( + "spawn_failed", + internal_error(format!("cannot capture shell snapshot: {err}")), + ) + })?; + let stdout = child.stdout.take().ok_or_else(|| { + ( + "missing_output", + internal_error("missing shell snapshot output".to_string()), + ) + })?; + let capture = async { + let mut output = Vec::new(); + stdout + .take((MAX_SNAPSHOT_CAPTURE_BYTES + 1) as u64) + .read_to_end(&mut output) + .await + .map_err(|err| { + ( + "read_failed", + internal_error(format!("cannot read shell snapshot: {err}")), + ) + })?; + if output.len() > MAX_SNAPSHOT_CAPTURE_BYTES { + return Err(( + "too_large", + internal_error(format!( + "shell snapshot capture exceeds {MAX_SNAPSHOT_CAPTURE_BYTES} bytes" + )), + )); + } + let status = child.wait().await.map_err(|err| { + ( + "wait_failed", + internal_error(format!("cannot finish shell snapshot: {err}")), + ) + })?; + if !status.success() { + return Err(( + "nonzero_exit", + internal_error(format!("shell snapshot capture exited with {status}")), + )); + } + Ok(output) + }; + let output = tokio::time::timeout(SNAPSHOT_TIMEOUT, capture) + .await + .map_err(|_| { + ( + "timeout", + internal_error("shell snapshot capture timed out".to_string()), + ) + })??; + + parse_snapshot(shell_type, &output, params.env_policy.as_ref()) +} + +fn parse_snapshot( + shell_type: ShellType, + output: &[u8], + env_policy: Option<&ExecEnvPolicy>, +) -> CaptureResult { + let captured = CapturedSnapshot::parse(shell_type, output).ok_or_else(|| { + ( + "invalid_capture", + internal_error("invalid shell snapshot capture".to_string()), + ) + })?; + let state = captured.render_state(); + if state.len().saturating_add(captured.environment.len()) > MAX_SNAPSHOT_BYTES { + return Err(( + "too_large", + internal_error(format!("shell snapshot exceeds {MAX_SNAPSHOT_BYTES} bytes")), + )); + } + + let mut environment = captured + .environment + .split(|byte| *byte == 0) + .filter(|entry| !entry.is_empty()) + .filter_map(|entry| { + let (name, value) = std::str::from_utf8(entry).ok()?.split_once('=')?; + Some((name.to_string(), value.to_string())) + }) + .collect::>(); + if environment.contains_key(PROXY_ACTIVE_ENV_KEY) { + strip_managed_proxy_env(&mut environment); + } + let mut environment = match env_policy { + Some(policy) => { + let mut policy = shell_environment_policy(policy); + policy.inherit = ShellEnvironmentPolicyInherit::All; + shell_environment::create_env_from_vars(environment, &policy, /*thread_id*/ None) + } + None => environment, + }; + environment.remove("PWD"); + environment.remove("OLDPWD"); + environment.retain(|name, _| !shell_environment::is_non_inheritable_env_var(name)); + + Ok(ShellSnapshot { state, environment }) +} + +#[cfg(test)] +#[path = "shell_snapshot_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/shell_snapshot_tests.rs b/codex-rs/exec-server/src/shell_snapshot_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..cdab0f750857db4a5b6b0a7b6955b25db94ea84a --- /dev/null +++ b/codex-rs/exec-server/src/shell_snapshot_tests.rs @@ -0,0 +1,327 @@ +use std::collections::BTreeMap; +use std::collections::HashMap; + +use codex_otel::MetricsClient; +use codex_otel::MetricsConfig; +use codex_protocol::config_types::ShellEnvironmentPolicyInherit; +use opentelemetry_sdk::metrics::InMemoryMetricExporter; +use opentelemetry_sdk::metrics::data::AggregatedMetrics; +use opentelemetry_sdk::metrics::data::MetricData; +use pretty_assertions::assert_eq; +use test_case::test_case; + +use super::CapturePurpose; +use super::MAX_SNAPSHOT_ATTEMPTS; +use super::MAX_SNAPSHOT_BYTES; +use super::SNAPSHOT_RETRY_BACKOFF; +use super::ShellSnapshotCache; +use super::parse_snapshot; +use crate::process_sandbox::prepare_exec_request_with_telemetry; +use crate::process_telemetry::ProcessTelemetry; +use crate::protocol::ExecEnvPolicy; +use crate::protocol::ExecParams; +use crate::protocol::ProcessId; +use crate::protocol::ShellInfo; +use crate::protocol::ShellSnapshotRequest; +use crate::telemetry::ExecServerTelemetry; + +#[test_case(1, CapturePurpose::Execution; "succeeds_on_first_attempt")] +#[test_case(2, CapturePurpose::Execution; "recovers_on_second_attempt")] +#[test_case(3, CapturePurpose::Execution; "recovers_on_last_attempt")] +#[test_case(4, CapturePurpose::Execution; "stops_after_three_failures")] +#[test_case(2, CapturePurpose::Prewarm; "failed_prewarm_releases_waiting_command")] +#[test_case(3, CapturePurpose::Prewarm; "failed_prewarm_preserves_last_attempt")] +#[test_case(4, CapturePurpose::Prewarm; "failed_prewarm_preserves_retry_limit")] +#[tokio::test] +async fn snapshot_failure_retries_are_bounded_and_single_flight( + recovery_attempt: usize, + initial_purpose: CapturePurpose, +) -> anyhow::Result<()> { + let prewarm_fails_first = initial_purpose == CapturePurpose::Prewarm; + let home = tempfile::TempDir::new()?; + let profile = home.path().join(".bashrc"); + std::fs::write(&profile, "printf x >> \"$HOME/captures\"\nexit 7\n")?; + let params = ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("snapshot-retry"), + argv: vec![ + "/bin/bash".to_string(), + "-lc".to_string(), + "true".to_string(), + ], + cwd: codex_utils_path_uri::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: "attachment-1".to_string(), + shell: ShellInfo { + name: "bash".to_string(), + path: "/bin/bash".to_string(), + }, + }), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }; + let cache = ShellSnapshotCache::default(); + let metrics = MetricsClient::new( + MetricsConfig::in_memory( + "test", + "codex-exec-server", + env!("CARGO_PKG_VERSION"), + InMemoryMetricExporter::default(), + ) + .with_runtime_reader(), + )?; + let telemetry = ExecServerTelemetry::new(metrics.clone()); + + for attempt in 1..=5 { + if attempt == recovery_attempt { + std::fs::write( + &profile, + "printf x >> \"$HOME/captures\"\nprofile_helper() { printf recovered; }\n", + )?; + } + let mut prepared = prepare_exec_request_with_telemetry( + ¶ms, + params.env.clone(), + /*runtime_paths*/ None, + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + &ProcessTelemetry::default(), + ) + .await + .expect("prepare capture"); + let mut concurrent = prepare_exec_request_with_telemetry( + ¶ms, + params.env.clone(), + /*runtime_paths*/ None, + /*network_policy_decider*/ None, + /*network_policy_audit_observer*/ None, + &ProcessTelemetry::default(), + ) + .await + .expect("prepare concurrent capture"); + let prewarming = prewarm_fails_first && attempt == 1; + let purpose = if prewarming { + CapturePurpose::Prewarm + } else { + CapturePurpose::Execution + }; + let (first, second) = tokio::join!( + biased; + cache.prepare(¶ms, &mut prepared, &telemetry, purpose), + cache.prepare(¶ms, &mut concurrent, &telemetry, CapturePurpose::Execution), + ); + if prewarming { + first.expect_err("prewarm must report failure without caching it"); + } else { + first.expect("capture failure must preserve command fallback"); + } + second.expect("waiting command must complete even when prewarm fails"); + assert_eq!( + (&prepared.command, &prepared.env), + (&concurrent.command, &concurrent.env) + ); + + tokio::time::pause(); + if attempt < recovery_attempt || recovery_attempt > MAX_SNAPSHOT_ATTEMPTS { + cache + .prepare( + ¶ms, + &mut prepared, + &telemetry, + CapturePurpose::Execution, + ) + .await + .expect("capture must stay cached during backoff"); + assert_eq!( + (&prepared.command, &prepared.env), + (¶ms.argv, ¶ms.env) + ); + } else { + assert_ne!(prepared.command, params.argv); + } + assert_eq!( + std::fs::read_to_string(home.path().join("captures"))?, + "x".repeat( + attempt.min(recovery_attempt).min(MAX_SNAPSHOT_ATTEMPTS) + + usize::from(prewarm_fails_first) + ) + ); + tokio::time::advance(SNAPSHOT_RETRY_BACKOFF).await; + tokio::time::resume(); + } + + let snapshot = metrics.snapshot()?; + let mut counters = BTreeMap::new(); + let mut durations = BTreeMap::new(); + for metric in snapshot + .scope_metrics() + .flat_map(opentelemetry_sdk::metrics::data::ScopeMetrics::metrics) + { + match metric.name() { + "codex.shell_snapshot" => { + let AggregatedMetrics::U64(MetricData::Sum(sum)) = metric.data() else { + panic!("expected shell snapshot counter"); + }; + for point in sum.data_points() { + let tags = point + .attributes() + .map(|attribute| (attribute.key.to_string(), attribute.value.to_string())) + .collect::>(); + counters.insert(tags, point.value()); + } + } + "codex.shell_snapshot.duration_ms" => { + let AggregatedMetrics::F64(MetricData::Histogram(histogram)) = metric.data() else { + panic!("expected shell snapshot duration histogram"); + }; + for point in histogram.data_points() { + let tags = point + .attributes() + .map(|attribute| (attribute.key.to_string(), attribute.value.to_string())) + .collect::>(); + durations.insert(tags, point.count()); + } + } + _ => {} + } + } + let mut expected_counters = BTreeMap::new(); + let mut expected_durations = BTreeMap::new(); + let captures = prewarm_fails_first + .then_some(("prewarm", 1)) + .into_iter() + .chain( + (1..=recovery_attempt.min(MAX_SNAPSHOT_ATTEMPTS)).map(|attempt| ("execution", attempt)), + ); + for (purpose, attempt) in captures { + let success = purpose == "execution" && attempt == recovery_attempt; + let mut tags = BTreeMap::from([ + ("version".to_string(), "v2".to_string()), + ("success".to_string(), success.to_string()), + ("purpose".to_string(), purpose.to_string()), + ("attempt".to_string(), attempt.to_string()), + ("shell".to_string(), "bash".to_string()), + ("sandbox".to_string(), "none".to_string()), + ]); + if !success { + tags.insert("failure_reason".to_string(), "nonzero_exit".to_string()); + } + expected_durations.insert(tags.clone(), /*value*/ 1); + expected_counters.insert(tags, /*value*/ 1); + } + assert_eq!( + (counters, durations), + (expected_counters, expected_durations) + ); + Ok(()) +} + +#[test] +fn snapshot_size_limit_counts_state_and_environment_before_filtering() { + let half = "x".repeat(MAX_SNAPSHOT_BYTES / 2); + let oversized = format!("# Snapshot file\n# {half}\n\0\0\0FILTERED={half}\0"); + let policy = ExecEnvPolicy { + inherit: ShellEnvironmentPolicyInherit::All, + ignore_default_excludes: false, + exclude: vec!["FILTERED".to_string()], + r#set: HashMap::new(), + include_only: Vec::new(), + }; + assert!(parse_snapshot(ShellType::Bash, oversized.as_bytes(), Some(&policy)).is_err()); +} + +#[test] +fn snapshot_filters_profile_exports_after_capture() { + let policy = ExecEnvPolicy { + inherit: ShellEnvironmentPolicyInherit::All, + ignore_default_excludes: false, + exclude: vec!["PROFILE_DENIED".to_string()], + r#set: HashMap::from([("PROFILE_ALLOWED".to_string(), "override".to_string())]), + include_only: vec!["PROFILE_*".to_string()], + }; + let snapshot = parse_snapshot( + ShellType::Bash, + b"profile \xff noise\n# Snapshot file\nfunction profile_helper() { :; }\n\0alias profile_alias='profile_helper'\n\0PROFILE_DENIED\0export PROFILE_DENIED=denied\n\0NON_UTF8\0export NON_UTF8='\xff'\n\0\0PROFILE_ALLOWED=profile\0PROFILE_DENIED=denied\0PROFILE_SECRET=secret\0PWD=/tmp\0NON_UTF8=\xff\0", + Some(&policy), + ) + .expect("snapshot should parse"); + + assert_eq!( + snapshot.environment, + HashMap::from([("PROFILE_ALLOWED".to_string(), "override".to_string())]) + ); + assert_eq!( + snapshot.state, + "# Snapshot file\nfunction profile_helper() { :; }\nalias profile_alias='profile_helper'\n" + ); +} + +#[test] +fn snapshot_preserves_profile_exports_with_restrictive_inheritance() { + for inherit in [ + ShellEnvironmentPolicyInherit::None, + ShellEnvironmentPolicyInherit::Core, + ] { + let policy = ExecEnvPolicy { + inherit, + ignore_default_excludes: false, + exclude: vec!["PROFILE_DENIED".to_string()], + r#set: HashMap::new(), + include_only: Vec::new(), + }; + let snapshot = parse_snapshot( + ShellType::Bash, + b"# Snapshot file\n\0\0\0PROFILE_ALLOWED=profile\0SDKROOT=/sdk\0PROFILE_SECRET=secret\0PROFILE_DENIED=denied\0", + Some(&policy), + ) + .expect("snapshot should parse"); + + assert_eq!( + snapshot.environment, + HashMap::from([ + ("PROFILE_ALLOWED".to_string(), "profile".to_string()), + ("SDKROOT".to_string(), "/sdk".to_string()), + ]) + ); + } +} + +#[test] +fn snapshot_caches_only_unmanaged_proxy_state() { + for (exports, expected) in [ + ( + "PROFILE_ALLOWED=profile\0HTTP_PROXY=http://127.0.0.1:4321\0CODEX_NETWORK_PROXY_ACTIVE=1\0CODEX_NETWORK_PROXY_CREDENTIAL_BROKER_ACTIVE=1\0", + HashMap::from([("PROFILE_ALLOWED".to_string(), "profile".to_string())]), + ), + ( + "PROFILE_ALLOWED=profile\0HTTP_PROXY=http://user-proxy.example\0", + HashMap::from([ + ("PROFILE_ALLOWED".to_string(), "profile".to_string()), + ( + "HTTP_PROXY".to_string(), + "http://user-proxy.example".to_string(), + ), + ]), + ), + ] { + let output = format!("# Snapshot file\n\0\0\0{exports}"); + let snapshot = parse_snapshot(ShellType::Bash, output.as_bytes(), /*env_policy*/ None) + .expect("snapshot should parse"); + + assert_eq!(snapshot.environment, expected); + } +} +use codex_shell_command::shell_detect::ShellType; diff --git a/codex-rs/exec-server/src/telemetry.rs b/codex-rs/exec-server/src/telemetry.rs new file mode 100644 index 0000000000000000000000000000000000000000..16fe4d138a9986e5c60fdb3e0f8ad78e799ec363 --- /dev/null +++ b/codex-rs/exec-server/src/telemetry.rs @@ -0,0 +1,429 @@ +use std::sync::Arc; +use std::sync::Mutex; +use std::time::Duration; +use std::time::Instant; + +use codex_otel::MetricsClient; +use tracing::warn; + +/// Registry-issued identity captured from the executor's authenticated relay connection. +pub(crate) struct ExecutorRegistration { + pub(crate) environment_id: String, + pub(crate) executor_registration_id: String, +} + +impl ExecutorRegistration { + pub(crate) fn new(environment_id: String, executor_registration_id: String) -> Option { + if [&environment_id, &executor_registration_id] + .iter() + .any(|id| id.trim().is_empty() || id.len() > 256 || id.chars().any(char::is_control)) + { + return None; + } + Some(Self { + environment_id, + executor_registration_id, + }) + } +} + +const CONNECTIONS_ACTIVE_METRIC: &str = "exec_server_connections_active"; +const CONNECTIONS_ACTIVE_DESCRIPTION: &str = "Number of active exec-server connections."; +const CONNECTIONS_TOTAL_METRIC: &str = "exec_server_connections_total"; +const CONNECTIONS_TOTAL_DESCRIPTION: &str = "Total number of accepted exec-server connections."; +const REQUESTS_TOTAL_METRIC: &str = "exec_server_requests_total"; +const REQUESTS_TOTAL_DESCRIPTION: &str = "Total number of exec-server requests."; +const REQUEST_DURATION_METRIC: &str = "exec_server_request_duration_seconds"; +const REQUEST_DURATION_DESCRIPTION: &str = "Duration of exec-server requests in seconds."; +const REQUEST_TOTAL_DURATION_METRIC: &str = "exec_server_request_total_duration_seconds"; +const REQUEST_TOTAL_DURATION_DESCRIPTION: &str = "Total exec-server request duration in seconds, including queueing, from decoded receipt until response enqueue or disconnection."; +const REQUEST_QUEUE_DURATION_METRIC: &str = "exec_server_request_queue_duration_seconds"; +const REQUEST_QUEUE_DURATION_DESCRIPTION: &str = + "Time exec-server requests spend queued before execution in seconds."; +const PROCESSES_ACTIVE_METRIC: &str = "exec_server_processes_active"; +const PROCESSES_ACTIVE_DESCRIPTION: &str = "Number of active exec-server processes."; +const PROCESSES_FINISHED_TOTAL_METRIC: &str = "exec_server_processes_finished_total"; +const PROCESSES_FINISHED_TOTAL_DESCRIPTION: &str = + "Total number of finished exec-server processes."; +const PROCESS_DURATION_METRIC: &str = "exec_server_process_duration_seconds"; +const PROCESS_DURATION_DESCRIPTION: &str = "Duration of exec-server processes in seconds."; +const REMOTE_REGISTRATION_METRICS: OperationMetrics = OperationMetrics { + total_name: "exec_server_remote_registration_total", + total_description: "Total number of remote exec-server registration attempts.", + duration_name: "exec_server_remote_registration_duration_seconds", + duration_description: "Duration of remote exec-server registration attempts in seconds.", +}; +const REMOTE_RENDEZVOUS_METRICS: OperationMetrics = OperationMetrics { + total_name: "exec_server_remote_rendezvous_connect_total", + total_description: "Total number of remote exec-server rendezvous connection attempts.", + duration_name: "exec_server_remote_rendezvous_connect_duration_seconds", + duration_description: "Duration of remote exec-server rendezvous connection attempts in seconds.", +}; +const REMOTE_RECONNECTS_TOTAL_METRIC: &str = "exec_server_remote_reconnects_total"; +const REMOTE_RECONNECTS_TOTAL_DESCRIPTION: &str = "Total number of remote exec-server reconnects."; + +#[derive(Clone, Copy)] +struct OperationMetrics { + total_name: &'static str, + total_description: &'static str, + duration_name: &'static str, + duration_description: &'static str, +} + +#[derive(Clone, Copy)] +pub(crate) enum ConnectionTransport { + Relay, + Stdio, + WebSocket, +} + +impl ConnectionTransport { + fn metric_tag(self) -> &'static str { + match self { + Self::Relay => "relay", + Self::Stdio => "stdio", + Self::WebSocket => "websocket", + } + } +} + +#[derive(Clone, Default)] +pub struct ExecServerTelemetry { + inner: Option>, +} + +struct ExecServerTelemetryInner { + metrics: MetricsClient, + active: Arc>, +} + +#[derive(Default)] +struct ActiveCounts { + relay_connections: i64, + stdio_connections: i64, + websocket_connections: i64, + processes: i64, +} + +impl ActiveCounts { + fn connections(&self, transport: ConnectionTransport) -> i64 { + match transport { + ConnectionTransport::Relay => self.relay_connections, + ConnectionTransport::Stdio => self.stdio_connections, + ConnectionTransport::WebSocket => self.websocket_connections, + } + } +} + +pub(crate) struct ConnectionMetricGuard { + telemetry: ExecServerTelemetry, + transport: ConnectionTransport, +} + +pub(crate) struct ProcessMetricGuard { + telemetry: ExecServerTelemetry, + span: tracing::Span, + started_at: Instant, + result: &'static str, +} + +impl ExecServerTelemetry { + pub fn new(metrics: MetricsClient) -> Self { + let active = Arc::new(Mutex::new(ActiveCounts::default())); + register_active_gauges(&metrics, &active); + Self { + inner: Some(Arc::new(ExecServerTelemetryInner { metrics, active })), + } + } + + pub(crate) fn connection_started( + &self, + transport: ConnectionTransport, + ) -> ConnectionMetricGuard { + self.with_inner(|inner| { + inner.adjust_connection_count(transport, /*delta*/ 1); + inner.counter( + CONNECTIONS_TOTAL_METRIC, + CONNECTIONS_TOTAL_DESCRIPTION, + &[("transport", transport.metric_tag())], + ); + }); + ConnectionMetricGuard { + telemetry: self.clone(), + transport, + } + } + + pub(crate) fn request_completed( + &self, + method: &'static str, + result: &'static str, + duration: Duration, + total_duration: Duration, + ) { + self.with_inner(|inner| { + let tags = [("method", method), ("result", result)]; + inner.counter(REQUESTS_TOTAL_METRIC, REQUESTS_TOTAL_DESCRIPTION, &tags); + inner.duration( + REQUEST_DURATION_METRIC, + REQUEST_DURATION_DESCRIPTION, + duration, + &tags, + ); + inner.duration( + REQUEST_TOTAL_DURATION_METRIC, + REQUEST_TOTAL_DURATION_DESCRIPTION, + total_duration, + &tags, + ); + }); + } + + pub(crate) fn request_queue_completed(&self, method: &'static str, duration: Duration) { + self.with_inner(|inner| { + inner.duration( + REQUEST_QUEUE_DURATION_METRIC, + REQUEST_QUEUE_DURATION_DESCRIPTION, + duration, + &[("method", method)], + ); + }); + } + + #[cfg(unix)] + pub(crate) fn shell_snapshot_captured( + &self, + duration: Duration, + result: Result<(), &'static str>, + capture_tags: &[(&str, &str)], + ) { + // Local execution has no exec-server telemetry owner. Use the host's + // configured metrics client while preserving an explicit server client. + let Some(metrics) = self + .inner + .as_ref() + .map(|inner| inner.metrics.clone()) + .or_else(codex_otel::global) + else { + return; + }; + let success = if result.is_ok() { "true" } else { "false" }; + let mut tags = vec![("version", "v2"), ("success", success)]; + tags.extend_from_slice(capture_tags); + if let Err(failure_reason) = result { + tags.push(("failure_reason", failure_reason)); + } + let _ = metrics.record_duration("codex.shell_snapshot.duration_ms", duration, &tags); + let _ = metrics.counter("codex.shell_snapshot", /*inc*/ 1, &tags); + } + + pub(crate) fn remote_registration_completed(&self, result: &'static str, duration: Duration) { + self.record_operation(REMOTE_REGISTRATION_METRICS, result, duration); + } + + pub(crate) fn remote_rendezvous_completed(&self, result: &'static str, duration: Duration) { + self.record_operation(REMOTE_RENDEZVOUS_METRICS, result, duration); + } + + pub(crate) fn remote_reconnect(&self, reason: &'static str) { + self.with_inner(|inner| { + inner.counter( + REMOTE_RECONNECTS_TOTAL_METRIC, + REMOTE_RECONNECTS_TOTAL_DESCRIPTION, + &[("reason", reason)], + ); + }); + } + + pub(crate) fn process_started(&self, process_id: &str) -> ProcessMetricGuard { + self.with_inner(|inner| { + inner.adjust_process_count(/*delta*/ 1); + }); + let parent = codex_otel::current_span_w3c_trace_context(); + // `parent:` accepts a local tracing span/ID, not a W3C context. A local + // parent would keep the request span alive until process exit and delay + // its export. Use `parent: None`, then set the W3C parent below to link + // the spans without retaining the request span. + let span = tracing::info_span!( + parent: None, + "codex.exec_server.process", + otel.kind = "internal", + process.id = process_id, + result = tracing::field::Empty, + ); + if let Some(parent) = parent { + codex_otel::set_parent_from_w3c_trace_context(&span, &parent); + } + ProcessMetricGuard { + telemetry: self.clone(), + span, + started_at: Instant::now(), + result: "unknown", + } + } + + fn process_finished(&self, result: &'static str, duration: Duration) { + self.with_inner(|inner| { + inner.adjust_process_count(/*delta*/ -1); + inner.counter( + PROCESSES_FINISHED_TOTAL_METRIC, + PROCESSES_FINISHED_TOTAL_DESCRIPTION, + &[("result", result)], + ); + inner.duration( + PROCESS_DURATION_METRIC, + PROCESS_DURATION_DESCRIPTION, + duration, + &[("result", result)], + ); + }); + } + + fn connection_finished(&self, transport: ConnectionTransport) { + self.with_inner(|inner| { + inner.adjust_connection_count(transport, /*delta*/ -1); + }); + } + + fn with_inner(&self, emit: impl FnOnce(&ExecServerTelemetryInner)) { + if let Some(inner) = &self.inner { + emit(inner); + } + } + + fn record_operation( + &self, + metrics: OperationMetrics, + result: &'static str, + duration: Duration, + ) { + self.with_inner(|inner| { + let tags = [("result", result)]; + inner.counter(metrics.total_name, metrics.total_description, &tags); + inner.duration( + metrics.duration_name, + metrics.duration_description, + duration, + &tags, + ); + }); + } +} + +impl Drop for ConnectionMetricGuard { + fn drop(&mut self) { + self.telemetry.connection_finished(self.transport); + } +} + +impl ProcessMetricGuard { + pub(crate) fn finish(mut self, result: &'static str) { + self.result = result; + } +} + +impl Drop for ProcessMetricGuard { + fn drop(&mut self) { + self.span.record("result", self.result); + self.telemetry + .process_finished(self.result, self.started_at.elapsed()); + } +} + +impl ExecServerTelemetryInner { + fn active_counts(&self) -> std::sync::MutexGuard<'_, ActiveCounts> { + // These are independent integer counts, so a panic cannot leave a cross-field invariant + // half-updated. Recovering a poisoned lock preserves the last completed count update. + self.active + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + } + + fn adjust_connection_count(&self, transport: ConnectionTransport, delta: i64) { + let mut active = self.active_counts(); + let count = match transport { + ConnectionTransport::Relay => &mut active.relay_connections, + ConnectionTransport::Stdio => &mut active.stdio_connections, + ConnectionTransport::WebSocket => &mut active.websocket_connections, + }; + *count += delta; + } + + fn adjust_process_count(&self, delta: i64) { + let mut active = self.active_counts(); + active.processes += delta; + } + + fn counter(&self, name: &str, description: &str, tags: &[(&str, &str)]) { + if self + .metrics + .counter_with_description(name, description, /*inc*/ 1, tags) + .is_err() + { + warn!(metric = name, "failed to emit exec-server counter"); + } + } + + fn duration(&self, name: &str, description: &str, duration: Duration, tags: &[(&str, &str)]) { + if self + .metrics + .record_duration_seconds_with_description(name, description, duration, tags) + .is_err() + { + warn!(metric = name, "failed to emit exec-server duration"); + } + } +} + +fn register_active_gauges(metrics: &MetricsClient, active: &Arc>) { + for transport in [ + ConnectionTransport::Relay, + ConnectionTransport::Stdio, + ConnectionTransport::WebSocket, + ] { + register_active_gauge( + metrics, + active, + CONNECTIONS_ACTIVE_METRIC, + CONNECTIONS_ACTIVE_DESCRIPTION, + &[("transport", transport.metric_tag())], + move |active| active.connections(transport), + ); + } + + register_active_gauge( + metrics, + active, + PROCESSES_ACTIVE_METRIC, + PROCESSES_ACTIVE_DESCRIPTION, + &[], + |active| active.processes, + ); +} + +fn register_active_gauge( + metrics: &MetricsClient, + active: &Arc>, + name: &str, + description: &str, + tags: &[(&str, &str)], + read: impl Fn(&ActiveCounts) -> i64 + Send + Sync + 'static, +) { + let active = Arc::clone(active); + if metrics + .register_observable_gauge_with_description( + name, + description, + move || { + let active = active + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + read(&active) + }, + tags, + ) + .is_err() + { + warn!(metric = name, "failed to register exec-server gauge"); + } +} diff --git a/codex-rs/exec-server/src/trace_context.rs b/codex-rs/exec-server/src/trace_context.rs new file mode 100644 index 0000000000000000000000000000000000000000..f6fd898c0575dcadd1d652c0f79eedf3f623267c --- /dev/null +++ b/codex-rs/exec-server/src/trace_context.rs @@ -0,0 +1,40 @@ +use http::HeaderMap; +use http::HeaderValue; + +pub(crate) fn current_rendezvous_headers() -> HeaderMap { + let mut headers = current_trace_context_headers(); + for (header, variable) in [ + ("x-cluster-name", "OPENAI_CLUSTER"), + ("x-openai-internal-caller", "DD_SERVICE"), + ] { + if let Ok(value) = std::env::var(variable) + && !value.is_empty() + && let Ok(value) = HeaderValue::try_from(value) + { + headers.insert(header, value); + } + } + headers +} + +pub(crate) fn current_trace_context_headers() -> HeaderMap { + let mut headers = HeaderMap::new(); + let Some(trace) = codex_otel::current_span_w3c_trace_context() else { + return headers; + }; + if let Some(traceparent) = trace.traceparent + && let Ok(value) = HeaderValue::try_from(traceparent) + { + headers.insert("traceparent", value); + } + if let Some(tracestate) = trace.tracestate + && let Ok(value) = HeaderValue::try_from(tracestate) + { + headers.insert("tracestate", value); + } + headers +} + +#[cfg(test)] +#[path = "trace_context_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/trace_context_tests.rs b/codex-rs/exec-server/src/trace_context_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..92a820afc43a4bdf2eb78104a6ab7cdec144ef44 --- /dev/null +++ b/codex-rs/exec-server/src/trace_context_tests.rs @@ -0,0 +1,27 @@ +use opentelemetry::trace::TracerProvider as _; +use opentelemetry_sdk::trace::SdkTracerProvider; +use tracing_subscriber::prelude::*; + +use super::current_trace_context_headers; + +#[test] +fn creates_traceparent_header_from_current_span() { + let provider = SdkTracerProvider::builder().build(); + let tracer = provider.tracer("exec-server-test"); + let subscriber = + tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer)); + let _guard = subscriber.set_default(); + tracing::callsite::rebuild_interest_cache(); + let span = tracing::info_span!("outbound-request"); + let _entered = span.enter(); + + let headers = current_trace_context_headers(); + + let traceparent = headers + .get("traceparent") + .expect("traceparent header") + .to_str() + .expect("valid traceparent header"); + assert!(traceparent.starts_with("00-")); + assert_eq!(traceparent.len(), 55); +} diff --git a/codex-rs/exec-server/src/websocket_pong_watchdog.rs b/codex-rs/exec-server/src/websocket_pong_watchdog.rs new file mode 100644 index 0000000000000000000000000000000000000000..b656772b3d276155353ed4e80b9eb72bb8286e50 --- /dev/null +++ b/codex-rs/exec-server/src/websocket_pong_watchdog.rs @@ -0,0 +1,44 @@ +use std::time::Duration; + +use tokio::time::Instant; + +#[cfg(test)] +pub(crate) const WEBSOCKET_PONG_TIMEOUT: Duration = Duration::from_millis(100); +#[cfg(not(test))] +pub(crate) const WEBSOCKET_PONG_TIMEOUT: Duration = Duration::from_secs(60); +pub(crate) const WEBSOCKET_PONG_TIMEOUT_REASON: &str = "pong_timeout"; + +/// Tracks whether a WebSocket peer has acknowledged a keepalive ping. +pub(crate) struct WebSocketPongWatchdog { + timeout: Duration, + deadline: Option, +} + +impl WebSocketPongWatchdog { + pub(crate) fn new(timeout: Duration) -> Self { + Self { + timeout, + deadline: None, + } + } + + pub(crate) fn ping_sent(&mut self, now: Instant) { + self.deadline.get_or_insert(now + self.timeout); + } + + pub(crate) fn deadline(&self) -> Option { + self.deadline + } + + pub(crate) fn write_deadline(&self, now: Instant) -> Instant { + self.deadline.unwrap_or(now + self.timeout) + } + + pub(crate) fn received_pong(&mut self) { + self.deadline = None; + } +} + +#[cfg(test)] +#[path = "websocket_pong_watchdog_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/websocket_pong_watchdog_tests.rs b/codex-rs/exec-server/src/websocket_pong_watchdog_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..2b15a0d95223a71e148e09721a12a7821ee8488f --- /dev/null +++ b/codex-rs/exec-server/src/websocket_pong_watchdog_tests.rs @@ -0,0 +1,34 @@ +use std::time::Duration; + +use pretty_assertions::assert_eq; +use tokio::time::Instant; + +use super::WebSocketPongWatchdog; + +#[test] +fn repeated_ping_does_not_extend_deadline() { + let started_at = Instant::now(); + let timeout = Duration::from_secs(2); + let mut watchdog = WebSocketPongWatchdog::new(timeout); + + watchdog.ping_sent(started_at); + watchdog.ping_sent(started_at + Duration::from_secs(1)); + assert_eq!( + watchdog.write_deadline(started_at + Duration::from_secs(1)), + started_at + timeout + ); + assert_eq!(watchdog.deadline(), Some(started_at + timeout)); +} + +#[test] +fn pong_starts_a_fresh_deadline() { + let started_at = Instant::now(); + let timeout = Duration::from_secs(2); + let mut watchdog = WebSocketPongWatchdog::new(timeout); + + watchdog.ping_sent(started_at); + watchdog.received_pong(); + assert_eq!(watchdog.deadline(), None); + watchdog.ping_sent(started_at + timeout); + assert_eq!(watchdog.deadline(), Some(started_at + timeout + timeout)); +} diff --git a/codex-rs/exec-server/testing/BUILD.bazel b/codex-rs/exec-server/testing/BUILD.bazel new file mode 100644 index 0000000000000000000000000000000000000000..93cb4ceab9127c9792e349e1138632da760f6d05 --- /dev/null +++ b/codex-rs/exec-server/testing/BUILD.bazel @@ -0,0 +1,46 @@ +load("@rules_rust//rust:defs.bzl", "rust_binary", "rust_library") + +rust_library( + name = "wine-exec-server-test-support", + testonly = True, + srcs = ["wine_exec_server.rs"], + crate_name = "wine_exec_server_test_support", + crate_root = "wine_exec_server.rs", + target_compatible_with = ["@platforms//os:linux"], + visibility = ["//codex-rs/core/tests/remote_env_windows:__pkg__"], + deps = [ + "//bazel/rules/testing/wine:wine_test_support", + "//codex-rs/utils/cargo-bin", + "@crates//:anyhow", + "@crates//:tokio", + ], +) + +rust_binary( + name = "wine-exec-test-runner", + testonly = True, + srcs = ["wine_remote_test_runner.rs"], + crate_name = "wine_exec_test_runner", + crate_root = "wine_remote_test_runner.rs", + target_compatible_with = ["@platforms//os:linux"], + visibility = ["//visibility:public"], + deps = [ + ":wine-exec-server-test-support", + "@crates//:anyhow", + "@crates//:tokio", + ], +) + +rust_binary( + name = "exec-server", + testonly = True, + srcs = ["exec_server.rs"], + crate_name = "exec_server", + crate_root = "exec_server.rs", + visibility = ["//visibility:public"], + deps = [ + "//codex-rs/exec-server", + "//codex-rs/http-client", + "@crates//:tokio", + ], +) diff --git a/codex-rs/exec-server/testing/README.md b/codex-rs/exec-server/testing/README.md new file mode 100644 index 0000000000000000000000000000000000000000..77f8e990d430bef30ce7c7c21cd34cd878910e77 --- /dev/null +++ b/codex-rs/exec-server/testing/README.md @@ -0,0 +1,5 @@ +# Windows exec-server fixture + +This directory contains the small Windows exec-server binary used by +foreign-OS tests. It links only `codex-exec-server` because the full Codex +Windows graph does not yet cross-build with Bazel. diff --git a/codex-rs/exec-server/testing/exec_server.rs b/codex-rs/exec-server/testing/exec_server.rs new file mode 100644 index 0000000000000000000000000000000000000000..9c6773995be9641acab41aba080c0d007b33a582 --- /dev/null +++ b/codex-rs/exec-server/testing/exec_server.rs @@ -0,0 +1,38 @@ +//! Minimal exec-server fixture for Bazel-only integration tests. +//! +//! Linking only exec-server avoids depending on the full Codex CLI binary +//! when a test only needs a WebSocket executor endpoint. It handles the arg0 +//! helper mode because sandboxed process requests re-exec this binary. + +use codex_exec_server::ExecServerRuntimePaths; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use std::ffi::OsStr; + +const CODEX_LINUX_SANDBOX_EXE_ENV_VAR: &str = "CODEX_TEST_LINUX_SANDBOX_EXE"; + +fn main() -> Result<(), Box> { + let mut args = std::env::args_os(); + let _ = args.next(); + let argv1 = args.next(); + #[cfg(unix)] + if argv1.as_deref() == Some(OsStr::new(codex_exec_server::CODEX_ARG0_EXEC_HELPER_ARG1)) { + codex_exec_server::run_arg0_exec_helper_main(); + } + if argv1.as_deref() == Some(OsStr::new(codex_exec_server::CODEX_FS_HELPER_ARG1)) { + codex_exec_server::run_fs_helper_main(); + } + + let current_exe = std::env::current_exe()?; + let codex_linux_sandbox_exe = + std::env::var_os(CODEX_LINUX_SANDBOX_EXE_ENV_VAR).map(std::path::PathBuf::from); + let runtime_paths = ExecServerRuntimePaths::new(current_exe, codex_linux_sandbox_exe)?; + tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build()? + .block_on(codex_exec_server::run_main( + "ws://127.0.0.1:0", + runtime_paths, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )) +} diff --git a/codex-rs/exec-server/testing/wine_exec_server.rs b/codex-rs/exec-server/testing/wine_exec_server.rs new file mode 100644 index 0000000000000000000000000000000000000000..a3643f8668c99baec60a6d1f382021b1607a8d42 --- /dev/null +++ b/codex-rs/exec-server/testing/wine_exec_server.rs @@ -0,0 +1,46 @@ +//! Test support for running the Windows exec-server under Wine. + +use std::future::Future; +use std::path::PathBuf; + +use anyhow::Context; +use anyhow::Result; +use tokio::io::AsyncBufReadExt; +use tokio::io::BufReader; +use wine_test_support::WineTestCommand; + +/// Runs the Windows exec-server under Wine for the duration of a scoped operation. +pub struct WineExecServer; + +impl WineExecServer { + /// Starts the server, passes its WebSocket URL and Wine prefix to `operation`, and tears it + /// down afterward. + pub async fn scope(self, operation: F) -> Result + where + F: FnOnce(String, PathBuf) -> Fut, + Fut: Future>, + { + let executable = codex_utils_cargo_bin::cargo_bin("wine-windows-exec-server")?; + let mut exec_server = WineTestCommand::new(executable) + .env("CODEX_HOME", r"C:\codex-home") + .spawn()?; + let wine_prefix = exec_server.prefix_path().to_path_buf(); + let stdout = exec_server.take_stdout(); + + exec_server + .scope(async move { + let mut lines = BufReader::new(stdout).lines(); + let exec_server_url = loop { + let line = lines + .next_line() + .await? + .context("Wine exec-server exited before reporting its URL")?; + if line.starts_with("ws://") { + break line; + } + }; + operation(exec_server_url, wine_prefix).await + }) + .await + } +} diff --git a/codex-rs/exec-server/testing/wine_remote_test_runner.rs b/codex-rs/exec-server/testing/wine_remote_test_runner.rs new file mode 100644 index 0000000000000000000000000000000000000000..883dab3cbcb14fae61f683a83d78a07db0e8e13f --- /dev/null +++ b/codex-rs/exec-server/testing/wine_remote_test_runner.rs @@ -0,0 +1,65 @@ +use std::ffi::OsStr; +use std::ffi::OsString; +use std::path::PathBuf; +use std::process::Stdio; + +use anyhow::Context; +use anyhow::Result; +use tokio::process::Command; +use wine_exec_server_test_support::WineExecServer; + +const TEST_BINARY_ENV_VAR: &str = "CODEX_WINE_EXEC_TEST_BINARY"; +const TEST_ENVIRONMENT_ENV_VAR: &str = "CODEX_TEST_ENVIRONMENT"; +const REMOTE_EXEC_SERVER_URL_ENV_VAR: &str = "CODEX_TEST_REMOTE_EXEC_SERVER_URL"; +const LEGACY_REMOTE_ENV_ENV_VAR: &str = "CODEX_TEST_REMOTE_ENV"; +const DOCKER_CONTAINER_ENV_VAR: &str = "CODEX_TEST_REMOTE_ENV_CONTAINER_NAME"; + +#[tokio::main(flavor = "multi_thread", worker_threads = 2)] +async fn main() -> Result<()> { + let test_binary = + PathBuf::from(std::env::var_os(TEST_BINARY_ENV_VAR).with_context(|| { + format!("{TEST_BINARY_ENV_VAR} must be set by the Bazel test rule") + })?); + let forwarded_args = std::env::args_os().skip(1).collect::>(); + + if is_terse_list_request(&forwarded_args) { + let status = Command::new(&test_binary) + .args(&forwarded_args) + .status() + .await + .context("list integration tests")?; + anyhow::ensure!( + status.success(), + "listing integration tests exited with {status}" + ); + return Ok(()); + } + + WineExecServer + .scope(|exec_server_url, _wine_prefix| async move { + let mut command = Command::new(test_binary); + command + .env(TEST_ENVIRONMENT_ENV_VAR, "wine-exec") + .env(REMOTE_EXEC_SERVER_URL_ENV_VAR, exec_server_url) + .env_remove(LEGACY_REMOTE_ENV_ENV_VAR) + .env_remove(DOCKER_CONTAINER_ENV_VAR) + .args(forwarded_args) + .stdin(Stdio::null()) + .stdout(Stdio::inherit()) + .stderr(Stdio::inherit()) + .kill_on_drop(true); + + let status = command.status().await.context("run integration tests")?; + anyhow::ensure!(status.success(), "integration tests exited with {status}"); + Ok(()) + }) + .await +} + +fn is_terse_list_request(args: &[OsString]) -> bool { + args.iter().map(OsString::as_os_str).eq([ + OsStr::new("--list"), + OsStr::new("--format"), + OsStr::new("terse"), + ]) +} diff --git a/codex-rs/exec-server/tests/accepted_websocket.rs b/codex-rs/exec-server/tests/accepted_websocket.rs new file mode 100644 index 0000000000000000000000000000000000000000..a3caf7b3de8f8cd0afdc73355679a1f99c88a5b5 --- /dev/null +++ b/codex-rs/exec-server/tests/accepted_websocket.rs @@ -0,0 +1,807 @@ +mod common; + +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use axum::Router; +use axum::extract::State; +use axum::extract::WebSocketUpgrade; +use axum::response::IntoResponse; +use axum::routing::any; +use codex_api::AuthProvider; +#[cfg(unix)] +use codex_exec_server::EnvironmentConnectionState; +use codex_exec_server::EnvironmentInfo; +use codex_exec_server::EnvironmentManager; +use codex_exec_server::EnvironmentObservedStatus; +use codex_exec_server::EnvironmentStatus; +use codex_exec_server::EnvironmentStatusKind; +use codex_exec_server::ExecParams; +#[cfg(unix)] +use codex_exec_server::ExecProcessEvent; +use codex_exec_server::ExecResponse; +use codex_exec_server::ExecServerClientConnectOptions; +use codex_exec_server::ExecServerRuntimePaths; +use codex_exec_server::InitializeParams; +use codex_exec_server::InitializeResponse; +use codex_exec_server::ProcessId; +use codex_exec_server::ReadParams; +use codex_exec_server::ReadResponse; +use codex_exec_server::RemoteEnvironmentConfig; +use codex_exec_server::RemoteEnvironmentTransport; +#[cfg(unix)] +use codex_exec_server::WriteStatus; +use codex_exec_server_protocol::JSONRPCError; +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCNotification; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_utils_path_uri::PathUri; +use common::exec_server::DisconnectableWebSocketProxy; +use futures::SinkExt; +use futures::StreamExt; +use http::HeaderMap; +use http::HeaderValue; +use pretty_assertions::assert_eq; +use tokio::net::TcpListener; +use tokio::sync::mpsc; +use tokio::task::JoinHandle; +use tokio::time::timeout; +use tokio_tungstenite::MaybeTlsStream; +use tokio_tungstenite::WebSocketStream; +use tokio_tungstenite::connect_async; +use tokio_tungstenite::tungstenite::Message; +use tokio_util::task::AbortOnDropHandle; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +const TEST_TIMEOUT: Duration = Duration::from_secs(5); + +type AcceptedSocket = axum::extract::ws::WebSocket; +const SESSION_ALREADY_ATTACHED_ERROR_CODE: i64 = -32010; + +#[derive(Debug)] +struct DirectExecutorAuth; + +impl AuthProvider for DirectExecutorAuth { + fn add_auth_headers(&self, headers: &mut HeaderMap) { + headers.insert( + http::header::AUTHORIZATION, + HeaderValue::from_static("AWS4-HMAC-SHA256 test-signature"), + ); + } +} + +#[tokio::test] +async fn accepted_websocket_rejects_initial_resume_session_id() -> Result<()> { + let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?; + let (_socket, _) = connect_async(&websocket_url).await?; + let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv()) + .await? + .context("accepted websocket channel should remain open")?; + let mut options = accepted_options(); + options.resume_session_id = Some("session-1".to_string()); + + let error = EnvironmentManager::from_accepted_websocket( + "environment-1".to_string(), + accepted_websocket, + options, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .await + .expect_err("initial accepted websocket should reject a resume session ID"); + + assert!( + error + .to_string() + .contains("initial connection cannot resume a session"), + "unexpected error: {error}" + ); + + server_task.abort(); + let _ = server_task.await; + Ok(()) +} + +#[tokio::test] +async fn accepted_websocket_environment_info_uses_initialization_metadata() -> Result<()> { + let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?; + let (_socket, manager) = + connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?; + let environment = manager + .default_environment() + .context("accepted environment")?; + + assert_eq!( + timeout(TEST_TIMEOUT, environment.info()).await??, + EnvironmentInfo::local() + ); + + server_task.abort(); + let _ = server_task.await; + Ok(()) +} + +#[tokio::test] +async fn accepted_websocket_interoperates_and_recovers_with_real_direct_executor() -> Result<()> { + let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?; + let proxy = DisconnectableWebSocketProxy::new(&websocket_url).await?; + let registry = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/cloud/environment/environment-1/direct/register")) + .and(header("authorization", "AWS4-HMAC-SHA256 test-signature")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "environment_id": "environment-1", + "transport": "direct_jsonrpc_v1", + "registration_id": "registration-1", + "url": proxy.websocket_url(), + }))) + .expect(1) + .mount(®istry) + .await; + + let http_client_factory = HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault); + let config = RemoteEnvironmentConfig::new_with_transport( + registry.uri(), + "environment-1".to_string(), + RemoteEnvironmentTransport::Direct, + Arc::new(DirectExecutorAuth), + http_client_factory.clone(), + )?; + let (codex_exe, codex_linux_sandbox_exe) = common::current_test_binary_helper_paths()?; + let runtime_paths = ExecServerRuntimePaths::new(codex_exe, codex_linux_sandbox_exe)?; + let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel(); + let executor_task = AbortOnDropHandle::new(tokio::spawn( + codex_exec_server::run_remote_environment_until_shutdown( + config, + runtime_paths, + async move { + let _ = shutdown_rx.await; + }, + ), + )); + let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv()) + .await? + .context("direct executor websocket should be accepted")?; + let manager = timeout( + TEST_TIMEOUT, + EnvironmentManager::from_accepted_websocket( + "environment-1".to_string(), + accepted_websocket, + accepted_options(), + http_client_factory, + ), + ) + .await??; + let environment = manager + .default_environment() + .context("direct executor environment should be installed")?; + + let expected_info = EnvironmentInfo::local(); + assert_eq!( + timeout(TEST_TIMEOUT, environment.force_info()).await??, + expected_info + ); + let files = tempfile::tempdir()?; + let large_file_path = files.path().join("large-response.bin"); + let large_file_contents = vec![0x5a; 128 * 1024]; + tokio::fs::write(&large_file_path, &large_file_contents).await?; + assert_eq!( + timeout( + TEST_TIMEOUT, + environment.get_filesystem().read_file( + &PathUri::from_host_native_path(&large_file_path)?, + Default::default(), + /*sandbox*/ None, + ) + ) + .await??, + large_file_contents, + ); + + // The process fixture uses a POSIX shell; metadata and shutdown remain tested on all platforms. + #[cfg(unix)] + { + let mut proxy = proxy; + let backend = environment.get_exec_backend(); + let temp_dir = tempfile::TempDir::new()?; + let gate_path = temp_dir.path().join("release-output"); + let emitted_path = temp_dir.path().join("output-emitted"); + let session = timeout( + TEST_TIMEOUT, + backend.start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-recover"), + argv: vec![ + "/bin/sh".to_string(), + "-c".to_string(), + concat!( + "printf 'ready:%s\\n' \"$$\"; ", + "while [ ! -f \"$GATE\" ]; do /bin/sleep 0.01; done; ", + "printf 'during:%s\\n' \"$$\"; ", + ": > \"$EMITTED\"; ", + "IFS= read -r line; ", + "printf 'after:%s:%s\\n' \"$$\" \"$line\"; ", + "exit 7", + ) + .to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: HashMap::from([ + ( + "GATE".to_string(), + gate_path.to_string_lossy().into_owned(), + ), + ( + "EMITTED".to_string(), + emitted_path.to_string_lossy().into_owned(), + ), + ]), + tty: false, + pipe_stdin: true, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }), + ) + .await??; + + let process = Arc::clone(&session.process); + let mut events = process.subscribe_events(); + let mut output = Vec::new(); + let mut last_seq = 0; + while !output.ends_with(b"\n") { + match timeout(Duration::from_secs(5), events.recv()).await?? { + ExecProcessEvent::Output(chunk) => { + assert_eq!(chunk.seq, last_seq + 1); + last_seq = chunk.seq; + output.extend_from_slice(&chunk.chunk.into_inner()); + } + event => anyhow::bail!("expected ready output before disconnect, got {event:?}"), + } + } + let ready = String::from_utf8(output.clone())?; + let pid = ready + .strip_prefix("ready:") + .and_then(|line| line.strip_suffix('\n')) + .context("ready output should contain the process id")? + .to_string(); + + let mut connection_state = environment + .subscribe_connection_state() + .context("direct environment connection state")?; + assert_eq!( + *connection_state.borrow_and_update(), + EnvironmentConnectionState::Connected + ); + proxy.pause_and_disconnect().await?; + timeout( + TEST_TIMEOUT, + connection_state.wait_for(|state| *state == EnvironmentConnectionState::Disconnected), + ) + .await??; + tokio::fs::write(&gate_path, b"").await?; + timeout(Duration::from_secs(5), async { + while tokio::fs::metadata(&emitted_path).await.is_err() { + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .context("process did not emit output while disconnected")?; + + let process_for_read = Arc::clone(&process); + let mut pending_read = tokio::spawn(async move { + process_for_read + .read( + /*after_seq*/ Some(last_seq), + /*max_bytes*/ None, + /*wait_ms*/ Some(0), + ) + .await + }); + assert!( + timeout(Duration::from_millis(200), &mut pending_read) + .await + .is_err(), + "process reads should wait while recovery is in progress" + ); + proxy.resume()?; + let replacement = timeout(TEST_TIMEOUT, accepted_sockets.recv()) + .await? + .context("real Direct executor should reconnect")?; + timeout( + TEST_TIMEOUT, + manager.replace_accepted_websocket("environment-1", replacement), + ) + .await??; + timeout( + TEST_TIMEOUT, + connection_state.wait_for(|state| *state == EnvironmentConnectionState::Connected), + ) + .await??; + assert!(Arc::ptr_eq( + &environment, + &manager + .default_environment() + .context("same environment should remain installed")?, + )); + assert_eq!( + timeout(TEST_TIMEOUT, environment.force_info()).await??, + expected_info + ); + + let recovered_read = timeout(Duration::from_secs(5), pending_read) + .await + .context("timed out waiting for a read after recovery")??; + let recovered_read = recovered_read?; + assert_eq!(recovered_read.failure, None); + let recovered_output = recovered_read + .chunks + .into_iter() + .flat_map(|chunk| chunk.chunk.into_inner()) + .collect::>(); + assert_eq!( + String::from_utf8(recovered_output)?, + format!("during:{pid}\n") + ); + + let write = timeout(Duration::from_secs(5), process.write(b"hello\n".to_vec())) + .await + .context("timed out waiting for a write after recovery")??; + assert_eq!(write.status, WriteStatus::Accepted); + + let mut saw_exit = false; + loop { + match timeout(Duration::from_secs(5), events.recv()).await?? { + ExecProcessEvent::Output(chunk) => { + assert_eq!(chunk.seq, last_seq + 1); + last_seq = chunk.seq; + output.extend_from_slice(&chunk.chunk.into_inner()); + } + ExecProcessEvent::Exited { seq, exit_code, .. } => { + assert_eq!(seq, last_seq + 1); + assert_eq!(exit_code, 7); + last_seq = seq; + saw_exit = true; + } + ExecProcessEvent::Closed { seq } => { + assert!(saw_exit, "closed must be delivered after exit"); + assert_eq!(seq, last_seq + 1); + break; + } + ExecProcessEvent::Failed(message) => { + anyhow::bail!("process recovery failed: {message}"); + } + } + } + assert_eq!( + String::from_utf8(output)?, + format!("ready:{pid}\nduring:{pid}\nafter:{pid}:hello\n") + ); + } + + registry.verify().await; + let registrations = registry + .received_requests() + .await + .context("registration requests")?; + assert_eq!( + registrations[0].body, + br#"{"transport":"direct_jsonrpc_v1"}"# + ); + + let _ = shutdown_tx.send(()); + timeout(TEST_TIMEOUT, executor_task).await???; + server_task.abort(); + let _ = server_task.await; + Ok(()) +} + +#[tokio::test] +async fn accepted_websocket_environment_is_ready_immediately() -> Result<()> { + let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?; + let (mut socket, manager) = + connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?; + + let status_task = + tokio::spawn(async move { manager.get_environment_status("environment-1").await }); + let request = receive_jsonrpc(&mut socket).await?; + let JSONRPCMessage::Request(JSONRPCRequest { id, method, .. }) = request else { + anyhow::bail!("expected environment status request, got {request:?}"); + }; + assert_eq!(method, "environment/status"); + send_jsonrpc( + &mut socket, + JSONRPCMessage::Response(JSONRPCResponse { + id, + result: serde_json::to_value(EnvironmentStatus { + status: EnvironmentStatusKind::Ready, + })?, + }), + ) + .await?; + + assert_eq!( + timeout(TEST_TIMEOUT, status_task).await??, + Some(EnvironmentObservedStatus::Ready) + ); + + server_task.abort(); + let _ = server_task.await; + Ok(()) +} + +#[tokio::test] +async fn accepted_websocket_replacement_retires_old_socket_and_retries() -> Result<()> { + let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?; + let (mut first_socket, manager) = + connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?; + let (mut rejected_socket, rejected_websocket) = + connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?; + manager + .replace_accepted_websocket("environment-1", rejected_websocket) + .await?; + let previous_socket_event = timeout(TEST_TIMEOUT, first_socket.next()) + .await + .context("the previous accepted websocket should be retired before replacement")?; + assert!( + matches!( + previous_socket_event, + None | Some(Ok(Message::Close(_))) | Some(Err(_)) + ), + "the previous accepted websocket should close before replacement: {previous_socket_event:?}" + ); + let initialize = receive_jsonrpc(&mut rejected_socket).await?; + let JSONRPCMessage::Request(JSONRPCRequest { id, method, .. }) = initialize else { + anyhow::bail!("expected replacement initialize request, got {initialize:?}"); + }; + assert_eq!(method, "initialize"); + + let (_overlapping_socket, overlapping_websocket) = + connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?; + manager + .replace_accepted_websocket("environment-1", overlapping_websocket) + .await + .expect_err("an overlapping replacement should be rejected"); + + send_jsonrpc( + &mut rejected_socket, + JSONRPCMessage::Error(JSONRPCError { + id, + error: JSONRPCErrorError { + code: SESSION_ALREADY_ATTACHED_ERROR_CODE, + message: "session session-1 is already attached to another connection".to_string(), + data: None, + }, + }), + ) + .await?; + let rejected_socket_event = timeout(TEST_TIMEOUT, rejected_socket.next()) + .await + .context("rejected replacement websocket should close")?; + assert!( + matches!( + rejected_socket_event, + None | Some(Ok(Message::Close(_))) | Some(Err(_)) + ), + "rejected replacement websocket should close: {rejected_socket_event:?}" + ); + + let (mut replacement_socket, replacement_websocket) = + connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?; + manager + .replace_accepted_websocket("environment-1", replacement_websocket) + .await?; + complete_initialize(&mut replacement_socket, "session-1", Some("session-1")).await?; + + server_task.abort(); + let _ = server_task.await; + Ok(()) +} + +#[tokio::test] +async fn accepted_websocket_reconnect_recovers_running_process_and_output() -> Result<()> { + let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?; + let (mut first_socket, manager) = + connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?; + let environment = manager + .default_environment() + .context("default environment should be installed")?; + let backend = environment.get_exec_backend(); + let process_id = ProcessId::from("process-1"); + let process_task = tokio::spawn({ + let process_id = process_id.clone(); + async move { + backend + .start(ExecParams { + metadata: Default::default(), + process_id, + argv: vec!["test-command".to_string()], + cwd: PathUri::parse("file:///workspace")?, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + shell_snapshot: None, + }) + .await + .map_err(anyhow::Error::from) + } + }); + let request = receive_jsonrpc(&mut first_socket).await?; + let JSONRPCMessage::Request(JSONRPCRequest { + id, method, params, .. + }) = request + else { + anyhow::bail!("expected process start request, got {request:?}"); + }; + assert_eq!(method, "process/start"); + assert_eq!( + serde_json::from_value::(params.context("process params should exist")?)? + .process_id, + process_id + ); + send_jsonrpc( + &mut first_socket, + JSONRPCMessage::Response(JSONRPCResponse { + id, + result: serde_json::to_value(ExecResponse { + process_id: process_id.clone(), + sandbox_type: None, + })?, + }), + ) + .await?; + let process = timeout(TEST_TIMEOUT, process_task).await???.process; + + first_socket.close(/*close_frame*/ None).await?; + let (mut replacement_socket, replacement_websocket) = + connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?; + manager + .replace_accepted_websocket("environment-1", replacement_websocket) + .await?; + complete_initialize(&mut replacement_socket, "session-1", Some("session-1")).await?; + + let request = receive_jsonrpc(&mut replacement_socket).await?; + let JSONRPCMessage::Request(JSONRPCRequest { + id, method, params, .. + }) = request + else { + anyhow::bail!("expected recovery process read request, got {request:?}"); + }; + assert_eq!(method, "process/read"); + assert_eq!( + serde_json::from_value::(params.context("read params should exist")?)?, + ReadParams { + process_id: process_id.clone(), + after_seq: Some(0), + max_bytes: None, + wait_ms: Some(0), + } + ); + send_jsonrpc( + &mut replacement_socket, + JSONRPCMessage::Response(JSONRPCResponse { + id, + result: serde_json::to_value(ReadResponse { + chunks: Vec::new(), + next_seq: 1, + exited: false, + exit_code: None, + closed: false, + failure: None, + sandbox_denied: false, + })?, + }), + ) + .await?; + + let read_task = tokio::spawn(async move { + process.read(Some(0), /*max_bytes*/ None, Some(0)).await + }); + let request = receive_jsonrpc(&mut replacement_socket).await?; + let JSONRPCMessage::Request(JSONRPCRequest { + id, method, params, .. + }) = request + else { + anyhow::bail!("expected existing process read request, got {request:?}"); + }; + assert_eq!(method, "process/read"); + assert_eq!( + serde_json::from_value::(params.context("read params should exist")?)?, + ReadParams { + process_id, + after_seq: Some(0), + max_bytes: None, + wait_ms: Some(0), + } + ); + let response = ReadResponse { + chunks: Vec::new(), + next_seq: 1, + exited: false, + exit_code: None, + closed: false, + failure: None, + sandbox_denied: false, + }; + send_jsonrpc( + &mut replacement_socket, + JSONRPCMessage::Response(JSONRPCResponse { + id, + result: serde_json::to_value(&response)?, + }), + ) + .await?; + assert_eq!(timeout(TEST_TIMEOUT, read_task).await???, response); + + server_task.abort(); + let _ = server_task.await; + Ok(()) +} + +async fn start_acceptor() -> Result<( + String, + mpsc::UnboundedReceiver, + JoinHandle<()>, +)> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let local_addr = listener.local_addr()?; + let (accepted_tx, accepted_rx) = mpsc::unbounded_channel(); + let app = Router::new() + .route("/", any(accept_websocket)) + .with_state(accepted_tx); + let server_task = tokio::spawn(async move { + let result = axum::serve(listener, app).await; + assert!( + result.is_ok(), + "accepted websocket test server should run: {result:?}" + ); + }); + Ok((format!("ws://{local_addr}/"), accepted_rx, server_task)) +} + +async fn accept_websocket( + websocket: WebSocketUpgrade, + State(accepted_tx): State>, +) -> impl IntoResponse { + websocket.on_upgrade(move |websocket| async move { + let _ = accepted_tx.send(websocket); + }) +} + +async fn connect_executor( + websocket_url: &str, + accepted_sockets: &mut mpsc::UnboundedReceiver, + session_id: &str, +) -> Result<( + WebSocketStream>, + EnvironmentManager, +)> { + let (mut websocket, _) = connect_async(websocket_url).await?; + let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv()) + .await? + .context("accepted websocket channel should remain open")?; + let manager_task = tokio::spawn(EnvironmentManager::from_accepted_websocket( + "environment-1".to_string(), + accepted_websocket, + accepted_options(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )); + complete_initialize(&mut websocket, session_id, /*resume_session_id*/ None).await?; + let manager = timeout(TEST_TIMEOUT, manager_task).await???; + Ok((websocket, manager)) +} + +async fn connect_replacement_executor( + websocket_url: &str, + accepted_sockets: &mut mpsc::UnboundedReceiver, +) -> Result<( + WebSocketStream>, + AcceptedSocket, +)> { + let (websocket, _) = connect_async(websocket_url).await?; + let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv()) + .await? + .context("accepted websocket channel should remain open")?; + Ok((websocket, accepted_websocket)) +} + +fn accepted_options() -> ExecServerClientConnectOptions { + ExecServerClientConnectOptions { + client_name: "host-test".to_string(), + initialize_timeout: TEST_TIMEOUT, + resume_session_id: None, + } +} + +async fn complete_initialize( + websocket: &mut WebSocketStream, + session_id: &str, + resume_session_id: Option<&str>, +) -> Result<()> +where + S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin, +{ + let initialize = receive_jsonrpc(&mut *websocket).await?; + let JSONRPCMessage::Request(JSONRPCRequest { + id, method, params, .. + }) = initialize + else { + anyhow::bail!("expected initialize request, got {initialize:?}"); + }; + assert_eq!(method, "initialize"); + assert_eq!( + serde_json::from_value::( + params.context("initialize request should contain params")? + )?, + InitializeParams { + client_name: "host-test".to_string(), + resume_session_id: resume_session_id.map(str::to_string), + } + ); + send_jsonrpc( + &mut *websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id, + result: serde_json::to_value(InitializeResponse { + session_id: session_id.to_string(), + environment_info: Some(EnvironmentInfo::local()), + })?, + }), + ) + .await?; + let initialized = receive_jsonrpc(&mut *websocket).await?; + assert_eq!( + initialized, + JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: Some(serde_json::json!({})), + }) + ); + Ok(()) +} + +async fn send_jsonrpc(websocket: &mut WebSocketStream, message: JSONRPCMessage) -> Result<()> +where + S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin, +{ + websocket + .send(Message::Text(serde_json::to_string(&message)?.into())) + .await?; + Ok(()) +} + +async fn receive_jsonrpc(websocket: &mut WebSocketStream) -> Result +where + S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin, +{ + loop { + let message = websocket + .next() + .await + .context("accepted websocket should remain open")??; + if let Message::Text(text) = message { + return Ok(serde_json::from_str(&text)?); + } + } +} diff --git a/codex-rs/exec-server/tests/capability_discovery.rs b/codex-rs/exec-server/tests/capability_discovery.rs new file mode 100644 index 0000000000000000000000000000000000000000..0be51f9fbb28dac284cb579aa3cd251ddf207af2 --- /dev/null +++ b/codex-rs/exec-server/tests/capability_discovery.rs @@ -0,0 +1,554 @@ +mod common; +#[cfg(target_os = "linux")] +#[path = "common/fake_bwrap.rs"] +mod fake_bwrap; + +#[cfg(target_os = "linux")] +use anyhow::Context as _; +use codex_exec_server::CAPABILITY_ROOTS_DISCOVER_METHOD; +use codex_exec_server::CapabilityRootDiscovery; +use codex_exec_server::CapabilityRootsDiscoverParams; +use codex_exec_server::CapabilityRootsDiscoverResponse; +use codex_exec_server::FileSystemSandboxContext; +use codex_exec_server::InitializeParams; +use codex_exec_server::InitializeResponse; +use codex_exec_server::WindowsSandboxSelection; +use codex_exec_server_protocol::CapabilityRootDiscoverRequest; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_protocol::models::PermissionProfile; +use codex_protocol::permissions::FileSystemAccessMode; +use codex_protocol::permissions::FileSystemSandboxEntry; +use codex_protocol::permissions::FileSystemSandboxPolicy; +use codex_protocol::permissions::NetworkSandboxPolicy; +use codex_utils_absolute_path::AbsolutePathBuf; +use codex_utils_path_uri::PathUri; +use common::exec_server::exec_server; +#[cfg(target_os = "linux")] +use common::exec_server::exec_server_with_env; +#[cfg(target_os = "linux")] +use fake_bwrap::write_fake_bwrap; +use pretty_assertions::assert_eq; + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn discovers_a_complete_capability_bundle_in_one_request() -> anyhow::Result<()> { + let root = tempfile::tempdir()?; + write_file( + &root.path().join(".codex-plugin/plugin.json"), + r#"{ + "name": "demo", + "interface": {"displayName": "Demo Plugin"}, + "mcpServers": "./config/mcp.json", + "apps": "./config/apps.json" +}"#, + )?; + write_file( + &root.path().join(".claude-plugin/plugin.json"), + r#"{"name":"lower-priority-claude"}"#, + )?; + write_file( + &root.path().join(".cursor-plugin/plugin.json"), + r#"{"name":"lower-priority-cursor"}"#, + )?; + write_file( + &root.path().join("config/mcp.json"), + r#"{"mcpServers":{"demo":{"command":"demo-server"}}}"#, + )?; + write_file( + &root.path().join("config/apps.json"), + r#"{"apps":{"demo":{"connector_id":"connector-demo"}}}"#, + )?; + write_file( + &root.path().join("skills/deploy/SKILL.md"), + "---\nname: deploy\ndescription: Deploy the service.\n---\n\nDeploy instructions.\n", + )?; + write_file( + &root.path().join("skills/deploy/agents/openai.yaml"), + "policy:\n allow_implicit_invocation: false\n", + )?; + write_file( + &root.path().join("nested/.claude-plugin/plugin.json"), + r#"{"name":"nested"}"#, + )?; + write_file( + &root.path().join("nested/skills/audit/SKILL.md"), + "---\nname: audit\ndescription: Audit the service.\n---\n", + )?; + write_file( + &root.path().join("nested-cursor/.cursor-plugin/plugin.json"), + r#"{"name":"cursor-nested"}"#, + )?; + write_file( + &root.path().join("nested-cursor/skills/review/SKILL.md"), + "---\nname: review\ndescription: Review the service.\n---\n", + )?; + + let mut server = exec_server().await?; + initialize(&mut server).await?; + let root_uri = PathUri::from_host_native_path(root.path())?; + let discovery = discover_root(&mut server, "demo@1", root_uri.clone()).await?; + + assert_eq!(discovery.id, "demo@1"); + assert_eq!(discovery.path, root_uri); + assert_eq!(discovery.error, None); + assert_eq!(discovery.warnings, Vec::::new()); + let plugin = discovery.plugin.as_ref().expect("root plugin"); + assert_eq!( + plugin.manifest.path, + root_uri.join(".codex-plugin/plugin.json")? + ); + assert!(plugin.manifest.contents.contains("Demo Plugin")); + assert_eq!( + plugin.mcp_config.as_ref().map(|file| &file.path), + Some(&root_uri.join("config/mcp.json")?) + ); + assert_eq!( + plugin.apps_config.as_ref().map(|file| &file.path), + Some(&root_uri.join("config/apps.json")?) + ); + assert_eq!( + discovery + .namespace_manifests + .iter() + .map(|file| file.path.clone()) + .collect::>(), + vec![ + root_uri.join(".codex-plugin/plugin.json")?, + root_uri.join("nested/.claude-plugin/plugin.json")?, + root_uri.join("nested-cursor/.cursor-plugin/plugin.json")?, + ] + ); + assert_eq!( + discovery + .skills + .iter() + .map(|skill| ( + skill.instructions.path.clone(), + skill + .metadata + .as_ref() + .map(|metadata| metadata.path.clone()), + )) + .collect::>(), + vec![ + (root_uri.join("nested-cursor/skills/review/SKILL.md")?, None,), + (root_uri.join("nested/skills/audit/SKILL.md")?, None,), + ( + root_uri.join("skills/deploy/SKILL.md")?, + Some(root_uri.join("skills/deploy/agents/openai.yaml")?), + ), + ] + ); + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn discovers_cursor_plugin_without_reading_default_mcp_for_inline_servers() +-> anyhow::Result<()> { + let root = tempfile::tempdir()?; + write_file( + &root.path().join(".cursor-plugin/plugin.json"), + r#"{"name":"cursor-demo","mcpServers":{"inline":{"command":"inline"}}}"#, + )?; + write_file( + &root.path().join(".mcp.json"), + r#"{"mcpServers":{"should-not-load":{"command":"wrong"}}}"#, + )?; + + let mut server = exec_server().await?; + initialize(&mut server).await?; + let root_uri = PathUri::from_host_native_path(root.path())?; + let discovery = discover_root(&mut server, "cursor@1", root_uri.clone()).await?; + + assert_eq!(discovery.error, None); + assert_eq!(discovery.warnings, Vec::::new()); + let plugin = discovery.plugin.expect("cursor plugin"); + assert_eq!( + plugin.manifest.path, + root_uri.join(".cursor-plugin/plugin.json")? + ); + assert_eq!(plugin.mcp_config, None); + assert_eq!( + discovery + .namespace_manifests + .iter() + .map(|manifest| manifest.path.clone()) + .collect::>(), + vec![root_uri.join(".cursor-plugin/plugin.json")?] + ); + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn sandboxed_discovery_batches_roots_without_combining_different_permissions() +-> anyhow::Result<()> { + #[cfg(windows)] + crate::skip_if_mxc_unavailable!(Ok(())); + let workspace = tempfile::tempdir()?; + let first_root = workspace.path().join("first"); + let second_root = workspace.path().join("second"); + write_file( + &first_root.join("skills/first/SKILL.md"), + "---\nname: first\ndescription: First skill.\n---\n", + )?; + write_file( + &second_root.join("skills/second/SKILL.md"), + "---\nname: second\ndescription: Second skill.\n---\n", + )?; + + let workspace_uri = PathUri::from_host_native_path(workspace.path())?; + let first_uri = PathUri::from_host_native_path(&first_root)?; + let second_uri = PathUri::from_host_native_path(&second_root)?; + let read_workspace = FileSystemSandboxEntry::new( + AbsolutePathBuf::from_absolute_path(workspace.path())?.into(), + FileSystemAccessMode::Read, + ); + let policy = FileSystemSandboxPolicy::restricted(vec![read_workspace]); + let shared_sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + workspace_uri, + ); + let shared_sandbox = with_native_sandbox(shared_sandbox); + + #[cfg(target_os = "linux")] + let fake_bwrap_directory = tempfile::tempdir()?; + #[cfg(target_os = "linux")] + let (mut server, fake_bwrap) = { + let fake_bin_dir = fake_bwrap_directory.path().to_path_buf(); + let fake_bwrap = write_fake_bwrap(&fake_bin_dir)?; + let mut path_entries = vec![fake_bin_dir]; + if let Some(path) = std::env::var_os("PATH") { + path_entries.extend(std::env::split_paths(&path)); + } + let helper_path = std::env::join_paths(path_entries)?; + ( + exec_server_with_env([("PATH", helper_path.as_os_str())], &[]).await?, + fake_bwrap, + ) + }; + #[cfg(not(target_os = "linux"))] + let mut server = exec_server().await?; + initialize(&mut server).await?; + let response = discover_roots( + &mut server, + vec![ + CapabilityRootDiscoverRequest { + id: "first".to_string(), + path: first_uri.clone(), + sandbox: Some(shared_sandbox.clone()), + }, + CapabilityRootDiscoverRequest { + id: "second".to_string(), + path: second_uri.clone(), + sandbox: Some(shared_sandbox.clone()), + }, + ], + ) + .await?; + assert_eq!( + response + .roots + .into_iter() + .map(|root| ( + root.id, + root.path, + root.skills + .into_iter() + .map(|skill| skill.instructions.path) + .collect::>(), + root.error, + )) + .collect::>(), + vec![ + ( + "first".to_string(), + first_uri.clone(), + vec![first_uri.join("skills/first/SKILL.md")?], + None, + ), + ( + "second".to_string(), + second_uri.clone(), + vec![second_uri.join("skills/second/SKILL.md")?], + None, + ), + ] + ); + + #[cfg(target_os = "linux")] + { + let launch_log = fake_bwrap.with_file_name("bwrap.log"); + let launch_count = std::fs::read_to_string(&launch_log) + .with_context(|| format!("expected fake bwrap launch log at {}", launch_log.display()))? + .lines() + .count(); + assert_eq!(launch_count, 1); + + std::fs::write(fake_bwrap.with_file_name("bwrap.fail-once"), "")?; + let fallback = discover_roots( + &mut server, + vec![ + CapabilityRootDiscoverRequest { + id: "fallback-first".to_string(), + path: first_uri.clone(), + sandbox: Some(shared_sandbox.clone()), + }, + CapabilityRootDiscoverRequest { + id: "fallback-second".to_string(), + path: second_uri.clone(), + sandbox: Some(shared_sandbox.clone()), + }, + ], + ) + .await?; + assert_eq!( + fallback + .roots + .into_iter() + .map(|root| (root.id, root.skills.len(), root.error)) + .collect::>(), + vec![ + ("fallback-first".to_string(), 1, None), + ("fallback-second".to_string(), 1, None), + ] + ); + assert!(std::fs::read_to_string(&launch_log)?.lines().count() > 2); + + server.shutdown().await?; + server = exec_server().await?; + initialize(&mut server).await?; + } + + let read_first_root = FileSystemSandboxEntry::new( + AbsolutePathBuf::from_absolute_path(&first_root)?.into(), + FileSystemAccessMode::Read, + ); + let first_policy = FileSystemSandboxPolicy::restricted(vec![read_first_root]); + let first_only_sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions( + &first_policy, + NetworkSandboxPolicy::Restricted, + ), + first_uri.clone(), + ); + let first_only_sandbox = with_native_sandbox(first_only_sandbox); + let read_second_root = FileSystemSandboxEntry::new( + AbsolutePathBuf::from_absolute_path(&second_root)?.into(), + FileSystemAccessMode::Read, + ); + let second_policy = FileSystemSandboxPolicy::restricted(vec![read_second_root]); + let second_only_sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions( + &second_policy, + NetworkSandboxPolicy::Restricted, + ), + second_uri.clone(), + ); + let second_only_sandbox = with_native_sandbox(second_only_sandbox); + let response = discover_roots( + &mut server, + vec![ + CapabilityRootDiscoverRequest { + id: "first-isolated".to_string(), + path: first_uri, + sandbox: Some(first_only_sandbox), + }, + CapabilityRootDiscoverRequest { + id: "second-isolated".to_string(), + path: second_uri, + sandbox: Some(second_only_sandbox), + }, + ], + ) + .await?; + assert_eq!( + response + .roots + .into_iter() + .map(|root| (root.id, root.skills.len(), root.error)) + .collect::>(), + vec![ + ("first-isolated".to_string(), 1, None), + ("second-isolated".to_string(), 1, None), + ] + ); + + server.shutdown().await?; + Ok(()) +} + +#[cfg(unix)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn sandboxed_discovery_follows_only_permitted_external_symlinks() -> anyhow::Result<()> { + let root = tempfile::tempdir()?; + let external = tempfile::tempdir()?; + write_file( + &root.path().join(".codex-plugin/plugin.json"), + r#"{"name":"linked-plugin","mcpServers":"./external-mcp.json"}"#, + )?; + write_file( + &external.path().join("mcp.json"), + r#"{"mcpServers":{"linked":{"command":"linked-server"}}}"#, + )?; + write_file( + &external.path().join("skill/SKILL.md"), + "---\nname: linked\ndescription: Linked external skill.\n---\n", + )?; + std::fs::create_dir_all(root.path().join("skills"))?; + std::os::unix::fs::symlink( + external.path().join("skill"), + root.path().join("skills/linked"), + )?; + std::os::unix::fs::symlink( + external.path().join("mcp.json"), + root.path().join("external-mcp.json"), + )?; + + let mut server = exec_server().await?; + initialize(&mut server).await?; + let root_uri = PathUri::from_host_native_path(root.path())?; + let root_path = AbsolutePathBuf::from_absolute_path(root.path())?; + let external_root = AbsolutePathBuf::from_absolute_path(external.path())?; + let path_entry = + |path: AbsolutePathBuf, access| FileSystemSandboxEntry::new(path.into(), access); + let read_root = path_entry(root_path, FileSystemAccessMode::Read); + let read_external = path_entry(external_root.clone(), FileSystemAccessMode::Read); + let deny_external_skill = path_entry(external_root.join("skill"), FileSystemAccessMode::Deny); + let cases = [ + ( + "permitted symlinks", + vec![read_root.clone(), read_external.clone()], + true, + true, + ), + ( + "denied external root", + vec![read_root.clone()], + false, + false, + ), + ( + "denied external skill", + vec![read_root, read_external, deny_external_skill], + false, + true, + ), + ]; + + for (scenario, entries, has_skill, has_mcp) in cases { + let policy = FileSystemSandboxPolicy::restricted(entries); + let sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + root_uri.clone(), + ); + let discovery = + discover_root_with_sandbox(&mut server, "linked", root_uri.clone(), Some(sandbox)) + .await?; + + assert_eq!(discovery.error, None, "{scenario}"); + assert_eq!(discovery.skills.len(), usize::from(has_skill), "{scenario}"); + assert_eq!( + discovery + .plugin + .and_then(|plugin| plugin.mcp_config) + .is_some_and(|config| config.contents.contains("linked-server")), + has_mcp, + "{scenario}" + ); + } + + server.shutdown().await?; + Ok(()) +} + +async fn discover_root( + server: &mut common::exec_server::ExecServerHarness, + id: &str, + path: PathUri, +) -> anyhow::Result { + discover_root_with_sandbox(server, id, path, /*sandbox*/ None).await +} + +async fn discover_root_with_sandbox( + server: &mut common::exec_server::ExecServerHarness, + id: &str, + path: PathUri, + sandbox: Option, +) -> anyhow::Result { + let response = discover_roots( + server, + vec![CapabilityRootDiscoverRequest { + id: id.to_string(), + path, + sandbox, + }], + ) + .await?; + let [discovery] = response.roots.as_slice() else { + anyhow::bail!("expected exactly one discovered root"); + }; + Ok(discovery.clone()) +} + +async fn discover_roots( + server: &mut common::exec_server::ExecServerHarness, + roots: Vec, +) -> anyhow::Result { + let request_id = server + .send_request( + CAPABILITY_ROOTS_DISCOVER_METHOD, + serde_json::to_value(CapabilityRootsDiscoverParams { roots })?, + ) + .await?; + let response = server.next_event().await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = response else { + anyhow::bail!("expected discovery response, received {response:?}"); + }; + assert_eq!(id, request_id); + Ok(serde_json::from_value(result)?) +} + +async fn initialize(server: &mut common::exec_server::ExecServerHarness) -> anyhow::Result<()> { + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "capability-discovery-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!(event, JSONRPCMessage::Response(response) if response.id == initialize_id) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { result, .. }) = response else { + unreachable!("wait predicate only accepts a response"); + }; + let _: InitializeResponse = serde_json::from_value(result)?; + server + .send_notification("initialized", serde_json::json!({})) + .await?; + Ok(()) +} + +fn write_file(path: &std::path::Path, contents: &str) -> anyhow::Result<()> { + let parent = path + .parent() + .ok_or_else(|| anyhow::anyhow!("test file should have a parent"))?; + std::fs::create_dir_all(parent)?; + std::fs::write(path, contents)?; + Ok(()) +} + +fn with_native_sandbox(mut sandbox: FileSystemSandboxContext) -> FileSystemSandboxContext { + if cfg!(windows) { + sandbox.windows_sandbox_selection = WindowsSandboxSelection::Mxc; + } + sandbox +} diff --git a/codex-rs/exec-server/tests/chatgpt_cloudflare_affinity.rs b/codex-rs/exec-server/tests/chatgpt_cloudflare_affinity.rs new file mode 100644 index 0000000000000000000000000000000000000000..d754fa8c4b7fec9cc4b7413896cca3d57c11141a --- /dev/null +++ b/codex-rs/exec-server/tests/chatgpt_cloudflare_affinity.rs @@ -0,0 +1,405 @@ +#![cfg(unix)] + +mod common; + +use std::collections::BTreeMap; +use std::ffi::OsString; +use std::fs; +use std::io; +use std::io::Read; +use std::io::Write; +use std::net::TcpListener; +use std::net::TcpStream; +use std::sync::Arc; +use std::sync::mpsc; +use std::thread; +use std::time::Duration; + +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::InitializeParams; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::RequestId; +use common::exec_server::ExecServerHarness; +use common::exec_server::exec_server_with_env; +use pretty_assertions::assert_eq; +use rcgen::BasicConstraints; +use rcgen::CertificateParams; +use rcgen::CertifiedIssuer; +use rcgen::DistinguishedName; +use rcgen::DnType; +use rcgen::ExtendedKeyUsagePurpose; +use rcgen::IsCa; +use rcgen::KeyPair; +use rcgen::KeyUsagePurpose; +use rcgen::PKCS_ECDSA_P256_SHA256; +use rustls::pki_types::CertificateDer; +use rustls::pki_types::PrivateKeyDer; +use serde::de::DeserializeOwned; +use serde_json::Value; +use tempfile::TempDir; + +const CHATGPT_MCP_URL: &str = "https://chatgpt.com/backend-api/ps/mcp"; +const NON_CHATGPT_MCP_URL: &str = "https://api.openai.com/backend-api/ps/mcp"; + +#[derive(Debug)] +struct CapturedRequest { + connect_authority: String, + request_line: String, + headers: BTreeMap>, +} + +struct TlsMaterial { + ca_cert_pem: String, + server_cert: CertificateDer<'static>, + server_key: PrivateKeyDer<'static>, +} + +struct TlsInterceptingProxy { + ca_cert_pem: String, + request_rx: mpsc::Receiver>, + thread: thread::JoinHandle>, + url: String, +} + +/// Exercises the same `http/request` route used by remotely executed Streamable HTTP MCP calls. +/// Each RPC uses the shared route-aware client. The first response sets `__cflb`, and the second response +/// replaces it, proving cross-client persistence through the shared cookie store. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_replays_only_chatgpt_cloudflare_cookies() -> anyhow::Result<()> { + let proxy = TlsInterceptingProxy::start(/*expected_requests*/ 4)?; + let temp_dir = TempDir::new()?; + let ca_path = temp_dir.path().join("cloudflare-affinity-test-ca.pem"); + fs::write(&ca_path, &proxy.ca_cert_pem)?; + let proxy_url = OsString::from(&proxy.url); + let empty = OsString::new(); + let env = vec![ + ( + OsString::from("CODEX_CA_CERTIFICATE"), + ca_path.as_os_str().to_owned(), + ), + (OsString::from("HTTPS_PROXY"), proxy_url.clone()), + (OsString::from("https_proxy"), proxy_url.clone()), + (OsString::from("ALL_PROXY"), proxy_url.clone()), + (OsString::from("all_proxy"), proxy_url), + (OsString::from("NO_PROXY"), empty.clone()), + (OsString::from("no_proxy"), empty), + ]; + let mut server = exec_server_with_env(env, &[]).await?; + initialize_exec_server(&mut server).await?; + + let first_response = execute_http_request(&mut server, CHATGPT_MCP_URL, "first").await?; + assert_eq!( + (first_response.status, first_response.body.into_inner()), + (200, b"ok".to_vec()) + ); + let first = proxy.next_request()?; + assert_eq!( + ( + first.connect_authority.as_str(), + first.request_line.as_str(), + first.headers.get("cookie"), + ), + ("chatgpt.com:443", "POST /backend-api/ps/mcp HTTP/1.1", None,) + ); + + let west_response = execute_http_request(&mut server, CHATGPT_MCP_URL, "west").await?; + assert_eq!(west_response.status, 200); + let request_with_west_affinity = proxy.next_request()?; + assert_eq!( + request_with_west_affinity + .headers + .get("cookie") + .cloned() + .unwrap_or_default(), + vec!["__cflb=west".to_string()] + ); + + let central_response = execute_http_request(&mut server, CHATGPT_MCP_URL, "central").await?; + assert_eq!(central_response.status, 200); + let request_with_central_affinity = proxy.next_request()?; + assert_eq!( + ( + request_with_central_affinity.request_line.as_str(), + request_with_central_affinity + .headers + .get("cookie") + .cloned() + .unwrap_or_default(), + ), + ( + "POST /backend-api/ps/mcp HTTP/1.1", + vec!["__cflb=central".to_string()], + ) + ); + let other_host_response = + execute_http_request(&mut server, NON_CHATGPT_MCP_URL, "other-host").await?; + assert_eq!(other_host_response.status, 200); + let other_host = proxy.next_request()?; + assert_eq!( + ( + other_host.connect_authority.as_str(), + other_host.request_line.as_str(), + other_host.headers.get("cookie"), + ), + ( + "api.openai.com:443", + "POST /backend-api/ps/mcp HTTP/1.1", + None, + ) + ); + + server.shutdown().await?; + proxy.finish()?; + Ok(()) +} + +impl TlsInterceptingProxy { + fn start(expected_requests: usize) -> anyhow::Result { + codex_utils_rustls_provider::ensure_rustls_crypto_provider(); + let material = generate_tls_material()?; + let listener = TcpListener::bind(("127.0.0.1", 0))?; + let address = listener.local_addr()?; + let config = Arc::new( + rustls::ServerConfig::builder_with_protocol_versions(&[&rustls::version::TLS13]) + .with_no_client_auth() + .with_single_cert(vec![material.server_cert], material.server_key)?, + ); + let (request_tx, request_rx) = mpsc::channel(); + let thread = thread::spawn(move || { + run_tls_intercepting_proxy(listener, config, request_tx, expected_requests) + .map_err(|error| error.to_string()) + }); + + Ok(Self { + ca_cert_pem: material.ca_cert_pem, + request_rx, + thread, + url: format!("http://{address}"), + }) + } + + fn next_request(&self) -> anyhow::Result { + self.request_rx + .recv_timeout(Duration::from_secs(5)) + .map_err(anyhow::Error::from)? + .map_err(anyhow::Error::msg) + } + + fn finish(self) -> anyhow::Result<()> { + self.thread + .join() + .map_err(|_| anyhow::anyhow!("TLS proxy thread panicked"))? + .map_err(anyhow::Error::msg) + } +} + +fn generate_tls_material() -> anyhow::Result { + let mut ca_params = CertificateParams::default(); + ca_params.is_ca = IsCa::Ca(BasicConstraints::Unconstrained); + ca_params.key_usages = vec![KeyUsagePurpose::KeyCertSign, KeyUsagePurpose::CrlSign]; + let mut ca_distinguished_name = DistinguishedName::new(); + ca_distinguished_name.push(DnType::CommonName, "Codex affinity test CA"); + ca_params.distinguished_name = ca_distinguished_name; + let ca_key_pair = KeyPair::generate_for(&PKCS_ECDSA_P256_SHA256)?; + let ca = CertifiedIssuer::self_signed(ca_params, ca_key_pair)?; + + let mut server_params = CertificateParams::new(vec![ + "chatgpt.com".to_string(), + "api.openai.com".to_string(), + ])?; + server_params.extended_key_usages = vec![ExtendedKeyUsagePurpose::ServerAuth]; + server_params.key_usages = vec![ + KeyUsagePurpose::DigitalSignature, + KeyUsagePurpose::KeyEncipherment, + ]; + let server_key_pair = KeyPair::generate_for(&PKCS_ECDSA_P256_SHA256)?; + let server_cert = server_params.signed_by(&server_key_pair, &ca)?; + + Ok(TlsMaterial { + ca_cert_pem: ca.pem(), + server_cert: server_cert.der().clone(), + server_key: PrivateKeyDer::from(server_key_pair), + }) +} + +fn run_tls_intercepting_proxy( + listener: TcpListener, + config: Arc, + request_tx: mpsc::Sender>, + expected_requests: usize, +) -> io::Result<()> { + for request_index in 0..expected_requests { + let (mut stream, _) = listener.accept()?; + configure_stream(&stream)?; + let connect_head = read_http_head(&mut stream)?; + let connect_authority = connect_head + .lines() + .next() + .and_then(|line| line.split_whitespace().nth(1)) + .ok_or_else(|| io::Error::other("CONNECT request line should include an authority"))? + .to_string(); + stream.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")?; + stream.flush()?; + + let connection = + rustls::ServerConnection::new(Arc::clone(&config)).map_err(io::Error::other)?; + let mut tls = rustls::StreamOwned::new(connection, stream); + let request = capture_http_request(&mut tls, connect_authority); + match request { + Ok(request) => request_tx + .send(Ok(request)) + .map_err(|_| io::Error::other("request receiver was dropped"))?, + Err(error) => { + let message = error.to_string(); + let _ = request_tx.send(Err(message)); + return Err(error); + } + } + + let response = match request_index { + 0 => concat!( + "HTTP/1.1 200 OK\r\n", + "content-length: 2\r\n", + "connection: close\r\n", + "set-cookie: __cflb=west; Path=/; Secure; HttpOnly\r\n", + "set-cookie: chatgpt_session=secret; Path=/; Secure; HttpOnly\r\n", + "\r\n", + "ok", + ), + 1 => concat!( + "HTTP/1.1 200 OK\r\n", + "content-length: 2\r\n", + "connection: close\r\n", + "set-cookie: __cflb=central; Path=/; Secure; HttpOnly\r\n", + "set-cookie: chatgpt_session=still-secret; Path=/; Secure; HttpOnly\r\n", + "\r\n", + "ok", + ), + _ => concat!( + "HTTP/1.1 200 OK\r\n", + "content-length: 2\r\n", + "connection: close\r\n", + "\r\n", + "ok", + ), + }; + tls.write_all(response.as_bytes())?; + tls.flush()?; + } + Ok(()) +} + +fn configure_stream(stream: &TcpStream) -> io::Result<()> { + stream.set_read_timeout(Some(Duration::from_secs(5)))?; + stream.set_write_timeout(Some(Duration::from_secs(5))) +} + +fn capture_http_request( + stream: &mut impl Read, + connect_authority: String, +) -> io::Result { + let request_head = read_http_head(stream)?; + let mut lines = request_head.lines(); + let request_line = lines + .next() + .ok_or_else(|| io::Error::other("HTTP request should include a request line"))? + .to_string(); + let mut headers: BTreeMap> = BTreeMap::new(); + for line in lines.filter(|line| !line.is_empty()) { + let (name, value) = line + .split_once(':') + .ok_or_else(|| io::Error::other(format!("invalid HTTP header: {line}")))?; + headers + .entry(name.to_ascii_lowercase()) + .or_default() + .push(value.trim().to_string()); + } + Ok(CapturedRequest { + connect_authority, + request_line, + headers, + }) +} + +fn read_http_head(stream: &mut impl Read) -> io::Result { + const MAX_HEADER_BYTES: usize = 64 * 1024; + let mut bytes = Vec::new(); + while !bytes.ends_with(b"\r\n\r\n") { + if bytes.len() == MAX_HEADER_BYTES { + return Err(io::Error::other("HTTP headers exceeded test limit")); + } + let mut byte = [0]; + stream.read_exact(&mut byte)?; + bytes.push(byte[0]); + } + String::from_utf8(bytes).map_err(io::Error::other) +} + +async fn initialize_exec_server(server: &mut ExecServerHarness) -> anyhow::Result<()> { + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "cloudflare-affinity-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + let _: Value = wait_for_response(server, initialize_id).await?; + server + .send_notification("initialized", serde_json::json!({})) + .await +} + +async fn execute_http_request( + server: &mut ExecServerHarness, + url: &str, + request_id: &str, +) -> anyhow::Result { + let response_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "POST".to_string(), + url: url.to_string(), + headers: Vec::new(), + body: None, + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: request_id.to_string(), + stream_response: false, + })?, + ) + .await?; + wait_for_response(server, response_id).await +} + +async fn wait_for_response( + server: &mut ExecServerHarness, + request_id: RequestId, +) -> anyhow::Result +where + T: DeserializeOwned, +{ + let message = server + .wait_for_event(|event| match event { + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) + | JSONRPCMessage::Error(codex_exec_server_protocol::JSONRPCError { id, .. }) => { + id == &request_id + } + _ => false, + }) + .await?; + match message { + JSONRPCMessage::Response(JSONRPCResponse { result, .. }) => { + Ok(serde_json::from_value(result)?) + } + JSONRPCMessage::Error(error) => { + anyhow::bail!("exec-server returned an error for {request_id:?}: {error:?}") + } + _ => unreachable!("predicate only accepts responses for the requested id"), + } +} diff --git a/codex-rs/exec-server/tests/common/exec_server.rs b/codex-rs/exec-server/tests/common/exec_server.rs new file mode 100644 index 0000000000000000000000000000000000000000..d0fc3892fc800c4ef81fc29fc504e5837903f0e6 --- /dev/null +++ b/codex-rs/exec-server/tests/common/exec_server.rs @@ -0,0 +1,413 @@ +#![allow(dead_code)] + +use std::path::PathBuf; +use std::process::Stdio; +use std::time::Duration; + +use anyhow::anyhow; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCNotification; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::RequestId; +use futures::SinkExt; +use futures::StreamExt; +use tempfile::TempDir; +use tokio::io::AsyncBufReadExt; +use tokio::io::BufReader; +use tokio::io::copy_bidirectional; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::process::Child; +use tokio::process::Command; +use tokio::sync::oneshot; +use tokio::task::JoinHandle; +use tokio::time::Instant; +use tokio::time::sleep; +use tokio::time::timeout; +use tokio_tungstenite::connect_async; +use tokio_tungstenite::tungstenite::Message; + +const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); +const CONNECT_RETRY_INTERVAL: Duration = Duration::from_millis(25); +const EVENT_TIMEOUT: Duration = Duration::from_secs(5); + +pub(crate) struct ExecServerHarness { + codex_home: TempDir, + child: Child, + websocket_url: String, + websocket: tokio_tungstenite::WebSocketStream< + tokio_tungstenite::MaybeTlsStream, + >, + next_request_id: i64, +} + +impl Drop for ExecServerHarness { + fn drop(&mut self) { + let _ = self.child.start_kill(); + } +} + +pub(crate) struct TestCodexHelperPaths { + pub(crate) codex_exe: PathBuf, + pub(crate) codex_linux_sandbox_exe: Option, +} + +pub(crate) struct DisconnectableWebSocketProxy { + websocket_url: String, + pause_tx: Option>, + blocked_connection_rx: Option>, + resume_tx: Option>, + task: JoinHandle<()>, +} + +impl Drop for DisconnectableWebSocketProxy { + fn drop(&mut self) { + self.task.abort(); + } +} + +pub(crate) fn test_codex_helper_paths() -> anyhow::Result { + let (helper_binary, codex_linux_sandbox_exe) = super::current_test_binary_helper_paths()?; + Ok(TestCodexHelperPaths { + codex_exe: helper_binary, + codex_linux_sandbox_exe, + }) +} + +pub(crate) async fn exec_server() -> anyhow::Result { + exec_server_with_env(std::iter::empty::<(&str, &str)>(), &[]).await +} + +pub(crate) async fn exec_server_with_env( + env: I, + args: &[&str], +) -> anyhow::Result +where + I: IntoIterator, + K: AsRef, + V: AsRef, +{ + let helper_paths = test_codex_helper_paths()?; + let mut child = Command::new(&helper_paths.codex_exe); + child.args(["exec-server", "--listen", "ws://127.0.0.1:0"]); + child.args(args); + child.envs(env); + ExecServerHarness::start(child).await +} + +impl ExecServerHarness { + pub(crate) async fn start(mut command: Command) -> anyhow::Result { + let codex_home = TempDir::new()?; + command.stdin(Stdio::null()); + command.stdout(Stdio::piped()); + command.stderr(Stdio::inherit()); + command.kill_on_drop(true); + if !command + .as_std() + .get_envs() + .any(|(key, value)| key == "CODEX_HOME" && value.is_some()) + { + command.env("CODEX_HOME", codex_home.path()); + } + let mut child = command.spawn()?; + + let websocket_url = read_listen_url_from_stdout(&mut child).await?; + let (websocket, _) = connect_websocket_when_ready(&websocket_url).await?; + Ok(Self { + codex_home, + child, + websocket_url, + websocket, + next_request_id: 1, + }) + } + + pub(crate) fn codex_home(&self) -> &std::path::Path { + self.codex_home.path() + } + + pub(crate) fn websocket_url(&self) -> &str { + &self.websocket_url + } + + pub(crate) async fn disconnect_websocket(&mut self) -> anyhow::Result<()> { + self.websocket.close(None).await?; + Ok(()) + } + + pub(crate) async fn reconnect_websocket(&mut self) -> anyhow::Result<()> { + let (websocket, _) = connect_websocket_when_ready(&self.websocket_url).await?; + self.websocket = websocket; + Ok(()) + } + + pub(crate) async fn disconnectable_websocket_proxy( + &self, + ) -> anyhow::Result { + DisconnectableWebSocketProxy::new(&self.websocket_url).await + } + + pub(crate) async fn send_request( + &mut self, + method: &str, + params: serde_json::Value, + ) -> anyhow::Result { + let id = RequestId::Integer(self.next_request_id); + self.next_request_id += 1; + self.send_message(JSONRPCMessage::Request(JSONRPCRequest { + id: id.clone(), + method: method.to_string(), + params: Some(params), + trace: None, + })) + .await?; + Ok(id) + } + + pub(crate) async fn send_notification( + &mut self, + method: &str, + params: serde_json::Value, + ) -> anyhow::Result<()> { + self.send_message(JSONRPCMessage::Notification(JSONRPCNotification { + method: method.to_string(), + params: Some(params), + })) + .await + } + + pub(crate) async fn send_raw_text(&mut self, text: &str) -> anyhow::Result<()> { + self.websocket + .send(Message::Text(text.to_string().into())) + .await?; + Ok(()) + } + + pub(crate) async fn send_raw_binary(&mut self, bytes: Vec) -> anyhow::Result<()> { + self.websocket.send(Message::Binary(bytes.into())).await?; + Ok(()) + } + + pub(crate) async fn next_event(&mut self) -> anyhow::Result { + self.next_event_with_timeout(EVENT_TIMEOUT).await + } + + pub(crate) async fn wait_for_event( + &mut self, + mut predicate: F, + ) -> anyhow::Result + where + F: FnMut(&JSONRPCMessage) -> bool, + { + let deadline = Instant::now() + EVENT_TIMEOUT; + loop { + let now = Instant::now(); + if now >= deadline { + return Err(anyhow!( + "timed out waiting for matching exec-server event after {EVENT_TIMEOUT:?}" + )); + } + let remaining = deadline.duration_since(now); + let event = self.next_event_with_timeout(remaining).await?; + if predicate(&event) { + return Ok(event); + } + } + } + + pub(crate) async fn shutdown(&mut self) -> anyhow::Result<()> { + self.child.start_kill()?; + timeout(CONNECT_TIMEOUT, self.child.wait()) + .await + .map_err(|_| anyhow!("timed out waiting for exec-server shutdown"))??; + Ok(()) + } + + async fn send_message(&mut self, message: JSONRPCMessage) -> anyhow::Result<()> { + let encoded = serde_json::to_string(&message)?; + self.websocket.send(Message::Text(encoded.into())).await?; + Ok(()) + } + + async fn next_event_with_timeout( + &mut self, + timeout_duration: Duration, + ) -> anyhow::Result { + loop { + let frame = timeout(timeout_duration, self.websocket.next()) + .await + .map_err(|_| anyhow!("timed out waiting for exec-server websocket event"))? + .ok_or_else(|| anyhow!("exec-server websocket closed"))??; + + match frame { + Message::Text(text) => { + return Ok(serde_json::from_str(text.as_ref())?); + } + Message::Binary(bytes) => { + return Ok(serde_json::from_slice(bytes.as_ref())?); + } + Message::Close(_) => return Err(anyhow!("exec-server websocket closed")), + Message::Ping(_) | Message::Pong(_) => {} + _ => {} + } + } + } +} + +impl DisconnectableWebSocketProxy { + pub(crate) async fn new(websocket_url: &str) -> anyhow::Result { + let upstream = websocket_url + .strip_prefix("ws://") + .ok_or_else(|| anyhow!("exec-server websocket URL must use ws://"))? + .trim_end_matches('/') + .to_string(); + let listener = TcpListener::bind("127.0.0.1:0").await?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let (pause_tx, pause_rx) = oneshot::channel(); + let (blocked_connection_tx, blocked_connection_rx) = oneshot::channel(); + let (resume_tx, resume_rx) = oneshot::channel(); + let task = tokio::spawn(run_disconnectable_proxy( + listener, + upstream, + pause_rx, + blocked_connection_tx, + resume_rx, + )); + Ok(DisconnectableWebSocketProxy { + websocket_url, + pause_tx: Some(pause_tx), + blocked_connection_rx: Some(blocked_connection_rx), + resume_tx: Some(resume_tx), + task, + }) + } + + pub(crate) fn websocket_url(&self) -> &str { + &self.websocket_url + } + + pub(crate) async fn pause_and_disconnect(&mut self) -> anyhow::Result<()> { + self.pause_tx + .take() + .ok_or_else(|| anyhow!("disconnectable websocket proxy is already paused"))? + .send(()) + .map_err(|_| anyhow!("disconnectable websocket proxy stopped"))?; + let blocked_connection_rx = self + .blocked_connection_rx + .take() + .ok_or_else(|| anyhow!("disconnectable websocket proxy is already paused"))?; + timeout(CONNECT_TIMEOUT, blocked_connection_rx) + .await + .map_err(|_| anyhow!("timed out waiting for client reconnect attempt"))? + .map_err(|_| anyhow!("disconnectable websocket proxy stopped"))?; + Ok(()) + } + + pub(crate) fn resume(&mut self) -> anyhow::Result<()> { + self.resume_tx + .take() + .ok_or_else(|| anyhow!("disconnectable websocket proxy is already resumed"))? + .send(()) + .map_err(|_| anyhow!("disconnectable websocket proxy stopped"))?; + Ok(()) + } +} + +async fn run_disconnectable_proxy( + listener: TcpListener, + upstream: String, + pause_rx: oneshot::Receiver<()>, + blocked_connection_tx: oneshot::Sender<()>, + mut resume_rx: oneshot::Receiver<()>, +) { + let Ok((mut downstream, _)) = listener.accept().await else { + return; + }; + let Ok(mut upstream_stream) = TcpStream::connect(&upstream).await else { + return; + }; + tokio::select! { + _ = copy_bidirectional(&mut downstream, &mut upstream_stream) => return, + _ = pause_rx => {} + } + drop(downstream); + drop(upstream_stream); + + let mut blocked_connection_tx = Some(blocked_connection_tx); + loop { + tokio::select! { + _ = &mut resume_rx => break, + accepted = listener.accept() => { + let Ok((blocked, _)) = accepted else { + break; + }; + drop(blocked); + if let Some(blocked_connection_tx) = blocked_connection_tx.take() { + let _ = blocked_connection_tx.send(()); + } + } + } + } + + loop { + let Ok((mut downstream, _)) = listener.accept().await else { + return; + }; + let Ok(mut upstream_stream) = TcpStream::connect(&upstream).await else { + continue; + }; + let _ = copy_bidirectional(&mut downstream, &mut upstream_stream).await; + } +} + +async fn connect_websocket_when_ready( + websocket_url: &str, +) -> anyhow::Result<( + tokio_tungstenite::WebSocketStream>, + tokio_tungstenite::tungstenite::handshake::client::Response, +)> { + let deadline = Instant::now() + CONNECT_TIMEOUT; + loop { + match connect_async(websocket_url).await { + Ok(websocket) => return Ok(websocket), + Err(err) + if Instant::now() < deadline + && matches!( + err, + tokio_tungstenite::tungstenite::Error::Io(ref io_err) + if io_err.kind() == std::io::ErrorKind::ConnectionRefused + ) => + { + sleep(CONNECT_RETRY_INTERVAL).await; + } + Err(err) => return Err(err.into()), + } + } +} + +async fn read_listen_url_from_stdout(child: &mut Child) -> anyhow::Result { + let stdout = child + .stdout + .take() + .ok_or_else(|| anyhow!("failed to capture exec-server stdout"))?; + let mut lines = BufReader::new(stdout).lines(); + let deadline = Instant::now() + CONNECT_TIMEOUT; + + loop { + let now = Instant::now(); + if now >= deadline { + return Err(anyhow!( + "timed out waiting for exec-server listen URL on stdout after {CONNECT_TIMEOUT:?}" + )); + } + let remaining = deadline.duration_since(now); + let line = timeout(remaining, lines.next_line()) + .await + .map_err(|_| anyhow!("timed out waiting for exec-server stdout"))?? + .ok_or_else(|| anyhow!("exec-server stdout closed before emitting listen URL"))?; + let listen_url = line.trim(); + if listen_url.starts_with("ws://") { + return Ok(listen_url.to_string()); + } + } +} diff --git a/codex-rs/exec-server/tests/common/fake_bwrap.rs b/codex-rs/exec-server/tests/common/fake_bwrap.rs new file mode 100644 index 0000000000000000000000000000000000000000..ddd293648026b30b567f20df20b6d6527aa98384 --- /dev/null +++ b/codex-rs/exec-server/tests/common/fake_bwrap.rs @@ -0,0 +1,62 @@ +use std::os::unix::fs::PermissionsExt; +use std::path::Path; +use std::path::PathBuf; + +pub(crate) fn write_fake_bwrap(bin_dir: &Path) -> anyhow::Result { + std::fs::create_dir_all(bin_dir)?; + let fake_bwrap = bin_dir.join("bwrap"); + std::fs::write( + &fake_bwrap, + r#"#!/bin/bash +set -euo pipefail + +for arg in "$@"; do + if [[ "${arg}" == "--help" ]]; then + echo "Usage: bwrap --argv0 --perms --as-pid-1" + exit 0 + fi +done + +args=("$@") +argv0="" +command_start=-1 +for i in "${!args[@]}"; do + if [[ "${args[$i]}" == "--argv0" && $((i + 1)) -lt ${#args[@]} ]]; then + argv0="${args[$((i + 1))]}" + fi + if [[ "${args[$i]}" == "--" ]]; then + command_start=$((i + 1)) + break + fi +done + +if [[ "${command_start}" -lt 0 || "${command_start}" -ge "${#args[@]}" ]]; then + echo "fake bwrap did not find an inner command" >&2 + exit 125 +fi + +cmd=("${args[@]:$command_start}") +case "${cmd[0]}" in + /usr/bin/true|/bin/true|true) + exec "${cmd[@]}" + ;; +esac + +printf '%s\n' "$*" >> "${0}.log" +if [[ -f "${0}.fail-once" ]]; then + rm -f "${0}.fail-once" + echo "forced fake bwrap failure" >&2 + exit 125 +fi + +if [[ -n "${argv0}" ]]; then + exec -a "${argv0}" "${cmd[@]}" +fi +exec "${cmd[@]}" +"#, + )?; + let mut permissions = std::fs::metadata(&fake_bwrap)?.permissions(); + permissions.set_mode(0o755); + std::fs::set_permissions(&fake_bwrap, permissions)?; + Ok(fake_bwrap) +} diff --git a/codex-rs/exec-server/tests/common/mod.rs b/codex-rs/exec-server/tests/common/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..39eeae10d6f0f30c92a129723519f0428dfc914e --- /dev/null +++ b/codex-rs/exec-server/tests/common/mod.rs @@ -0,0 +1,293 @@ +use std::env; +use std::io::Write; +use std::path::Path; +use std::path::PathBuf; +use std::process::Command; +use std::process::Stdio; +use std::time::Duration; + +use codex_exec_server::CODEX_ARG0_EXEC_HELPER_ARG1; +use codex_exec_server::CODEX_FS_HELPER_ARG1; +use codex_exec_server::ExecServerRuntimePaths; +use codex_exec_server::ExecServerTelemetry; +use codex_exec_server::RequestDispatchMode; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_sandboxing::landlock::CODEX_LINUX_SANDBOX_ARG0; +use codex_test_binary_support::TestBinaryDispatchGuard; +use codex_test_binary_support::TestBinaryDispatchMode; +use codex_test_binary_support::configure_test_binary_dispatch; +use ctor::ctor; + +pub(crate) mod exec_server; + +pub(crate) const TEST_BUILD_COMMIT: &str = "0123456789abcdef0123456789abcdef01234567"; + +pub(crate) const DELAYED_OUTPUT_AFTER_EXIT_PARENT_ARG: &str = + "--codex-test-delayed-output-after-exit-parent"; +pub(crate) const SYSTEM_PROXY_REQUEST_URL_ENV: &str = + "CODEX_EXEC_SERVER_TEST_SYSTEM_PROXY_REQUEST_URL"; +pub(crate) const SYSTEM_PROXY_URL_ENV: &str = "CODEX_EXEC_SERVER_TEST_SYSTEM_PROXY_URL"; + +const CODEX_WINDOWS_SANDBOX_ARG1: &str = "--run-as-windows-sandbox"; +const DELAYED_OUTPUT_AFTER_EXIT_CHILD_ARG: &str = "--codex-test-delayed-output-after-exit-child"; + +#[macro_export] +macro_rules! skip_if_mxc_unavailable { + ($return_value:expr $(,)?) => {{ + if !codex_sandboxing::windows_mxc_available() { + eprintln!("skipping test: native MXC is unavailable on this host"); + return $return_value; + } + }}; +} + +#[ctor] +pub static TEST_BINARY_DISPATCH_GUARD: Option = { + let guard = configure_test_binary_dispatch("codex-exec-server-tests", |exe_name, argv1| { + if argv1 == Some(CODEX_ARG0_EXEC_HELPER_ARG1) { + return TestBinaryDispatchMode::DispatchArg0Only; + } + if argv1 == Some(CODEX_FS_HELPER_ARG1) { + return TestBinaryDispatchMode::DispatchArg0Only; + } + if argv1 == Some(CODEX_WINDOWS_SANDBOX_ARG1) { + return TestBinaryDispatchMode::DispatchArg0Only; + } + if argv1 == Some(codex_sandboxing::CODEX_WINDOWS_MXC_ARG1) { + return TestBinaryDispatchMode::DispatchArg0Only; + } + if exe_name == CODEX_LINUX_SANDBOX_ARG0 { + return TestBinaryDispatchMode::DispatchArg0Only; + } + TestBinaryDispatchMode::InstallAliases + }); + maybe_run_delayed_output_after_exit_from_test_binary(); + maybe_run_exec_server_from_test_binary(guard.as_ref()); + guard +}; + +pub(crate) fn current_test_binary_helper_paths() -> anyhow::Result<(PathBuf, Option)> { + let current_exe = env::current_exe()?; + let codex_linux_sandbox_exe = if cfg!(target_os = "linux") { + TEST_BINARY_DISPATCH_GUARD + .as_ref() + .and_then(|guard| guard.paths().codex_linux_sandbox_exe.clone()) + .or_else(|| Some(current_exe.clone())) + } else { + None + }; + Ok((current_exe, codex_linux_sandbox_exe)) +} + +fn maybe_run_delayed_output_after_exit_from_test_binary() { + let mut args = env::args(); + let _program = args.next(); + let Some(command) = args.next() else { + return; + }; + match command.as_str() { + DELAYED_OUTPUT_AFTER_EXIT_PARENT_ARG => { + let release_path = next_release_path_arg(args); + run_delayed_output_after_exit_parent(&release_path); + } + DELAYED_OUTPUT_AFTER_EXIT_CHILD_ARG => { + let release_path = next_release_path_arg(args); + run_delayed_output_after_exit_child(&release_path); + } + _ => {} + } +} + +fn next_release_path_arg(mut args: impl Iterator) -> PathBuf { + let Some(release_path) = args.next() else { + eprintln!("expected release path"); + std::process::exit(1); + }; + if args.next().is_some() { + eprintln!("unexpected extra arguments"); + std::process::exit(1); + } + PathBuf::from(release_path) +} + +fn run_delayed_output_after_exit_parent(release_path: &Path) { + let current_exe = match env::current_exe() { + Ok(current_exe) => current_exe, + Err(error) => { + eprintln!("failed to resolve current test binary: {error}"); + std::process::exit(1); + } + }; + match Command::new(current_exe) + .arg(DELAYED_OUTPUT_AFTER_EXIT_CHILD_ARG) + .arg(release_path) + .stdin(Stdio::null()) + .spawn() + { + Ok(_) => std::process::exit(0), + Err(error) => { + eprintln!("failed to spawn delayed output child: {error}"); + std::process::exit(1); + } + } +} + +fn run_delayed_output_after_exit_child(release_path: &Path) { + for _ in 0..1_000 { + if release_path.exists() { + let mut stdout = std::io::stdout().lock(); + if let Err(error) = writeln!(stdout, "late output after exit") { + eprintln!("failed to write delayed output: {error}"); + std::process::exit(1); + } + if let Err(error) = stdout.flush() { + eprintln!("failed to flush delayed output: {error}"); + std::process::exit(1); + } + std::process::exit(0); + } + std::thread::sleep(Duration::from_millis(10)); + } + eprintln!( + "timed out waiting for release path {}", + release_path.display() + ); + std::process::exit(1); +} + +fn maybe_run_exec_server_from_test_binary(guard: Option<&TestBinaryDispatchGuard>) { + let mut args = env::args(); + let _program = args.next(); + let Some(command) = args.next() else { + return; + }; + if command != "exec-server" { + return; + } + // Initialize in the executor child, just as the real CLI does at startup. + codex_build_info::BuildInfo::initialize(TEST_BUILD_COMMIT); + + let Some(flag) = args.next() else { + eprintln!("expected --listen"); + std::process::exit(1); + }; + if flag != "--listen" { + eprintln!("expected --listen, got `{flag}`"); + std::process::exit(1); + } + let Some(listen_url) = args.next() else { + eprintln!("expected listen URL"); + std::process::exit(1); + }; + let remaining_args = args.collect::>(); + let request_dispatch_mode = match remaining_args.as_slice() { + [] => RequestDispatchMode::Inline, + [flag, value] if flag == "--concurrent-requests" => match value.parse() { + Ok(mode) => mode, + Err(error) => { + eprintln!("invalid concurrent request count: {error}"); + std::process::exit(1); + } + }, + args => { + eprintln!("unexpected exec-server arguments: {args:?}"); + std::process::exit(1); + } + }; + + let current_exe = match env::current_exe() { + Ok(current_exe) => current_exe, + Err(error) => { + eprintln!("failed to resolve current test binary: {error}"); + std::process::exit(1); + } + }; + let runtime_paths = match ExecServerRuntimePaths::new( + current_exe.clone(), + linux_sandbox_exe(guard, ¤t_exe), + ) { + Ok(runtime_paths) => runtime_paths, + Err(error) => { + eprintln!("failed to configure exec-server runtime paths: {error}"); + std::process::exit(1); + } + }; + let runtime = match tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build() + { + Ok(runtime) => runtime, + Err(error) => { + eprintln!("failed to build Tokio runtime: {error}"); + std::process::exit(1); + } + }; + let http_client_factory = match ( + env::var(SYSTEM_PROXY_REQUEST_URL_ENV), + env::var(SYSTEM_PROXY_URL_ENV), + ) { + (Ok(request_url), Ok(proxy_url)) => { + codex_http_client::cache_system_proxy_route_for_test(&request_url, proxy_url); + HttpClientFactory::new(OutboundProxyPolicy::RespectSystemProxy) + } + (Err(env::VarError::NotPresent), Err(env::VarError::NotPresent)) => { + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault) + } + _ => { + eprintln!("system proxy test configuration requires both request and proxy URLs"); + std::process::exit(1); + } + }; + let exit_code = match runtime.block_on(async { + #[cfg(target_os = "macos")] + let runtime_paths = { + let home = codex_utils_home_dir::find_codex_home()?; + let config = codex_config::loader::load_config_layers_state( + &codex_exec_server::LocalFileSystem::unsandboxed(), + home.as_path(), + /*cwd*/ None, + &[], + codex_config::LoaderOverrides::default(), + &codex_config::NoopThreadConfigLoader, + ) + .await?; + runtime_paths.with_allowed_symlinked_codex_home( + codex_config::allowed_symlinked_codex_home(&config, &home), + ) + }; + codex_exec_server::run_main_with_telemetry( + &listen_url, + runtime_paths, + ExecServerTelemetry::default(), + http_client_factory, + request_dispatch_mode, + ) + .await + }) { + Ok(()) => 0, + Err(error) => { + eprintln!("exec-server failed: {error}"); + 1 + } + }; + std::process::exit(exit_code); +} + +fn linux_sandbox_exe( + guard: Option<&TestBinaryDispatchGuard>, + current_exe: &std::path::Path, +) -> Option { + #[cfg(target_os = "linux")] + { + guard + .and_then(|guard| guard.paths().codex_linux_sandbox_exe.clone()) + .or_else(|| Some(current_exe.to_path_buf())) + } + #[cfg(not(target_os = "linux"))] + { + let _ = guard; + let _ = current_exe; + None + } +} diff --git a/codex-rs/exec-server/tests/common/relay.rs b/codex-rs/exec-server/tests/common/relay.rs new file mode 100644 index 0000000000000000000000000000000000000000..cca0f967a45009734f8c8dcfa8ef4f9c67d92fb3 --- /dev/null +++ b/codex-rs/exec-server/tests/common/relay.rs @@ -0,0 +1,156 @@ +use std::sync::Arc; +use std::sync::Mutex; + +use anyhow::Context; +use anyhow::Result; +use codex_api::AuthProvider; +use codex_exec_server::ExecServerClient; +use codex_exec_server::NoiseChannelIdentity; +use codex_exec_server::NoiseRendezvousConnectArgs; +use codex_exec_server::NoiseRendezvousConnectBundle; +use codex_exec_server::RemoteEnvironmentConfig; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use http::HeaderMap; +use http::HeaderValue; +use tokio::net::TcpListener; +use tokio::time::timeout; +use tokio_util::task::AbortOnDropHandle; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +pub(crate) use codex_exec_server_test_support::relay::TEST_TIMEOUT; +pub(crate) use codex_exec_server_test_support::relay::accept_websocket; +pub(crate) use codex_exec_server_test_support::relay::assert_relay_data_is_encrypted; +pub(crate) use codex_exec_server_test_support::relay::proxy_relay_frames; +pub(crate) use codex_exec_server_test_support::relay::registered_executor_public_key; + +pub(crate) const ENVIRONMENT_ID: &str = "env-noise-relay-test"; +pub(crate) const EXECUTOR_REGISTRATION_ID: &str = "registration-1"; +pub(crate) const HARNESS_KEY_AUTHORIZATION: &str = "harness-key-authorization"; +pub(crate) const REGISTRY_TOKEN: &str = "registry-token"; + +#[derive(Debug)] +struct StaticRegistryAuthProvider; + +impl AuthProvider for StaticRegistryAuthProvider { + fn add_auth_headers(&self, headers: &mut HeaderMap) { + let _ = headers.insert( + http::header::AUTHORIZATION, + HeaderValue::from_static("Bearer registry-token"), + ); + } +} + +pub(crate) fn static_registry_auth_provider() -> codex_api::SharedAuthProvider { + Arc::new(StaticRegistryAuthProvider) +} + +pub(crate) struct RelayTest { + registry: MockServer, + listener: TcpListener, +} + +pub(crate) struct RelayConnection { + pub(crate) client: ExecServerClient, + captured_frames: Arc>>>, + relay_task: AbortOnDropHandle>, +} + +impl RelayTest { + pub(crate) async fn new() -> Result { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let rendezvous_url = format!("ws://{}", listener.local_addr()?); + let registry = MockServer::start().await; + Mock::given(method("POST")) + .and(path(format!( + "/cloud/environment/{ENVIRONMENT_ID}/register" + ))) + .and(header("authorization", format!("Bearer {REGISTRY_TOKEN}"))) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "environment_id": ENVIRONMENT_ID, + "url": format!("{rendezvous_url}/relay?role=environment"), + "security_profile": "noise_hybrid_ik_v1", + "executor_registration_id": EXECUTOR_REGISTRATION_ID, + }))) + .mount(®istry) + .await; + Mock::given(method("POST")) + .and(path(format!( + "/cloud/environment/{ENVIRONMENT_ID}/validate" + ))) + .and(header("authorization", format!("Bearer {REGISTRY_TOKEN}"))) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "valid": true, + }))) + .mount(®istry) + .await; + Ok(Self { registry, listener }) + } + + pub(crate) fn config(&self) -> Result { + Ok(RemoteEnvironmentConfig::new( + self.registry.uri(), + ENVIRONMENT_ID.to_string(), + static_registry_auth_provider(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )?) + } + + pub(crate) async fn connect(&self) -> Result { + let rendezvous_url = format!("ws://{}", self.listener.local_addr()?); + let environment_websocket = accept_websocket(&self.listener, "environment").await?; + let executor_public_key = registered_executor_public_key(&self.registry).await?; + let harness_identity = NoiseChannelIdentity::generate()?; + let client_args = NoiseRendezvousConnectArgs { + bundle: NoiseRendezvousConnectBundle { + websocket_url: format!("{rendezvous_url}/relay?role=harness"), + environment_id: ENVIRONMENT_ID.to_string(), + executor_registration_id: EXECUTOR_REGISTRATION_ID.to_string(), + executor_public_key, + harness_key_authorization: HARNESS_KEY_AUTHORIZATION.to_string(), + }, + harness_identity, + client_name: "noise-relay-test".to_string(), + connect_timeout: TEST_TIMEOUT, + initialize_timeout: TEST_TIMEOUT, + resume_session_id: None, + http_client_factory: HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + }; + let client_task = + tokio::spawn( + async move { ExecServerClient::connect_noise_rendezvous(client_args).await }, + ); + let harness_websocket = accept_websocket(&self.listener, "harness").await?; + let captured_frames = Arc::new(Mutex::new(Vec::new())); + let relay_task = AbortOnDropHandle::new(tokio::spawn(proxy_relay_frames( + environment_websocket, + harness_websocket, + Arc::clone(&captured_frames), + ))); + let client = timeout(TEST_TIMEOUT, client_task) + .await + .context("Noise harness client should connect")???; + Ok(RelayConnection { + client, + captured_frames, + relay_task, + }) + } +} + +impl RelayConnection { + pub(crate) fn assert_encrypted(&self) -> Result<()> { + assert_relay_data_is_encrypted(&self.captured_frames) + } + + pub(crate) async fn close(self) { + drop(self.client); + self.relay_task.abort(); + let _ = self.relay_task.await; + } +} diff --git a/codex-rs/exec-server/tests/deferred_environment.rs b/codex-rs/exec-server/tests/deferred_environment.rs new file mode 100644 index 0000000000000000000000000000000000000000..e7bc32d243f4ebbd51c73e499cfea36cc960d256 --- /dev/null +++ b/codex-rs/exec-server/tests/deferred_environment.rs @@ -0,0 +1,458 @@ +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; + +use codex_exec_server::EnvironmentReadyInfo; +use codex_exec_server::ExecServerError; +use codex_exec_server::NoiseChannelPublicKey; +use codex_exec_server::NoiseRendezvousConnectBundle; +use codex_exec_server::NoiseRendezvousConnectProvider; +use codex_exec_server_test_support::environment_manager_without_environments; +use codex_protocol::capabilities::CapabilityRootLocation; +use codex_protocol::capabilities::SelectedCapabilityRoot; +use codex_utils_path_uri::PathUri; +use futures::FutureExt; +use futures::future::BoxFuture; +use futures::poll; +use pretty_assertions::assert_eq; + +#[derive(Default)] +struct FailingNoiseConnectProvider { + calls: AtomicUsize, +} + +impl FailingNoiseConnectProvider { + fn calls(&self) -> usize { + self.calls.load(Ordering::Relaxed) + } +} + +impl NoiseRendezvousConnectProvider for FailingNoiseConnectProvider { + fn connect_bundle( + &self, + _: NoiseChannelPublicKey, + ) -> BoxFuture<'_, Result> { + self.calls.fetch_add(1, Ordering::Relaxed); + async { + Err(ExecServerError::Protocol( + "test Noise provider called".to_string(), + )) + } + .boxed() + } +} + +fn ready_info(root_id: &str, environment_id: &str) -> anyhow::Result { + Ok(EnvironmentReadyInfo { + selected_capability_roots: vec![SelectedCapabilityRoot { + id: root_id.to_string(), + location: CapabilityRootLocation::Environment { + environment_id: environment_id.to_string(), + path: PathUri::parse("file:///plugins/root")?, + }, + }], + }) +} + +#[tokio::test] +async fn readiness_before_materialization_creates_the_stable_environment() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + let readiness_provider = Arc::new(FailingNoiseConnectProvider::default()); + let materialization_provider = Arc::new(FailingNoiseConnectProvider::default()); + let selected = ready_info("selected-root", "tools")?; + + let ready = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(selected.clone()), + readiness_provider.clone(), + )? + .expect("readiness report should create the environment"); + let materialized = manager.materialize_pending_noise_environment( + "tools".to_string(), + materialization_provider.clone(), + )?; + + assert!(Arc::ptr_eq(&ready, &materialized)); + assert_eq!( + ready.selected_capability_roots(), + selected.selected_capability_roots + ); + let error = ready.wait_until_ready().await.unwrap_err(); + assert!(error.to_string().contains("test Noise provider called")); + assert_eq!(readiness_provider.calls(), 1); + assert_eq!(materialization_provider.calls(), 0); + Ok(()) +} + +#[tokio::test] +async fn materialize_then_report_ready_reuses_the_pending_environment() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + let pending_provider = Arc::new(FailingNoiseConnectProvider::default()); + let pending = manager + .materialize_pending_noise_environment("tools".to_string(), pending_provider.clone())?; + assert_eq!(pending.last_ready_info(), None); + let mut pending_readiness = Box::pin(pending.wait_until_ready()); + assert!(poll!(&mut pending_readiness).is_pending()); + let ready = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(ready_info("selected-root", "tools")?), + Arc::new(FailingNoiseConnectProvider::default()), + )? + .expect("provisioning report should apply to the pending environment"); + + assert!(Arc::ptr_eq(&pending, &ready)); + let error = pending_readiness.await.unwrap_err(); + assert!(error.to_string().contains("test Noise provider called")); + assert_eq!(pending_provider.calls(), 1); + Ok(()) +} + +#[tokio::test] +async fn ordinary_environment_ignores_provisioning_reports() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + manager.upsert_environment( + "tools".to_string(), + "ws://127.0.0.1:1".to_string(), + Some(std::time::Duration::from_millis(1)), + )?; + let existing_environment = manager + .get_environment("tools") + .expect("existing environment"); + + let reported = manager.report_environment_provisioning_status( + "tools".to_string(), + Ok(ready_info("selected-root", "tools")?), + Arc::new(FailingNoiseConnectProvider::default()), + )?; + + let current_environment = manager + .get_environment("tools") + .expect("current environment"); + assert!(Arc::ptr_eq(&existing_environment, ¤t_environment)); + assert!(reported.is_none()); + assert!(existing_environment.selected_capability_roots().is_empty()); + assert_eq!(existing_environment.last_ready_info(), None); + Ok(()) +} + +#[tokio::test] +async fn failure_before_materialization_is_reported_without_connecting() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + let provider = Arc::new(FailingNoiseConnectProvider::default()); + + let failed = manager + .report_environment_provisioning_status( + "tools".to_string(), + Err("provisioning failed".to_string()), + provider.clone(), + )? + .expect("failure report should create the environment"); + let materialized = manager.materialize_pending_noise_environment( + "tools".to_string(), + Arc::new(FailingNoiseConnectProvider::default()), + )?; + + assert!(Arc::ptr_eq(&failed, &materialized)); + assert_eq!( + failed.status().await, + codex_exec_server::EnvironmentObservedStatus::Disconnected { + error: "environment unavailable: provisioning failed".to_string(), + } + ); + let error = failed.wait_until_ready().await.unwrap_err(); + assert!(error.to_string().ends_with("provisioning failed")); + assert!(failed.startup_finished()); + assert_eq!(provider.calls(), 0); + Ok(()) +} + +#[tokio::test] +async fn failure_releases_the_existing_pending_environment_without_connecting() -> anyhow::Result<()> +{ + let manager = environment_manager_without_environments(); + let provider = Arc::new(FailingNoiseConnectProvider::default()); + let pending = + manager.materialize_pending_noise_environment("tools".to_string(), provider.clone())?; + + let reported = manager + .report_environment_provisioning_status( + "tools".to_string(), + Err("provisioning failed".to_string()), + provider.clone(), + )? + .expect("failure report should apply to the pending environment"); + + assert!(Arc::ptr_eq(&pending, &reported)); + let error = pending.wait_until_ready().await.unwrap_err(); + assert!(error.to_string().ends_with("provisioning failed")); + assert_eq!(provider.calls(), 0); + Ok(()) +} + +#[tokio::test] +async fn repeated_failure_preserves_the_first_error_until_ready() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + let provider = Arc::new(FailingNoiseConnectProvider::default()); + let failed = manager + .report_environment_provisioning_status( + "tools".to_string(), + Err("first failure".to_string()), + provider.clone(), + )? + .expect("failure report should create the environment"); + + let repeated = manager + .report_environment_provisioning_status( + "tools".to_string(), + Err("different failure".to_string()), + provider.clone(), + )? + .expect("repeated failure should be idempotent"); + assert!(Arc::ptr_eq(&failed, &repeated)); + + let error = failed.wait_until_ready().await.unwrap_err(); + assert!(error.to_string().ends_with("first failure")); + assert_eq!(provider.calls(), 0); + let invalid_ready_error = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(ready_info("selected-root", "other")?), + provider.clone(), + ) + .unwrap_err(); + assert!(matches!(invalid_ready_error, ExecServerError::Protocol(_))); + assert_eq!( + failed.wait_until_ready().await.unwrap_err().to_string(), + error.to_string() + ); + assert!(failed.selected_capability_roots().is_empty()); + assert_eq!(failed.last_ready_info(), None); + assert_eq!(provider.calls(), 0); + let selected = ready_info("selected-root", "tools")?; + let ready = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(selected.clone()), + provider.clone(), + )? + .expect("successful provisioning should recover the same environment"); + assert!(Arc::ptr_eq(&failed, &ready)); + assert_eq!(failed.last_ready_info().as_deref(), Some(&selected)); + assert_eq!( + failed.selected_capability_roots(), + selected.selected_capability_roots + ); + let error = failed.wait_until_ready().await.unwrap_err(); + assert!(error.to_string().contains("test Noise provider called")); + assert_eq!(provider.calls(), 1); + Ok(()) +} + +#[tokio::test] +async fn ready_environment_rejects_a_later_failure() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + let provider = Arc::new(FailingNoiseConnectProvider::default()); + let ready = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(ready_info("selected-root", "tools")?), + provider.clone(), + )? + .expect("ready report should create the environment"); + + let error = manager + .report_environment_provisioning_status( + "tools".to_string(), + Err("late failure".to_string()), + provider, + ) + .unwrap_err(); + + assert!(error.to_string().contains("already ready")); + assert_eq!(ready.selected_capability_roots().len(), 1); + Ok(()) +} + +#[tokio::test] +async fn existing_environment_accepts_matching_readiness() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + let provider = Arc::new(FailingNoiseConnectProvider::default()); + let ready_info = ready_info("selected-root", "tools")?; + + let environment = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(ready_info.clone()), + provider.clone(), + )? + .expect("readiness report should create the environment"); + manager.report_environment_provisioning_status( + "tools".to_string(), + Ok(ready_info.clone()), + provider, + )?; + assert_eq!( + environment.selected_capability_roots(), + ready_info.selected_capability_roots + ); + Ok(()) +} + +#[tokio::test] +async fn existing_environment_overwrites_reported_readiness() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + let provider = Arc::new(FailingNoiseConnectProvider::default()); + let environment = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(ready_info("selected-root", "tools")?), + provider.clone(), + )? + .expect("readiness report should create the environment"); + + let updated_ready_info = ready_info("different-root", "tools")?; + manager.report_environment_provisioning_status( + "tools".to_string(), + Ok(updated_ready_info.clone()), + provider, + )?; + assert_eq!( + environment.selected_capability_roots(), + updated_ready_info.selected_capability_roots + ); + assert!(Arc::ptr_eq( + &environment, + &manager.get_environment("tools").expect("environment") + )); + Ok(()) +} + +#[tokio::test] +async fn last_ready_info_preserves_snapshots_through_replacement_and_clear() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + let provider = Arc::new(FailingNoiseConnectProvider::default()); + let selected = ready_info("selected-root", "tools")?; + let environment = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(selected.clone()), + provider.clone(), + )? + .expect("readiness report should create the environment"); + let snapshot = environment.last_ready_info(); + assert_eq!(snapshot.as_deref(), Some(&selected)); + + for replacement in [ + ready_info("different-root", "tools")?, + EnvironmentReadyInfo::default(), + ] { + manager.report_environment_provisioning_status( + "tools".to_string(), + Ok(replacement.clone()), + provider.clone(), + )?; + assert_eq!(environment.last_ready_info().as_deref(), Some(&replacement)); + assert_eq!(snapshot.as_deref(), Some(&selected)); + } + + let error = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(ready_info("invalid-root", "other")?), + provider.clone(), + ) + .unwrap_err(); + assert!(matches!(error, ExecServerError::Protocol(_))); + assert_eq!( + environment.last_ready_info().as_deref(), + Some(&EnvironmentReadyInfo::default()) + ); + assert_eq!(provider.calls(), 0); + Ok(()) +} + +#[tokio::test] +async fn invalid_ready_report_fails_the_provisioning_gate() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + let provider = Arc::new(FailingNoiseConnectProvider::default()); + let environment = + manager.materialize_pending_noise_environment("tools".to_string(), provider.clone())?; + let readiness = Box::pin(environment.wait_until_ready()); + + let error = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(ready_info("selected-root", "other")?), + provider.clone(), + ) + .unwrap_err(); + + assert!(matches!(error, ExecServerError::Protocol(_))); + let readiness_error = readiness.await.unwrap_err(); + assert!(readiness_error.to_string().contains(&error.to_string())); + assert!(environment.selected_capability_roots().is_empty()); + assert_eq!(provider.calls(), 0); + + let selected = ready_info("selected-root", "tools")?; + let reported = manager + .report_environment_provisioning_status( + "tools".to_string(), + Ok(selected.clone()), + provider.clone(), + )? + .expect("a corrected ready report should recover provisioning"); + assert!(Arc::ptr_eq(&environment, &reported)); + assert_eq!( + environment.selected_capability_roots(), + selected.selected_capability_roots + ); + assert_eq!(provider.calls(), 0); + Ok(()) +} + +#[tokio::test] +async fn duplicate_materialization_reuses_the_pending_environment() -> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + let provider = Arc::new(FailingNoiseConnectProvider::default()); + let environment = + manager.materialize_pending_noise_environment("tools".to_string(), provider.clone())?; + let replacement_provider = Arc::new(FailingNoiseConnectProvider::default()); + + let current = manager + .materialize_pending_noise_environment("tools".to_string(), replacement_provider.clone())?; + assert!(Arc::ptr_eq(&environment, ¤t)); + assert_eq!(provider.calls(), 0); + assert_eq!(replacement_provider.calls(), 0); + Ok(()) +} + +#[tokio::test] +async fn deferred_materialization_conflicts_with_an_existing_ordinary_environment() +-> anyhow::Result<()> { + let manager = environment_manager_without_environments(); + manager.upsert_environment( + "tools".to_string(), + "ws://127.0.0.1:1".to_string(), + Some(std::time::Duration::from_millis(1)), + )?; + let existing_environment = manager.get_environment("tools").expect("environment"); + let deferred_provider = Arc::new(FailingNoiseConnectProvider::default()); + + let error = manager + .materialize_pending_noise_environment("tools".to_string(), deferred_provider.clone()) + .unwrap_err(); + + assert!(matches!( + error, + ExecServerError::ProvisioningModeConflict { environment_id } + if environment_id == "tools" + )); + let current_environment = manager + .get_environment("tools") + .expect("ordinary environment should remain registered"); + assert!(Arc::ptr_eq(&existing_environment, ¤t_environment)); + assert_eq!(deferred_provider.calls(), 0); + Ok(()) +} diff --git a/codex-rs/exec-server/tests/environment.rs b/codex-rs/exec-server/tests/environment.rs new file mode 100644 index 0000000000000000000000000000000000000000..814ae7f2d25257501411a75df596e71c4e89209d --- /dev/null +++ b/codex-rs/exec-server/tests/environment.rs @@ -0,0 +1,315 @@ +mod common; + +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use anyhow::Context; +use codex_exec_server::EnvironmentManager; +use codex_exec_server::ExecutorCapabilityDiscoveryCache; +use codex_exec_server::REMOTE_ENVIRONMENT_ID; +use codex_exec_server::SelectedCapabilityRootsStatus; +use codex_exec_server_protocol::CAPABILITY_ROOTS_DISCOVER_METHOD; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_http_client::cache_system_proxy_route_for_test; +use codex_protocol::capabilities::CapabilityRootLocation; +use codex_protocol::capabilities::SelectedCapabilityRoot; +use codex_utils_path_uri::PathUri; +use common::exec_server::exec_server; +use futures::SinkExt; +use futures::StreamExt; +use pretty_assertions::assert_eq; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::sync::oneshot; +use tokio::time::sleep; +use tokio::time::timeout; +use tokio_tungstenite::accept_async; +use tokio_tungstenite::connect_async; +use tokio_tungstenite::tungstenite::Message; +use tokio_util::task::AbortOnDropHandle; + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial_test::serial(remote_exec_server)] +async fn prepared_remote_environment_uses_configured_system_proxy() -> anyhow::Result<()> { + let server = exec_server().await?; + let upstream = server + .websocket_url() + .strip_prefix("ws://") + .context("exec-server websocket should use ws://")? + .to_string(); + let proxy_listener = TcpListener::bind("127.0.0.1:0").await?; + let proxy_url = format!("http://{}", proxy_listener.local_addr()?); + let websocket_url = "ws://exec-server-system-proxy.invalid:8765/"; + let proxy_resolution_url = "http://exec-server-system-proxy.invalid:8765/"; + cache_system_proxy_route_for_test(proxy_resolution_url, proxy_url); + + let (request_tx, request_rx) = oneshot::channel(); + let _proxy_task = AbortOnDropHandle::new(tokio::spawn(async move { + let (mut client, _) = proxy_listener.accept().await?; + let mut request = Vec::new(); + let mut byte = [0_u8; 1]; + while !request.ends_with(b"\r\n\r\n") { + client.read_exact(&mut byte).await?; + request.push(byte[0]); + } + let request_line = String::from_utf8(request)? + .lines() + .next() + .context("system proxy should receive a CONNECT request")? + .to_string(); + request_tx + .send(request_line) + .map_err(|_| anyhow::anyhow!("system proxy request receiver was dropped"))?; + + let mut target = TcpStream::connect(upstream).await?; + client + .write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n") + .await?; + tokio::io::copy_bidirectional(&mut client, &mut target).await?; + Ok::<(), anyhow::Error>(()) + })); + + let codex_home = tempfile::tempdir()?; + std::fs::write( + codex_home.path().join("environments.toml"), + format!( + "default = \"{REMOTE_ENVIRONMENT_ID}\"\ninclude_local = false\n\n[[environments]]\nid = \"{REMOTE_ENVIRONMENT_ID}\"\nurl = \"{websocket_url}\"\n" + ), + )?; + + let prepared = EnvironmentManager::prepare_from_codex_home(codex_home.path()).await?; + assert!(prepared.default_environment_is_remote()); + let manager = prepared.build( + /*local_runtime_paths*/ None, + HttpClientFactory::new(OutboundProxyPolicy::RespectSystemProxy), + )?; + + let request_line = timeout(Duration::from_secs(5), request_rx) + .await + .context("prepared environment did not connect through the system proxy")??; + assert_eq!( + request_line, + "CONNECT exec-server-system-proxy.invalid:8765 HTTP/1.1" + ); + let environment = manager + .default_environment() + .context("prepared remote environment")?; + timeout(Duration::from_secs(5), environment.info()) + .await + .context("prepared remote environment did not initialize through the system proxy")??; + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial_test::serial(remote_exec_server)] +async fn selected_capability_inspection_tracks_connection_recovery() -> anyhow::Result<()> { + let server = exec_server().await?; + let mut proxy = server.disconnectable_websocket_proxy().await?; + let manager = EnvironmentManager::create_for_tests( + Some(proxy.websocket_url().to_string()), + /*local_runtime_paths*/ None, + ) + .await; + let environment = manager + .default_environment() + .context("remote environment")?; + environment.info().await?; + + let skill_root_path = PathUri::parse("file:///plugins/demo")?; + let selected_root = SelectedCapabilityRoot { + id: "demo@1".to_string(), + location: CapabilityRootLocation::Environment { + environment_id: REMOTE_ENVIRONMENT_ID.to_string(), + path: skill_root_path.clone(), + }, + }; + assert_eq!( + manager.inspect_selected_capability_roots(std::slice::from_ref(&selected_root)), + SelectedCapabilityRootsStatus { + ready_roots: vec![selected_root.clone()], + warnings: Vec::new(), + } + ); + let file_system = environment.get_filesystem_without_reconnect(); + + proxy.pause_and_disconnect().await?; + assert_eq!( + manager.inspect_selected_capability_roots(std::slice::from_ref(&selected_root)), + SelectedCapabilityRootsStatus::default() + ); + let read_result = timeout( + Duration::from_secs(1), + file_system.read_directory(&skill_root_path, /*sandbox*/ None), + ) + .await + .context("passive filesystem read waited for recovery")?; + assert!(read_result.is_err()); + + proxy.resume()?; + let recovered_status = timeout(Duration::from_secs(5), async { + loop { + let status = + manager.inspect_selected_capability_roots(std::slice::from_ref(&selected_root)); + if !status.ready_roots.is_empty() { + break status; + } + sleep(Duration::from_millis(10)).await; + } + }) + .await + .context("environment did not recover")?; + assert_eq!( + recovered_status, + SelectedCapabilityRootsStatus { + ready_roots: vec![selected_root], + warnings: Vec::new(), + } + ); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial_test::serial(remote_exec_server)] +async fn capability_discovery_retries_executor_disconnect_within_same_request() -> anyhow::Result<()> +{ + let server = exec_server().await?; + let proxy_listener = TcpListener::bind("127.0.0.1:0").await?; + let proxy_websocket_url = format!("ws://{}", proxy_listener.local_addr()?); + let upstream_websocket_url = server.websocket_url().to_string(); + let discovery_attempts = Arc::new(AtomicUsize::new(0)); + let proxy_discovery_attempts = Arc::clone(&discovery_attempts); + let _proxy_task = AbortOnDropHandle::new(tokio::spawn(async move { + while let Ok((downstream, _)) = proxy_listener.accept().await { + let mut downstream = accept_async(downstream).await?; + let (mut upstream, _) = connect_async(&upstream_websocket_url).await?; + + loop { + tokio::select! { + message = downstream.next() => { + let Some(message) = message.transpose()? else { + break; + }; + if let Message::Text(message_text) = &message { + let request = serde_json::from_str::(message_text.as_ref())?; + if request.get("method").and_then(serde_json::Value::as_str) + == Some(CAPABILITY_ROOTS_DISCOVER_METHOD) + { + let attempt = proxy_discovery_attempts.fetch_add(1, Ordering::SeqCst); + if attempt == 0 { + break; + } + sleep(Duration::from_secs(9)).await; + } + } + upstream.send(message).await?; + } + message = upstream.next() => { + let Some(message) = message.transpose()? else { + break; + }; + downstream.send(message).await?; + } + } + } + } + Ok::<(), anyhow::Error>(()) + })); + let manager = Arc::new( + EnvironmentManager::create_for_tests( + Some(proxy_websocket_url), + /*local_runtime_paths*/ None, + ) + .await, + ); + manager + .default_environment() + .context("remote environment")? + .info() + .await?; + + let cache = Arc::new(ExecutorCapabilityDiscoveryCache::new(Arc::clone(&manager))); + let skill_root = tempfile::tempdir()?; + let selected_roots = vec![SelectedCapabilityRoot { + id: "recovering-skill".to_string(), + location: CapabilityRootLocation::Environment { + environment_id: REMOTE_ENVIRONMENT_ID.to_string(), + path: PathUri::from_host_native_path(skill_root.path())?, + }, + }]; + + let snapshot = timeout( + Duration::from_secs(12), + cache.snapshot(&selected_roots, &HashMap::new()), + ) + .await + .context("capability discovery did not retry within the same request")?; + let discovery = snapshot.roots()[0] + .result + .as_ref() + .map_err(|error| anyhow::anyhow!("{error}"))?; + + assert_eq!(discovery.id, "recovering-skill"); + assert_eq!( + 2, + discovery_attempts.load(Ordering::SeqCst), + "same-request retry must issue a second capability discovery RPC" + ); + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn capability_discovery_retries_after_executor_reconnects() -> anyhow::Result<()> { + let server = exec_server().await?; + let manager = Arc::new(EnvironmentManager::default_for_tests()); + let cache = ExecutorCapabilityDiscoveryCache::new(Arc::clone(&manager)); + let skill_root = tempfile::tempdir()?; + let refused_listener = TcpListener::bind("127.0.0.1:0").await?; + let refused_address = refused_listener.local_addr()?; + drop(refused_listener); + manager.upsert_environment( + "recovering".to_string(), + format!("ws://{refused_address}"), + Some(Duration::from_millis(100)), + )?; + let selected_roots = vec![SelectedCapabilityRoot { + id: "recovering-skill".to_string(), + location: CapabilityRootLocation::Environment { + environment_id: "recovering".to_string(), + path: PathUri::from_host_native_path(skill_root.path())?, + }, + }]; + + let failed_snapshot = cache.snapshot(&selected_roots, &HashMap::new()).await; + assert!(failed_snapshot.roots()[0].result.is_err()); + assert!(!cache.take_recovered_discovery()); + + manager.upsert_environment( + "recovering".to_string(), + server.websocket_url().to_string(), + /*connect_timeout*/ None, + )?; + manager + .get_environment("recovering") + .context("recovered environment")? + .wait_until_ready() + .await?; + + let recovered_snapshot = cache.snapshot(&selected_roots, &HashMap::new()).await; + let discovery = recovered_snapshot.roots()[0] + .result + .as_ref() + .map_err(|error| anyhow::anyhow!("{error}"))?; + + assert_eq!(discovery.id, "recovering-skill"); + assert!(cache.take_recovered_discovery()); + assert!(!cache.take_recovered_discovery()); + Ok(()) +} diff --git a/codex-rs/exec-server/tests/environment_config.rs b/codex-rs/exec-server/tests/environment_config.rs new file mode 100644 index 0000000000000000000000000000000000000000..5df1bea910ff19abf12e50957f27ccdaf2e118ee --- /dev/null +++ b/codex-rs/exec-server/tests/environment_config.rs @@ -0,0 +1,137 @@ +mod common; + +use codex_config::CONFIG_TOML_FILE; +use codex_config::ConfigLayerSource; +use codex_config::format_config_layer_source; +use codex_config::loader::project_trust_key; +use codex_exec_server::Environment; +use codex_exec_server::EnvironmentConfigLayer; +use codex_exec_server::EnvironmentConfigLayerStack; +use codex_exec_server::EnvironmentConfigReadParams; +use codex_exec_server::EnvironmentConfigReadResponse; +use codex_exec_server::ExecServerError; +use codex_utils_absolute_path::AbsolutePathBuf; +use codex_utils_path_uri::PathUri; +use common::exec_server::exec_server; +use pretty_assertions::assert_eq; + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn remote_environment_reads_projected_executor_config() -> anyhow::Result<()> { + let mut server = exec_server().await?; + let codex_home = + AbsolutePathBuf::from_absolute_path(std::fs::canonicalize(server.codex_home())?)?; + let config_file = codex_home.join(CONFIG_TOML_FILE); + let project = codex_home.join("project"); + let dot_codex = project.join(".codex"); + tokio::fs::create_dir_all(dot_codex.as_path()).await?; + tokio::fs::write(project.join(".project-root").as_path(), "").await?; + let project_key = toml::Value::String(project_trust_key(project.as_path())).to_string(); + tokio::fs::write( + &config_file, + format!( + "project_root_markers = [\".project-root\"]\n[projects.{project_key}]\ntrust_level = \"trusted\"" + ), + ) + .await?; + tokio::fs::write( + dot_codex.join(CONFIG_TOML_FILE).as_path(), + r#" +[future_environment] +relative_path = "./executor-relative" +unselected = "do not return" +"#, + ) + .await?; + + let environment = Environment::create_for_tests(Some(server.websocket_url().to_string()))?; + let environment_info = environment.info().await?; + assert!(environment_info.capabilities.environment_config_read); + assert_eq!( + environment_info.user_home_dir, + dirs::home_dir().and_then(|home_dir| PathUri::from_host_native_path(home_dir).ok()), + ); + + let response = environment + .read_environment_config(EnvironmentConfigReadParams { + cwd: PathUri::from_abs_path(&project), + config_paths: vec![vec![ + "future_environment".to_string(), + "relative_path".to_string(), + ]], + requirements_paths: Vec::new(), + }) + .await?; + + let projected_toml = toml::toml! { + [future_environment] + relative_path = "./executor-relative" + }; + assert_eq!( + response, + EnvironmentConfigReadResponse { + user_home_dir: dirs::home_dir() + .and_then(|home_dir| PathUri::from_host_native_path(home_dir).ok()), + codex_home_dir: PathUri::from_abs_path(&codex_home), + hostname: codex_config::host_name(), + config: EnvironmentConfigLayerStack { + layers: vec![EnvironmentConfigLayer { + source: format_config_layer_source( + &ConfigLayerSource::Project { + dot_codex_folder: dot_codex.clone(), + }, + CONFIG_TOML_FILE, + ), + base_dir: PathUri::from_abs_path(&dot_codex), + toml: toml::to_string(&projected_toml)?, + }], + cloud_insertion_index: 0, + }, + requirements: EnvironmentConfigLayerStack { + layers: Vec::new(), + cloud_insertion_index: 0, + }, + } + ); + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn environment_config_read_rejects_empty_selectors() -> anyhow::Result<()> { + let mut server = exec_server().await?; + let codex_home = + AbsolutePathBuf::from_absolute_path(std::fs::canonicalize(server.codex_home())?)?; + let environment = Environment::create_for_tests(Some(server.websocket_url().to_string()))?; + + for (config_paths, expected_message) in [ + ( + Vec::new(), + "at least one config or requirements path is required", + ), + ( + vec![Vec::new()], + "TOML paths must contain at least one key segment", + ), + ] { + let error = environment + .read_environment_config(EnvironmentConfigReadParams { + cwd: PathUri::from_abs_path(&codex_home), + config_paths, + requirements_paths: Vec::new(), + }) + .await + .expect_err("invalid selectors should fail"); + assert!( + matches!( + error, + ExecServerError::Server { code: -32602, ref message } + if message == expected_message + ), + "unexpected error: {error:?}" + ); + } + + server.shutdown().await?; + Ok(()) +} diff --git a/codex-rs/exec-server/tests/exec_process.rs b/codex-rs/exec-server/tests/exec_process.rs new file mode 100644 index 0000000000000000000000000000000000000000..b4e861d50b91c042bf29ade4521dbb9d65200e68 --- /dev/null +++ b/codex-rs/exec-server/tests/exec_process.rs @@ -0,0 +1,1839 @@ +mod common; +#[path = "exec_process/windows_sandbox.rs"] +mod windows_sandbox; + +use std::collections::HashMap; +#[cfg(unix)] +use std::os::unix::fs::PermissionsExt; +use std::sync::Arc; + +use anyhow::Context; +use anyhow::Result; +use codex_exec_server::Environment; +use codex_exec_server::ExecBackend; +#[cfg(unix)] +use codex_exec_server::ExecEnvPolicy; +use codex_exec_server::ExecOutputStream; +use codex_exec_server::ExecParams; +use codex_exec_server::ExecProcess; +use codex_exec_server::ExecProcessEvent; +#[cfg(any(unix, windows))] +use codex_exec_server::FileSystemSandboxContext; +use codex_exec_server::ProcessId; +use codex_exec_server::ProcessSignal; +use codex_exec_server::ReadResponse; +#[cfg(unix)] +use codex_exec_server::ShellInfo; +#[cfg(unix)] +use codex_exec_server::ShellSnapshotRequest; +use codex_exec_server::StartedExecProcess; +#[cfg(any(unix, windows))] +use codex_exec_server::WindowsSandboxSelection; +use codex_exec_server::WriteStatus; +#[cfg(unix)] +use codex_network_proxy::NetworkProxyConfig; +#[cfg(unix)] +use codex_network_proxy::RemoteNetworkProxyConfig; +#[cfg(unix)] +use codex_network_proxy::RemoteNetworkProxyLaunchConfig; +#[cfg(unix)] +use codex_protocol::config_types::ShellEnvironmentPolicyInherit; +#[cfg(unix)] +use codex_protocol::models::PermissionProfile; +#[cfg(unix)] +use codex_protocol::permissions::FileSystemAccessMode; +#[cfg(unix)] +use codex_protocol::permissions::FileSystemPath; +#[cfg(unix)] +use codex_protocol::permissions::FileSystemSandboxEntry; +#[cfg(unix)] +use codex_protocol::permissions::FileSystemSandboxPolicy; +#[cfg(unix)] +use codex_protocol::permissions::FileSystemSpecialPath; +#[cfg(unix)] +use codex_protocol::permissions::NetworkSandboxPolicy; +use codex_protocol::protocol::SandboxPolicy; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +use tempfile::TempDir; +use test_case::test_case; +use tokio::sync::watch; +use tokio::time::Duration; +use tokio::time::sleep; +use tokio::time::timeout; + +use common::DELAYED_OUTPUT_AFTER_EXIT_PARENT_ARG; +use common::current_test_binary_helper_paths; +use common::exec_server::ExecServerHarness; +use common::exec_server::exec_server; + +struct ProcessContext { + backend: Arc, + _server: Option, +} + +#[derive(Debug, PartialEq, Eq)] +enum ProcessEventSnapshot { + Output { + seq: u64, + stream: ExecOutputStream, + text: String, + }, + Exited { + seq: u64, + exit_code: i32, + }, + Closed { + seq: u64, + }, +} + +async fn create_process_context(use_remote: bool) -> Result { + if use_remote { + let server = exec_server().await?; + let environment = Environment::create_for_tests(Some(server.websocket_url().to_string()))?; + Ok(ProcessContext { + backend: environment.get_exec_backend(), + _server: Some(server), + }) + } else { + let environment = Environment::create_for_tests(/*exec_server_url*/ None)?; + Ok(ProcessContext { + backend: environment.get_exec_backend(), + _server: None, + }) + } +} + +#[cfg(target_os = "macos")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn codex_home_symlink_opt_out_respects_host_config_and_scope() -> Result<()> { + use codex_exec_server::WriteFileOptions; + use common::exec_server::exec_server_with_env; + use std::os::unix::fs::symlink; + + let workspace = TempDir::new()?; + let home = TempDir::new()?; + let target = TempDir::new()?; + let alias = home.path().join("visualizations"); + let other_alias = workspace.path().join(".codex/visualizations"); + std::fs::create_dir(workspace.path().join(".codex"))?; + symlink(target.path(), &alias)?; + symlink(target.path(), &other_alias)?; + std::fs::write( + workspace.path().join(".codex/config.toml"), + "allow_symlinked_codex_home = true\n", + )?; + + for enabled in [None, Some(false), Some(true)] { + std::fs::write( + home.path().join("config.toml"), + enabled.map_or_else(String::new, |enabled| { + format!("allow_symlinked_codex_home = {enabled}\n") + }), + )?; + let mut server = exec_server_with_env([("CODEX_HOME", home.path())], &[]).await?; + let environment = Environment::create_for_tests(Some(server.websocket_url().to_string()))?; + for root in [alias.as_path(), other_alias.as_path(), workspace.path()] { + let mut policy = FileSystemSandboxPolicy::read_only(); + policy.entries.push(FileSystemSandboxEntry::new( + PathUri::from_host_native_path(root)?.into(), + FileSystemAccessMode::Write, + )); + let sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions( + &policy, + NetworkSandboxPolicy::Restricted, + ), + PathUri::from_host_native_path(workspace.path())?, + ); + let result = environment + .get_filesystem() + .write_file( + &PathUri::from_host_native_path(root.join("output"))?, + b"written".to_vec(), + WriteFileOptions::default(), + Some(&sandbox), + ) + .await; + assert_eq!( + result.is_ok(), + root == workspace.path() || (enabled == Some(true) && root == alias), + "root={root:?}, enabled={enabled:?}: {result:?}" + ); + if let Err(error) = result { + assert!( + error + .to_string() + .contains("symlinked writable roots are not supported"), + "{error}" + ); + } + } + server.shutdown().await?; + } + assert_eq!(std::fs::read(target.path().join("output"))?, b"written"); + assert_eq!(std::fs::read(workspace.path().join("output"))?, b"written"); + Ok(()) +} + +#[cfg(unix)] +#[test_case(false, false, false, false, "bash"; "local_pipe")] +#[test_case(false, true, false, false, "bash"; "local_tty")] +#[test_case(true, false, false, false, "bash"; "remote_pipe")] +#[test_case(true, true, false, false, "bash"; "remote_tty")] +#[test_case(true, false, true, false, "bash"; "remote_sandbox")] +#[test_case(false, false, false, false, "sh"; "local_sh_pipe")] +#[test_case(false, false, false, false, "bash-sh"; "local_bash_backed_sh")] +#[test_case(false, false, false, true, "bash"; "local_bash_env")] +#[test_case(true, false, false, true, "bash"; "remote_bash_env")] +#[cfg_attr( + target_os = "macos", + test_case(false, false, false, false, "zsh"; "local_zsh_pipe") +)] +#[cfg_attr( + target_os = "macos", + test_case(false, false, false, true, "zsh"; "local_zshenv") +)] +#[cfg_attr( + target_os = "macos", + test_case(true, false, false, true, "zsh"; "remote_zshenv") +)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn shell_snapshot_v2_filters_profile_exports_and_stays_in_memory( + use_remote: bool, + tty: bool, + use_sandbox: bool, + automatic_startup: bool, + shell_name: &str, +) -> Result<()> { + if use_sandbox + && let Some(warning) = + codex_sandboxing::system_bwrap_warning(&PermissionProfile::read_only()) + { + eprintln!("skipping sandbox test: {warning}"); + return Ok(()); + } + let context = create_process_context(use_remote).await?; + let home = TempDir::new()?; + let cwd = PathUri::from_host_native_path(home.path())?; + let (shell_path, profile_name) = match shell_name { + "bash" if automatic_startup => ("/bin/bash", ".bash-env"), + "bash" => ("/bin/bash", ".bashrc"), + "sh" => ("/bin/sh", ".snapshot-env"), + "bash-sh" => ("/bin/bash", ".snapshot-env"), + "zsh" if automatic_startup => ("/bin/zsh", ".zshenv"), + "zsh" => ("/bin/zsh", ".zshrc"), + name => anyhow::bail!("unsupported test shell {name}"), + }; + let profile_path = home.path().join(profile_name); + let profile_path_entry = home.path().join("profile-bin"); + let runtime_path_entry = home.path().join("runtime-bin"); + std::fs::create_dir(&profile_path_entry)?; + let wc = profile_path_entry.join("wc"); + std::fs::write( + &wc, + "#!/bin/sh\nprintf x >> \"$HOME/tool-captures\"\nexec /usr/bin/wc \"$@\"\n", + )?; + std::fs::set_permissions(&wc, std::fs::Permissions::from_mode(0o755))?; + let posix_shell = matches!(shell_name, "sh" | "bash-sh"); + let padding = if !use_remote && !tty && shell_name == "bash" { + format!( + "snapshot_padding() {{ printf '%s' '{}'; }}\n", + "🦀".repeat(20_000) + ) + } else { + String::new() + }; + let shadowed_builtins = if posix_shell { + "" + } else { + "unset() { exit 41; }\nbuiltin() { :; }\n" + }; + std::fs::write( + &profile_path, + format!( + "printf x >> \"$HOME/captures\"\nexport PATH=\"$HOME/profile-bin:/usr/bin:/bin\"\nexport PROFILE_ALLOWED=profile\nexport PROFILE_SECRET=secret\nexport PROFILE_DENIED=denied\nprofile_helper() {{ printf helper; }}\nif [ -n \"${{BASH_VERSION-}}\" ]; then\n shopt -s extglob nocasematch\n eval 'profile_helper() {{ case $1 in @(foo|bar)*) printf helper ;; *) return 1 ;; esac; }}'\nfi\nset -u\n{shadowed_builtins}{padding}" + ), + )?; + if shell_name == "zsh" && automatic_startup { + std::fs::write( + home.path().join(".zshrc"), + "export PATH=\"$HOME/profile-bin:/usr/bin:/bin\"\n", + )?; + } + let mut configured_environment = HashMap::from([( + "HOME".to_string(), + home.path().to_string_lossy().into_owned(), + )]); + if posix_shell { + configured_environment.insert( + "ENV".to_string(), + "${XDG_CONFIG_HOME:-$HOME}/.snapshot-env".to_string(), + ); + // Keep coverage for large values alongside the many-small-entry case below. + for index in 0..3 { + configured_environment.insert(format!("PROFILE_SDK_{index}"), "x".repeat(60 * 1024)); + } + } + if shell_name == "bash" && automatic_startup { + configured_environment.insert( + "BASH_ENV".to_string(), + profile_path.to_string_lossy().into_owned(), + ); + } + // Many small entries exercise capture overhead separately from the byte limit above. + let many_entries = !use_remote && !tty && !automatic_startup; + if many_entries { + configured_environment.extend( + (0..1_000).map(|index| (format!("PROFILE_ENTRY_{index}"), format!("value-{index}"))), + ); + } + let policy = ExecEnvPolicy { + inherit: ShellEnvironmentPolicyInherit::All, + ignore_default_excludes: false, + exclude: vec!["PROFILE_DENIED".to_string()], + r#set: configured_environment, + include_only: vec![ + "BASH_ENV".to_string(), + "ENV".to_string(), + "HOME".to_string(), + "PATH".to_string(), + "PROFILE_*".to_string(), + ], + }; + let (command_prefix, expected_prefix) = if shell_name == "sh" { + ("", "") + } else { + ("profile_helper FOObar; ", "helper") + }; + let entry_check = if many_entries { + "[ \"${PROFILE_ENTRY_999-missing}\" = value-999 ] || exit 43; " + } else { + "" + }; + let command = format!( + "case $- in *u*) ;; *) exit 42 ;; esac; {entry_check}export PATH='{}':\"$PATH\"; {command_prefix}printf '|%s|%s|%s|%s|%s|%s' \"$PROFILE_ALLOWED\" \"${{PROFILE_SECRET-missing}}\" \"${{PROFILE_DENIED-missing}}\" \"$PATH\" \"${{__CODEX_SHELL_SNAPSHOT_STATE_0-missing}}\" \"${{__CODEX_SHELL_SNAPSHOT_STATE_1-missing}}\"", + runtime_path_entry.display(), + ); + let expected_stdout = format!( + "{expected_prefix}|profile|missing|missing|{}:{}:/usr/bin:/bin|missing|missing", + runtime_path_entry.display(), + profile_path_entry.display(), + ); + + for attempt in 0..2 { + let started = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from(format!("snapshot-{attempt}")), + argv: vec![shell_path.to_string(), "-lc".to_string(), command.clone()], + cwd: cwd.clone(), + env_policy: Some(policy.clone()), + shell_snapshot: Some(ShellSnapshotRequest { + scope_id: "attachment-1".to_string(), + shell: ShellInfo { + name: if posix_shell { "sh" } else { shell_name }.to_string(), + path: shell_path.to_string(), + }, + }), + env: HashMap::new(), + tty, + pipe_stdin: false, + arg0: (shell_name == "bash-sh").then(|| "sh".to_string()), + sandbox: (use_sandbox && attempt == 0).then(|| { + FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::read_only(), + cwd.clone(), + ) + }), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + let (stdout, stderr, status, closed) = + collect_process_output_from_events(started.process).await?; + assert_eq!( + (stdout, stderr, status, closed), + (expected_stdout.clone(), String::new(), Some(0), true,) + ); + } + + assert_eq!(std::fs::read_to_string(home.path().join("captures"))?, "x"); + assert!(!std::fs::read(home.path().join("tool-captures"))?.is_empty()); + if let Some(server) = context._server { + assert!(!server.codex_home().join("shell_snapshots").exists()); + } + Ok(()) +} + +#[cfg(unix)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial_test::serial(remote_exec_server)] +async fn shell_snapshot_v2_remote_managed_proxy_uses_prepared_execution_context() -> Result<()> { + let context = create_process_context(/*use_remote*/ true).await?; + let home = TempDir::new()?; + let cwd = PathUri::from_host_native_path(home.path())?; + std::fs::write( + home.path().join(".bashrc"), + "printf '%s\\n' \"$HTTP_PROXY\" >> \"$HOME/captures\"\ntest \"$CODEX_NETWORK_PROXY_ACTIVE\" = 1 || exit 41\nexport PROFILE_ALLOWED=profile\nprofile_helper() { printf helper; }\n", + )?; + let policy = ExecEnvPolicy { + inherit: ShellEnvironmentPolicyInherit::All, + ignore_default_excludes: false, + exclude: Vec::new(), + r#set: HashMap::from([( + "HOME".to_string(), + home.path().to_string_lossy().into_owned(), + )]), + include_only: vec![ + "HOME".to_string(), + "PATH".to_string(), + "PROFILE_*".to_string(), + ], + }; + let proxy_config = RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig { + enabled: true, + ..NetworkProxyConfig::default() + })?; + let mut proxy_addresses = Vec::new(); + + for attempt in 0..2 { + let started = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from(format!("managed-snapshot-{attempt}")), + argv: vec![ + "/bin/bash".to_string(), + "-lc".to_string(), + "profile_helper; printf '|%s|%s|%s' \"$PROFILE_ALLOWED\" \"$CODEX_NETWORK_PROXY_ACTIVE\" \"$HTTP_PROXY\"".to_string(), + ], + cwd: cwd.clone(), + env_policy: Some(policy.clone()), + shell_snapshot: Some(ShellSnapshotRequest { + scope_id: "managed-attachment".to_string(), + shell: ShellInfo { + name: "bash".to_string(), + path: "/bin/bash".to_string(), + }, + }), + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: true, + managed_network: None, + network_proxy: Some( + RemoteNetworkProxyLaunchConfig::new(proxy_config.clone()).for_execution( + "remote-environment".to_string(), + format!("managed-snapshot-{attempt}"), + ), + ), + }) + .await?; + let (stdout, stderr, status, closed) = + collect_process_output_from_events(started.process).await?; + let proxy_address = stdout + .strip_prefix("helper|profile|1|") + .context("snapshot should restore profile functions and live proxy state")?; + assert!(proxy_address.starts_with("http://127.0.0.1:")); + assert_eq!((stderr, status, closed), (String::new(), Some(0), true)); + proxy_addresses.push(proxy_address.to_string()); + } + + assert_eq!( + std::fs::read_to_string(home.path().join("captures"))?, + format!("{}\n", proxy_addresses[0]) + ); + Ok(()) +} + +#[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"))] +#[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, +) -> Result<()> { + if use_remote + && let Some(warning) = + codex_sandboxing::system_bwrap_warning(&PermissionProfile::workspace_write()) + { + eprintln!("skipping sandbox test: {warning}"); + return Ok(()); + } + let context = create_process_context(use_remote).await?; + let home = TempDir::new()?; + let cwd = PathUri::from_host_native_path(home.path())?; + let (shell_path, profile_name) = match shell_name { + "bash" => ("/bin/bash", ".bashrc"), + "zsh" => ("/bin/zsh", ".zshrc"), + name => anyhow::bail!("unsupported test shell {name}"), + }; + std::fs::write( + home.path().join(profile_name), + "printf x >> \"$HOME/captures\"\nexit 7\n", + )?; + let policy = ExecEnvPolicy { + inherit: ShellEnvironmentPolicyInherit::All, + ignore_default_excludes: false, + exclude: Vec::new(), + r#set: HashMap::from([( + "HOME".to_string(), + home.path().to_string_lossy().into_owned(), + )]), + include_only: vec!["HOME".to_string(), "PATH".to_string()], + }; + let mut params = ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("snapshot-first"), + argv: vec![ + shell_path.to_string(), + "-lc".to_string(), + "if command -v profile_helper >/dev/null; then profile_helper; else printf original; fi".to_string(), + ], + cwd: cwd.clone(), + env_policy: Some(policy), + shell_snapshot: Some(ShellSnapshotRequest { + scope_id: "attachment-1".to_string(), + shell: ShellInfo { + name: shell_name.to_string(), + path: shell_path.to_string(), + }, + }), + env: HashMap::new(), + tty, + pipe_stdin: false, + arg0: None, + sandbox: use_remote.then(|| { + FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::workspace_write(), + cwd, + ) + }), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }; + + 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) + ); + // 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; + } + assert_eq!( + std::fs::read_to_string(home.path().join("captures"))?, + "x".repeat(failures_before_repair) + ); + + std::fs::write( + home.path().join(profile_name), + "printf x >> \"$HOME/captures\"\nprofile_helper() { printf recovered; }\n", + )?; + let (expected_output, expected_captures) = if failures_before_repair == 3 { + ("original", "xxx") + } else { + ("recovered", "xx") + }; + for attempt in 0..2 { + params.process_id = ProcessId::from(format!("snapshot-after-repair-{attempt}")); + let started = context.backend.start(params.clone()).await?; + assert_eq!( + collect_process_output_from_events(started.process).await?, + (expected_output.to_string(), String::new(), Some(0), true) + ); + } + assert_eq!( + std::fs::read_to_string(home.path().join("captures"))?, + expected_captures + ); + Ok(()) +} + +#[cfg(unix)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn remote_sandboxed_process_preserves_custom_arg0() -> Result<()> { + if let Some(warning) = codex_sandboxing::system_bwrap_warning(&PermissionProfile::read_only()) { + eprintln!("skipping bwrap test: {warning}"); + return Ok(()); + } + + let context = create_process_context(/*use_remote*/ true).await?; + let workspace = TempDir::new()?; + let outside_workspace = TempDir::new()?; + let denied_file = outside_workspace.path().join("denied.txt"); + std::fs::write(&denied_file, b"denied")?; + let cwd = PathUri::from_host_native_path(workspace.path())?; + let policy = FileSystemSandboxPolicy::restricted(vec![ + FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::Minimal, + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }, + FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }, + ]); + let sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + cwd.clone(), + ); + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-custom-arg0"), + argv: vec![ + "/bin/sh".to_string(), + "-c".to_string(), + "printf '%s' \"$0\"; if /bin/cat \"$CODEX_TEST_DENIED_FILE\" >/dev/null 2>&1; then exit 42; fi" + .to_string(), + ], + cwd, + shell_snapshot: None, + env_policy: None, + env: HashMap::from([ + ("PATH".to_string(), std::env::var("PATH")?), + ( + "CODEX_TEST_DENIED_FILE".to_string(), + denied_file.to_string_lossy().into_owned(), + ), + ]), + tty: false, + pipe_stdin: false, + arg0: Some("custom-arg0".to_string()), + sandbox: Some(sandbox), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + let output = collect_process_output_from_events(session.process).await?; + + assert_eq!( + output, + ("custom-arg0".to_string(), String::new(), Some(0), true) + ); + Ok(()) +} + +async fn assert_exec_process_starts_and_exits(use_remote: bool) -> Result<()> { + let context = create_process_context(use_remote).await?; + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-1"), + argv: vec!["true".to_string()], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + assert_eq!(session.process.process_id().as_str(), "proc-1"); + let wake_rx = session.process.subscribe_wake(); + let (_, exit_code, closed) = + collect_process_output_from_reads(session.process, wake_rx).await?; + + assert_eq!(exit_code, Some(0)); + assert!(closed); + Ok(()) +} + +#[cfg(target_os = "linux")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn remote_process_keeps_sandbox_helper_visible_with_restricted_reads() -> Result<()> { + if let Some(warning) = codex_sandboxing::system_bwrap_warning(&PermissionProfile::read_only()) { + eprintln!("skipping bwrap test: {warning}"); + return Ok(()); + } + + let context = create_process_context(/*use_remote*/ true).await?; + let workspace = TempDir::new()?; + let file = workspace.path().join("allowed.txt"); + std::fs::write(&file, b"allowed")?; + let cwd = PathUri::from_host_native_path(workspace.path())?; + let policy = FileSystemSandboxPolicy::restricted(vec![ + FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::Minimal, + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }, + FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }, + ]); + let sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + cwd.clone(), + ); + + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-restricted-helper"), + argv: vec!["/bin/cat".to_string(), file.to_string_lossy().into_owned()], + cwd, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: HashMap::from([("PATH".to_string(), std::env::var("PATH")?)]), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: Some(sandbox), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + let output = collect_process_output_from_events(session.process).await?; + + assert_eq!( + output, + ("allowed".to_string(), String::new(), Some(0), true) + ); + Ok(()) +} + +#[cfg(target_os = "linux")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn remote_tty_process_uses_configured_sandbox_helper_with_hostile_path() -> Result<()> { + if let Some(warning) = codex_sandboxing::system_bwrap_warning(&PermissionProfile::read_only()) { + eprintln!("skipping bwrap test: {warning}"); + return Ok(()); + } + + let context = create_process_context(/*use_remote*/ true).await?; + let workspace = TempDir::new()?; + let file = workspace.path().join("allowed.txt"); + std::fs::write(&file, b"allowed")?; + let hostile_helper = workspace.path().join("codex-linux-sandbox"); + std::fs::write(&hostile_helper, b"#!/bin/sh\nprintf hostile")?; + let mut permissions = std::fs::metadata(&hostile_helper)?.permissions(); + permissions.set_mode(0o755); + std::fs::set_permissions(&hostile_helper, permissions)?; + let path = std::env::var_os("PATH").context("PATH is not set")?; + let hostile_path = std::env::join_paths( + std::iter::once(workspace.path().to_path_buf()).chain(std::env::split_paths(&path)), + )?; + let cwd = PathUri::from_host_native_path(workspace.path())?; + let policy = FileSystemSandboxPolicy::restricted(vec![ + FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::Minimal, + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }, + FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }, + ]); + let sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + cwd.clone(), + ); + + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-hostile-helper-path"), + argv: vec!["/bin/cat".to_string(), file.to_string_lossy().into_owned()], + cwd, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: HashMap::from([( + "PATH".to_string(), + hostile_path.to_string_lossy().into_owned(), + )]), + tty: true, + pipe_stdin: false, + arg0: None, + sandbox: Some(sandbox), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + let output = collect_process_output_from_events(session.process).await?; + + assert_eq!( + output, + ("allowed".to_string(), String::new(), Some(0), true) + ); + Ok(()) +} + +#[cfg(unix)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn remote_process_preserves_empty_workspace_roots() -> Result<()> { + if let Some(warning) = codex_sandboxing::system_bwrap_warning(&PermissionProfile::read_only()) { + eprintln!("skipping bwrap test: {warning}"); + return Ok(()); + } + + let context = create_process_context(/*use_remote*/ true).await?; + let tmp = TempDir::new()?; + let file = tmp.path().join("excluded.txt"); + std::fs::write(&file, b"excluded")?; + let cwd = PathUri::from_host_native_path(tmp.path())?; + let policy = FileSystemSandboxPolicy::restricted(vec![FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }]); + let mut sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + cwd.clone(), + ); + sandbox.workspace_roots.clear(); + + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-empty-workspace-roots"), + argv: vec!["/bin/cat".to_string(), file.to_string_lossy().into_owned()], + cwd, + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: Some(sandbox), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + let (stdout, _stderr, exit_code, closed) = + collect_process_output_from_events(session.process).await?; + + assert!(!stdout.contains("excluded"), "unexpected stdout: {stdout}"); + assert_ne!(exit_code, Some(0)); + assert!(closed); + Ok(()) +} + +async fn read_process_until_change( + session: Arc, + wake_rx: &mut watch::Receiver, + after_seq: Option, +) -> Result { + let response = session + .read(after_seq, /*max_bytes*/ None, /*wait_ms*/ Some(0)) + .await?; + if !response.chunks.is_empty() || response.closed || response.failure.is_some() { + return Ok(response); + } + + timeout(Duration::from_secs(2), wake_rx.changed()).await??; + session + .read(after_seq, /*max_bytes*/ None, /*wait_ms*/ Some(0)) + .await + .map_err(Into::into) +} + +async fn collect_process_output_from_reads( + session: Arc, + mut wake_rx: watch::Receiver, +) -> Result<(String, Option, bool)> { + let mut output = String::new(); + let mut exit_code = None; + let mut after_seq = None; + loop { + let response = + read_process_until_change(Arc::clone(&session), &mut wake_rx, after_seq).await?; + if let Some(message) = response.failure { + anyhow::bail!("process failed before closed state: {message}"); + } + for chunk in response.chunks { + output.push_str(&String::from_utf8_lossy(&chunk.chunk.into_inner())); + after_seq = Some(chunk.seq); + } + if response.exited { + exit_code = response.exit_code; + } + if response.closed { + break; + } + after_seq = response.next_seq.checked_sub(1).or(after_seq); + } + drop(session); + Ok((output, exit_code, true)) +} + +async fn collect_process_output_from_events( + session: Arc, +) -> Result<(String, String, Option, bool)> { + collect_process_output_from_events_with_timeout(session, Duration::from_secs(2)).await +} + +async fn collect_process_output_from_events_with_timeout( + session: Arc, + event_timeout: Duration, +) -> Result<(String, String, Option, bool)> { + let mut events = session.subscribe_events(); + let mut stdout = String::new(); + let mut stderr = String::new(); + let mut exit_code = None; + loop { + match timeout(event_timeout, events.recv()).await?? { + ExecProcessEvent::Output(chunk) => match chunk.stream { + ExecOutputStream::Stdout | ExecOutputStream::Pty => { + stdout.push_str(&String::from_utf8_lossy(&chunk.chunk.into_inner())); + } + ExecOutputStream::Stderr => { + stderr.push_str(&String::from_utf8_lossy(&chunk.chunk.into_inner())); + } + }, + ExecProcessEvent::Exited { + seq: _, + exit_code: code, + .. + } => { + exit_code = Some(code); + } + ExecProcessEvent::Closed { seq: _ } => { + drop(session); + return Ok((stdout, stderr, exit_code, true)); + } + ExecProcessEvent::Failed(message) => { + anyhow::bail!("process failed before closed state: {message}"); + } + } + } +} + +async fn collect_process_event_snapshots( + session: Arc, +) -> Result> { + let mut events = session.subscribe_events(); + let mut snapshots = Vec::new(); + loop { + let snapshot = match timeout(Duration::from_secs(2), events.recv()).await?? { + ExecProcessEvent::Output(chunk) => ProcessEventSnapshot::Output { + seq: chunk.seq, + stream: chunk.stream, + text: String::from_utf8_lossy(&chunk.chunk.into_inner()).into_owned(), + }, + ExecProcessEvent::Exited { seq, exit_code, .. } => { + ProcessEventSnapshot::Exited { seq, exit_code } + } + ExecProcessEvent::Closed { seq } => ProcessEventSnapshot::Closed { seq }, + ExecProcessEvent::Failed(message) => { + anyhow::bail!("process failed before closed state: {message}"); + } + }; + let closed = matches!(snapshot, ProcessEventSnapshot::Closed { .. }); + snapshots.push(snapshot); + if closed { + drop(session); + return Ok(snapshots); + } + } +} + +async fn assert_exec_process_streams_output(use_remote: bool) -> Result<()> { + let context = create_process_context(use_remote).await?; + let process_id = "proc-stream".to_string(); + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: process_id.clone().into(), + argv: vec![ + "/bin/sh".to_string(), + "-c".to_string(), + "sleep 0.05; printf 'session output\\n'".to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + assert_eq!(session.process.process_id().as_str(), process_id); + + let StartedExecProcess { process, .. } = session; + let wake_rx = process.subscribe_wake(); + let (output, exit_code, closed) = collect_process_output_from_reads(process, wake_rx).await?; + assert_eq!(output, "session output\n"); + assert_eq!(exit_code, Some(0)); + assert!(closed); + Ok(()) +} + +async fn assert_exec_process_pushes_events(use_remote: bool) -> Result<()> { + let context = create_process_context(use_remote).await?; + let process_id = "proc-events".to_string(); + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: process_id.clone().into(), + argv: vec![ + "/bin/sh".to_string(), + "-c".to_string(), + "printf 'event output\\n'; sleep 0.1; printf 'event err\\n' >&2; sleep 0.1; exit 7".to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + assert_eq!(session.process.process_id().as_str(), process_id); + + let StartedExecProcess { process, .. } = session; + let actual = collect_process_event_snapshots(process).await?; + assert_eq!( + actual, + vec![ + ProcessEventSnapshot::Output { + seq: 1, + stream: ExecOutputStream::Stdout, + text: "event output\n".to_string(), + }, + ProcessEventSnapshot::Output { + seq: 2, + stream: ExecOutputStream::Stderr, + text: "event err\n".to_string(), + }, + ProcessEventSnapshot::Exited { + seq: 3, + exit_code: 7, + }, + ProcessEventSnapshot::Closed { seq: 4 }, + ] + ); + Ok(()) +} + +async fn assert_exec_process_replays_events_after_close(use_remote: bool) -> Result<()> { + let context = create_process_context(use_remote).await?; + let process_id = "proc-events-late".to_string(); + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: process_id.clone().into(), + argv: vec![ + "/bin/sh".to_string(), + "-c".to_string(), + "printf 'late one\\n'; printf 'late two\\n'".to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + assert_eq!(session.process.process_id().as_str(), process_id); + + let StartedExecProcess { process, .. } = session; + let wake_rx = process.subscribe_wake(); + let read_result = collect_process_output_from_reads(Arc::clone(&process), wake_rx).await?; + assert_eq!( + read_result, + ("late one\nlate two\n".to_string(), Some(0), true) + ); + + let event_result = collect_process_output_from_events(process).await?; + assert_eq!( + event_result, + ( + "late one\nlate two\n".to_string(), + String::new(), + Some(0), + true + ) + ); + Ok(()) +} + +async fn assert_exec_process_retains_output_after_exit_until_streams_close( + use_remote: bool, +) -> Result<()> { + let context = create_process_context(use_remote).await?; + let (helper_binary, _) = current_test_binary_helper_paths()?; + let release_dir = TempDir::new()?; + let release_path = release_dir.path().join("release-delayed-output"); + let process_id = "proc-output-after-exit".to_string(); + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: process_id.clone().into(), + argv: vec![ + helper_binary.to_string_lossy().into_owned(), + DELAYED_OUTPUT_AFTER_EXIT_PARENT_ARG.to_string(), + release_path.to_string_lossy().into_owned(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + assert_eq!(session.process.process_id().as_str(), process_id); + + let StartedExecProcess { process, .. } = session; + + let exit_response = timeout( + Duration::from_secs(2), + process.read( + /*after_seq*/ None, + /*max_bytes*/ None, + /*wait_ms*/ Some(2_000), + ), + ) + .await??; + assert!( + exit_response.chunks.is_empty(), + "parent should exit before child writes delayed output" + ); + assert_eq!(exit_response.exit_code, Some(0)); + assert!(!exit_response.closed); + let exit_seq = exit_response + .next_seq + .checked_sub(1) + .context("exit response should advance next_seq")?; + std::fs::write(&release_path, b"go")?; + + let late_response = timeout( + Duration::from_secs(2), + process.read( + /*after_seq*/ Some(exit_seq), + /*max_bytes*/ None, + /*wait_ms*/ Some(2_000), + ), + ) + .await??; + let mut late_output = String::new(); + for chunk in late_response.chunks { + assert_eq!(chunk.stream, ExecOutputStream::Stdout); + late_output.push_str(&String::from_utf8_lossy(&chunk.chunk.into_inner())); + } + assert_eq!(late_output, "late output after exit\n"); + + let wake_rx = process.subscribe_wake(); + let actual = collect_process_output_from_reads(process, wake_rx).await?; + assert_eq!( + actual, + ("late output after exit\n".to_string(), Some(0), true) + ); + Ok(()) +} + +async fn assert_exec_process_write_then_read(use_remote: bool) -> Result<()> { + let context = create_process_context(use_remote).await?; + let process_id = "proc-stdin".to_string(); + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: process_id.clone().into(), + argv: vec![ + // Use `/bin/sh` instead of Python so this stdin round-trip test + // stays portable across Bazel and non-macOS runners where + // `/usr/bin/python3` is not guaranteed to exist. + "/bin/sh".to_string(), + "-c".to_string(), + "IFS= read line; printf 'from-stdin:%s\\n' \"$line\"".to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: true, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + assert_eq!(session.process.process_id().as_str(), process_id); + + tokio::time::sleep(Duration::from_millis(200)).await; + session.process.write(b"hello\n".to_vec()).await?; + let StartedExecProcess { process, .. } = session; + let wake_rx = process.subscribe_wake(); + let (output, exit_code, closed) = collect_process_output_from_reads(process, wake_rx).await?; + + assert!( + output.contains("from-stdin:hello"), + "unexpected output: {output:?}" + ); + assert_eq!(exit_code, Some(0)); + assert!(closed); + Ok(()) +} + +async fn assert_exec_process_write_then_read_without_tty(use_remote: bool) -> Result<()> { + let context = create_process_context(use_remote).await?; + let process_id = "proc-stdin-pipe".to_string(); + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: process_id.clone().into(), + argv: vec![ + "/bin/sh".to_string(), + "-c".to_string(), + "IFS= read line; printf 'from-stdin:%s\\n' \"$line\"".to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: false, + pipe_stdin: true, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + assert_eq!(session.process.process_id().as_str(), process_id); + + tokio::time::sleep(Duration::from_millis(200)).await; + let write_response = session.process.write(b"hello\n".to_vec()).await?; + assert_eq!(write_response.status, WriteStatus::Accepted); + let StartedExecProcess { process, .. } = session; + let wake_rx = process.subscribe_wake(); + let actual = collect_process_output_from_reads(process, wake_rx).await?; + + assert_eq!(actual, ("from-stdin:hello\n".to_string(), Some(0), true)); + Ok(()) +} + +async fn assert_remote_windows_sandbox_process_write( + expected_sandbox_type: codex_sandboxing::SandboxType, + tty: bool, +) -> Result<()> { + if expected_sandbox_type == codex_sandboxing::SandboxType::WindowsMxc { + crate::skip_if_mxc_unavailable!(Ok(())); + } + let context = create_process_context(/*use_remote*/ true).await?; + let workspace = TempDir::new()?; + let blocked_file = workspace.path().join("blocked.txt"); + let cwd = PathUri::from_host_native_path(workspace.path())?; + let mut sandbox = FileSystemSandboxContext::from_legacy_sandbox_policy( + SandboxPolicy::new_read_only_policy(), + cwd.clone(), + )?; + match expected_sandbox_type { + codex_sandboxing::SandboxType::WindowsRestrictedToken => { + sandbox.windows_sandbox_selection = WindowsSandboxSelection::RestrictedToken; + } + codex_sandboxing::SandboxType::WindowsMxc => { + sandbox.windows_sandbox_selection = WindowsSandboxSelection::Mxc; + } + codex_sandboxing::SandboxType::None + | codex_sandboxing::SandboxType::MacosSeatbelt + | codex_sandboxing::SandboxType::LinuxSeccomp => { + anyhow::bail!("expected a Windows sandbox type") + } + } + + let session = match context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-windows-sandbox-stdin"), + argv: vec![ + r"C:\Windows\System32\cmd.exe".to_string(), + "/D".to_string(), + "/V:ON".to_string(), + "/S".to_string(), + "/C".to_string(), + format!( + "set /P line= & echo blocked > \"{}\" & echo from-stdin:!line!", + blocked_file.display() + ), + ], + cwd, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty, + pipe_stdin: !tty, + arg0: None, + sandbox: Some(sandbox), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await + { + Ok(session) => session, + Err(err) => return Err(err.into()), + }; + assert_eq!(session.sandbox_type, Some(expected_sandbox_type)); + + let input = if tty { b"hello\r" } else { b"hello\n" }; + let write_response = session.process.write(input.to_vec()).await?; + assert_eq!(write_response.status, WriteStatus::Accepted); + let StartedExecProcess { process, .. } = session; + let wake_rx = process.subscribe_wake(); + let (output, exit_code, closed) = collect_process_output_from_reads(process, wake_rx).await?; + + assert!( + output.contains("from-stdin:hello"), + "unexpected output: {output:?}" + ); + assert_eq!(exit_code, Some(0)); + assert!(closed); + assert!(!blocked_file.exists()); + Ok(()) +} + +async fn assert_exec_process_rejects_write_without_pipe_stdin(use_remote: bool) -> Result<()> { + let context = create_process_context(use_remote).await?; + let process_id = "proc-stdin-closed".to_string(); + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: process_id.clone().into(), + argv: vec![ + "/bin/sh".to_string(), + "-c".to_string(), + "sleep 0.3; if IFS= read -r line; then printf 'read:%s\\n' \"$line\"; else printf 'eof\\n'; fi".to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + assert_eq!(session.process.process_id().as_str(), process_id); + + let write_response = session.process.write(b"ignored\n".to_vec()).await?; + assert_eq!(write_response.status, WriteStatus::StdinClosed); + let StartedExecProcess { process, .. } = session; + let wake_rx = process.subscribe_wake(); + let (output, exit_code, closed) = collect_process_output_from_reads(process, wake_rx).await?; + + assert_eq!(output, "eof\n"); + assert_eq!(exit_code, Some(0)); + assert!(closed); + Ok(()) +} + +async fn assert_exec_process_signal_interrupts_process(use_remote: bool) -> Result<()> { + let context = create_process_context(use_remote).await?; + let process_id = "proc-signal".to_string(); + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: process_id.clone().into(), + argv: vec![ + "/bin/sh".to_string(), + "-c".to_string(), + "trap 'printf \"signal:2\\n\"; exit 7' INT; printf 'ready\\n'; while :; do :; done".to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + assert_eq!(session.process.process_id().as_str(), process_id); + + let StartedExecProcess { process, .. } = session; + let mut wake_rx = process.subscribe_wake(); + let mut ready_output = String::new(); + let mut after_seq = None; + loop { + let response = + read_process_until_change(Arc::clone(&process), &mut wake_rx, after_seq).await?; + for chunk in response.chunks { + ready_output.push_str(&String::from_utf8_lossy(&chunk.chunk.into_inner())); + after_seq = Some(chunk.seq); + } + if ready_output.contains("ready\n") { + break; + } + if response.closed { + anyhow::bail!("process closed before readiness marker: {ready_output:?}"); + } + after_seq = response.next_seq.checked_sub(1).or(after_seq); + } + + process.signal(ProcessSignal::Interrupt).await?; + let (output, exit_code, closed) = collect_process_output_from_reads(process, wake_rx).await?; + + assert!( + output.contains("signal:2"), + "expected signal handler output, got {output:?}" + ); + assert_eq!(exit_code, Some(7)); + assert!(closed); + Ok(()) +} + +async fn assert_exec_process_signal_terminates_on_windows(use_remote: bool) -> Result<()> { + let context = create_process_context(use_remote).await?; + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-windows-signal"), + argv: vec![ + "cmd".to_string(), + "/C".to_string(), + "echo ready && ping -n 30 127.0.0.1 >NUL".to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + + let StartedExecProcess { process, .. } = session; + let wake_rx = process.subscribe_wake(); + process.signal(ProcessSignal::Interrupt).await?; + let (_output, exit_code, closed) = collect_process_output_from_reads(process, wake_rx).await?; + + assert_eq!(exit_code, Some(1)); + assert!(closed); + Ok(()) +} + +async fn assert_exec_process_preserves_queued_events_before_subscribe( + use_remote: bool, +) -> Result<()> { + let context = create_process_context(use_remote).await?; + let session = context + .backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-queued"), + argv: vec![ + "/bin/sh".to_string(), + "-c".to_string(), + "printf 'queued output\\n'".to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + + tokio::time::sleep(Duration::from_millis(200)).await; + + let StartedExecProcess { process, .. } = session; + let wake_rx = process.subscribe_wake(); + let (output, exit_code, closed) = collect_process_output_from_reads(process, wake_rx).await?; + assert_eq!(output, "queued output\n"); + assert_eq!(exit_code, Some(0)); + assert!(closed); + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn remote_exec_process_recovers_after_transport_disconnect() -> Result<()> { + let server = exec_server().await?; + let mut proxy = server.disconnectable_websocket_proxy().await?; + let environment = Environment::create_for_tests(Some(proxy.websocket_url().to_string()))?; + let backend = environment.get_exec_backend(); + let temp_dir = TempDir::new()?; + let gate_path = temp_dir.path().join("release-output"); + let emitted_path = temp_dir.path().join("output-emitted"); + let session = backend + .start(ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-recover"), + argv: vec![ + "/bin/sh".to_string(), + "-c".to_string(), + concat!( + "printf 'ready:%s\\n' \"$$\"; ", + "while [ ! -f \"$GATE\" ]; do /bin/sleep 0.01; done; ", + "printf 'during:%s\\n' \"$$\"; ", + ": > \"$EMITTED\"; ", + "IFS= read -r line; ", + "printf 'after:%s:%s\\n' \"$$\" \"$line\"; ", + "exit 7", + ) + .to_string(), + ], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: /*env_policy*/ None, + env: HashMap::from([ + ( + "GATE".to_string(), + gate_path.to_string_lossy().into_owned(), + ), + ( + "EMITTED".to_string(), + emitted_path.to_string_lossy().into_owned(), + ), + ]), + tty: false, + pipe_stdin: true, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + + let process = Arc::clone(&session.process); + let mut events = process.subscribe_events(); + let mut output = Vec::new(); + let mut last_seq = 0; + while !output.ends_with(b"\n") { + match timeout(Duration::from_secs(5), events.recv()).await?? { + ExecProcessEvent::Output(chunk) => { + assert_eq!(chunk.seq, last_seq + 1); + last_seq = chunk.seq; + output.extend_from_slice(&chunk.chunk.into_inner()); + } + event => anyhow::bail!("expected ready output before disconnect, got {event:?}"), + } + } + let ready = String::from_utf8(output.clone())?; + let pid = ready + .strip_prefix("ready:") + .and_then(|line| line.strip_suffix('\n')) + .context("ready output should contain the process id")? + .to_string(); + + proxy.pause_and_disconnect().await?; + tokio::fs::write(&gate_path, b"").await?; + timeout(Duration::from_secs(5), async { + while tokio::fs::metadata(&emitted_path).await.is_err() { + sleep(Duration::from_millis(10)).await; + } + }) + .await + .context("process did not emit output while disconnected")?; + + let process_for_read = Arc::clone(&process); + let mut pending_read = tokio::spawn(async move { + process_for_read + .read( + /*after_seq*/ Some(last_seq), + /*max_bytes*/ None, + /*wait_ms*/ Some(0), + ) + .await + }); + assert!( + timeout(Duration::from_millis(200), &mut pending_read) + .await + .is_err(), + "process reads should wait while recovery is in progress" + ); + proxy.resume()?; + + let recovered_read = timeout(Duration::from_secs(5), pending_read) + .await + .context("timed out waiting for a read after recovery")??; + let recovered_read = recovered_read?; + assert_eq!(recovered_read.failure, None); + let recovered_output = recovered_read + .chunks + .into_iter() + .flat_map(|chunk| chunk.chunk.into_inner()) + .collect::>(); + assert_eq!( + String::from_utf8(recovered_output)?, + format!("during:{pid}\n") + ); + + let write = timeout(Duration::from_secs(5), process.write(b"hello\n".to_vec())) + .await + .context("timed out waiting for a write after recovery")??; + assert_eq!(write.status, WriteStatus::Accepted); + + let mut saw_exit = false; + loop { + match timeout(Duration::from_secs(5), events.recv()).await?? { + ExecProcessEvent::Output(chunk) => { + assert_eq!(chunk.seq, last_seq + 1); + last_seq = chunk.seq; + output.extend_from_slice(&chunk.chunk.into_inner()); + } + ExecProcessEvent::Exited { seq, exit_code, .. } => { + assert_eq!(seq, last_seq + 1); + assert_eq!(exit_code, 7); + last_seq = seq; + saw_exit = true; + } + ExecProcessEvent::Closed { seq } => { + assert!(saw_exit, "closed must be delivered after exit"); + assert_eq!(seq, last_seq + 1); + break; + } + ExecProcessEvent::Failed(message) => { + anyhow::bail!("process recovery failed: {message}"); + } + } + } + assert_eq!( + String::from_utf8(output)?, + format!("ready:{pid}\nduring:{pid}\nafter:{pid}:hello\n") + ); + + Ok(()) +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_starts_and_exits(use_remote: bool) -> Result<()> { + assert_exec_process_starts_and_exits(use_remote).await +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_streams_output(use_remote: bool) -> Result<()> { + assert_exec_process_streams_output(use_remote).await +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_pushes_events(use_remote: bool) -> Result<()> { + assert_exec_process_pushes_events(use_remote).await +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_replays_events_after_close(use_remote: bool) -> Result<()> { + assert_exec_process_replays_events_after_close(use_remote).await +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_retains_output_after_exit_until_streams_close( + use_remote: bool, +) -> Result<()> { + assert_exec_process_retains_output_after_exit_until_streams_close(use_remote).await +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_write_then_read(use_remote: bool) -> Result<()> { + assert_exec_process_write_then_read(use_remote).await +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_write_then_read_without_tty(use_remote: bool) -> Result<()> { + assert_exec_process_write_then_read_without_tty(use_remote).await +} + +#[test_case( + codex_sandboxing::SandboxType::WindowsRestrictedToken, + false; + "restricted_token" +)] +#[test_case( + codex_sandboxing::SandboxType::WindowsMxc, + false; + "mxc_pipe" +)] +#[test_case( + codex_sandboxing::SandboxType::WindowsMxc, + true; + "mxc_conpty" +)] +#[cfg_attr(not(windows), ignore = "Windows-only exec-server sandbox process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial_test::serial(remote_exec_server)] +async fn remote_windows_sandbox_process_accepts_process_write( + expected_sandbox_type: codex_sandboxing::SandboxType, + tty: bool, +) -> Result<()> { + assert_remote_windows_sandbox_process_write(expected_sandbox_type, tty).await +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_rejects_write_without_pipe_stdin(use_remote: bool) -> Result<()> { + assert_exec_process_rejects_write_without_pipe_stdin(use_remote).await +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_signal_interrupts_process(use_remote: bool) -> Result<()> { + assert_exec_process_signal_interrupts_process(use_remote).await +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(windows), ignore = "Windows-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_signal_terminates_on_windows(use_remote: bool) -> Result<()> { + assert_exec_process_signal_terminates_on_windows(use_remote).await +} + +#[test_case(false ; "local")] +#[test_case(true ; "remote")] +#[cfg_attr(not(unix), ignore = "Unix-only exec-server process test")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +// Serialize tests that launch a real exec-server process through the full CLI. +#[serial_test::serial(remote_exec_server)] +async fn exec_process_preserves_queued_events_before_subscribe(use_remote: bool) -> Result<()> { + assert_exec_process_preserves_queued_events_before_subscribe(use_remote).await +} diff --git a/codex-rs/exec-server/tests/exec_process/windows_sandbox.rs b/codex-rs/exec-server/tests/exec_process/windows_sandbox.rs new file mode 100644 index 0000000000000000000000000000000000000000..f7e725ff9ec4030fb390c8f68d3c6d535bbad2bd --- /dev/null +++ b/codex-rs/exec-server/tests/exec_process/windows_sandbox.rs @@ -0,0 +1,109 @@ +//! Shared Windows sandbox behavior over the real exec-server RPC connection. + +use super::*; +use codex_protocol::models::PermissionProfile; +use codex_protocol::permissions::FileSystemAccessMode; +use codex_protocol::permissions::FileSystemPath; +use codex_protocol::permissions::FileSystemSandboxEntry; +use codex_protocol::permissions::FileSystemSandboxPolicy; +use codex_protocol::permissions::FileSystemSpecialPath; +use codex_protocol::permissions::NetworkSandboxPolicy; +use pretty_assertions::assert_eq; + +#[cfg_attr(not(windows), ignore = "requires a native Windows sandbox")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial_test::serial(remote_exec_server)] +async fn mxc_tmpdir_uses_command_environment_over_rpc() -> Result<()> { + crate::skip_if_mxc_unavailable!(Ok(())); + let root = TempDir::new()?; + let command_temp = root.path().join("command temp"); + let server_temp = root.path().join("server temp"); + std::fs::create_dir(&command_temp)?; + std::fs::create_dir(&server_temp)?; + let outside = server_temp.join("outside.txt"); + std::fs::write(&outside, "original")?; + let server = common::exec_server::exec_server_with_env( + [ + ("TEMP", server_temp.as_os_str()), + ("TMP", server_temp.as_os_str()), + ], + &[], + ) + .await?; + let environment = Environment::create_for_tests(Some(server.websocket_url().to_owned()))?; + let cwd = PathUri::from_host_native_path(root.path())?; + let fs = FileSystemSandboxPolicy::restricted(vec![ + FileSystemSandboxEntry::new( + FileSystemPath::Special { + value: FileSystemSpecialPath::Root, + }, + FileSystemAccessMode::Read, + ), + FileSystemSandboxEntry::new( + FileSystemPath::Special { + value: FileSystemSpecialPath::Tmpdir, + }, + FileSystemAccessMode::Write, + ), + ]); + let mut sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&fs, NetworkSandboxPolicy::Restricted), + cwd.clone(), + ); + sandbox.windows_sandbox_selection = codex_exec_server::WindowsSandboxSelection::Mxc; + let command_temp = command_temp.to_string_lossy().into_owned(); + let started = environment + .get_exec_backend() + .start(ExecParams { + process_id: ProcessId::from("windows-sandbox-temp"), + metadata: None, + argv: vec![ + r"C:\Windows\System32\cmd.exe".to_owned(), + "/D".to_owned(), + "/S".to_owned(), + "/C".to_owned(), + format!( + "echo allowed>\"%TEMP%\\allowed.txt\" & 2>\"%TEMP%\\denied.txt\" echo modified>\"{}\" & exit /b 0", + outside.display() + ), + ], + cwd, + env_policy: None, + shell_snapshot: None, + env: HashMap::from([ + ("SystemRoot".to_owned(), std::env::var("SystemRoot")?), + ("TEMP".to_owned(), command_temp.clone()), + ("TMP".to_owned(), command_temp), + ]), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: Some(sandbox), + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await?; + assert_eq!( + started.sandbox_type, + Some(codex_sandboxing::SandboxType::WindowsMxc) + ); + assert_eq!( + collect_process_output_from_events_with_timeout( + started.process, + Duration::from_secs(/*secs*/ 30), + ) + .await?, + (String::new(), String::new(), Some(0), true) + ); + assert_eq!( + ( + std::fs::read_to_string(root.path().join("command temp").join("allowed.txt"))? + .trim_end() + .to_owned(), + std::fs::read_to_string(outside)? + ), + ("allowed".to_owned(), "original".to_owned()) + ); + Ok(()) +} diff --git a/codex-rs/exec-server/tests/file_stream.rs b/codex-rs/exec-server/tests/file_stream.rs new file mode 100644 index 0000000000000000000000000000000000000000..9fd3201eebf17ddcfeb911a34c680b1bac3d21d5 --- /dev/null +++ b/codex-rs/exec-server/tests/file_stream.rs @@ -0,0 +1,384 @@ +mod common; + +use anyhow::Result; +use codex_exec_server::Environment; +use codex_exec_server::ExecServerClient; +use codex_exec_server::ExecServerError; +use codex_exec_server::ExecutorFileSystem; +use codex_exec_server::FsCloseParams; +use codex_exec_server::FsOpenParams; +use codex_exec_server::FsReadBlockParams; +use codex_exec_server::FsReadBlockResponse; +use codex_exec_server::ReadFileOptions; +use codex_exec_server::RemoteExecServerConnectArgs; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_utils_path_uri::PathUri; +use futures::TryStreamExt; +use pretty_assertions::assert_eq; +use std::sync::Arc; +#[cfg(any(unix, windows))] +use std::time::Duration; +use tempfile::TempDir; +#[cfg(windows)] +use tokio::net::windows::named_pipe::ServerOptions; +#[cfg(any(unix, windows))] +use tokio::time::timeout; +use uuid::Uuid; + +use crate::common::exec_server::exec_server; + +const BLOCK_SIZE: usize = 1024 * 1024; +const OPEN_FILE_LIMIT: usize = 128; + +#[tokio::test] +async fn stream_stops_after_an_exact_block_boundary() -> Result<()> { + let server = exec_server().await?; + let file_system = connect_file_system(server.websocket_url())?; + let tmp = TempDir::new()?; + let path = tmp.path().join("exact-blocks.bin"); + std::fs::write(&path, vec![b'x'; BLOCK_SIZE * 2])?; + + let chunks = file_system + .read_file_stream( + &PathUri::from_host_native_path(path)?, + /*sandbox*/ None, + ) + .await? + .try_collect::>() + .await?; + + assert_eq!( + chunks.iter().map(bytes::Bytes::len).collect::>(), + vec![BLOCK_SIZE, BLOCK_SIZE] + ); + Ok(()) +} + +#[tokio::test] +async fn completed_streams_release_handle_capacity() -> Result<()> { + let server = exec_server().await?; + let file_system = connect_file_system(server.websocket_url())?; + let tmp = TempDir::new()?; + let path = tmp.path().join("repeated.txt"); + std::fs::write(&path, b"repeated")?; + let path = PathUri::from_host_native_path(path)?; + + for _ in 0..=OPEN_FILE_LIMIT { + let chunks = file_system + .read_file_stream(&path, /*sandbox*/ None) + .await? + .try_collect::>() + .await?; + assert_eq!(chunks, vec![bytes::Bytes::from_static(b"repeated")]); + } + + Ok(()) +} + +#[cfg(unix)] +#[tokio::test] +async fn file_reads_reject_fifo_without_waiting_for_a_writer() -> Result<()> { + let server = exec_server().await?; + let file_system = connect_file_system(server.websocket_url())?; + let tmp = TempDir::new()?; + let path = tmp.path().join("named-pipe"); + let output = std::process::Command::new("mkfifo").arg(&path).output()?; + if !output.status.success() { + anyhow::bail!( + "mkfifo failed: stdout={} stderr={}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + } + + let path_uri = PathUri::from_host_native_path(&path)?; + let read_error = timeout( + Duration::from_secs(1), + file_system.read_file(&path_uri, ReadFileOptions::default(), /*sandbox*/ None), + ) + .await + .expect("reading a FIFO should not wait for a writer") + .expect_err("reading a FIFO should be rejected"); + let stream_result = timeout( + Duration::from_secs(1), + file_system.read_file_stream(&path_uri, /*sandbox*/ None), + ) + .await + .expect("streaming a FIFO should not wait for a writer"); + let Err(stream_error) = stream_result else { + panic!("streaming a FIFO should be rejected"); + }; + let expected = format!("path `{}` is not a file", path.display()); + assert_eq!( + (read_error.to_string(), stream_error.to_string()), + (expected.clone(), expected) + ); + Ok(()) +} + +#[cfg(windows)] +#[tokio::test] +async fn file_reads_reject_named_pipes() -> Result<()> { + let server = exec_server().await?; + let file_system = connect_file_system(server.websocket_url())?; + + let read_path = format!(r"\\.\pipe\codex-fs-read-{}", Uuid::new_v4()); + let _read_pipe = ServerOptions::new() + .first_pipe_instance(true) + .create(&read_path)?; + let read_error = timeout( + Duration::from_secs(1), + file_system.read_file( + &PathUri::from_host_native_path(std::path::Path::new(&read_path))?, + ReadFileOptions::default(), + /*sandbox*/ None, + ), + ) + .await + .expect("reading a named pipe should not hang") + .expect_err("reading a named pipe should be rejected"); + + let stream_path = format!(r"\\.\pipe\codex-fs-stream-{}", Uuid::new_v4()); + let _stream_pipe = ServerOptions::new() + .first_pipe_instance(true) + .create(&stream_path)?; + let stream_result = timeout( + Duration::from_secs(1), + file_system.read_file_stream( + &PathUri::from_host_native_path(std::path::Path::new(&stream_path))?, + /*sandbox*/ None, + ), + ) + .await + .expect("streaming a named pipe should not hang"); + let Err(stream_error) = stream_result else { + panic!("streaming a named pipe should be rejected"); + }; + + assert_eq!( + (read_error.kind(), stream_error.kind()), + ( + std::io::ErrorKind::InvalidInput, + std::io::ErrorKind::InvalidInput, + ) + ); + Ok(()) +} + +#[cfg(unix)] +#[tokio::test] +async fn stream_keeps_reading_the_open_file_after_path_replacement() -> Result<()> { + let server = exec_server().await?; + let file_system = connect_file_system(server.websocket_url())?; + let tmp = TempDir::new()?; + let path = tmp.path().join("replaceable.bin"); + std::fs::write(&path, vec![b'a'; BLOCK_SIZE + 1])?; + let sandbox = read_only_sandbox(tmp.path().to_path_buf()); + let mut stream = file_system + .read_file_stream(&PathUri::from_host_native_path(&path)?, Some(&sandbox)) + .await?; + + assert_eq!( + stream.try_next().await?, + Some(bytes::Bytes::from(vec![b'a'; BLOCK_SIZE])) + ); + let replacement = tmp.path().join("replacement.bin"); + std::fs::write(&replacement, vec![b'b'; BLOCK_SIZE + 1])?; + std::fs::remove_file(&path)?; + std::fs::rename(replacement, &path)?; + + assert_eq!( + stream.try_next().await?, + Some(bytes::Bytes::from_static(b"a")) + ); + assert_eq!(stream.try_next().await?, None); + Ok(()) +} + +#[tokio::test] +async fn read_block_supports_non_sequential_offsets_and_lengths() -> Result<()> { + let mut server = exec_server().await?; + let client = ExecServerClient::connect_websocket(RemoteExecServerConnectArgs::new( + server.websocket_url().to_string(), + "file-stream-protocol-test".to_string(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )) + .await?; + let tmp = TempDir::new()?; + let path = tmp.path().join("non-sequential.bin"); + std::fs::write(&path, b"0123456789")?; + let open = client + .fs_open(FsOpenParams { + handle_id: Uuid::new_v4().simple().to_string(), + path: PathUri::from_host_native_path(path)?, + sandbox: None, + }) + .await?; + + let mut blocks = Vec::new(); + for (offset, len) in [(6, 3), (1, 2), (8, 4), (0, 2)] { + blocks.push( + client + .fs_read_block(FsReadBlockParams { + handle_id: open.handle_id.clone(), + offset, + len, + }) + .await?, + ); + } + assert_eq!( + blocks, + vec![ + FsReadBlockResponse { + chunk: b"678".to_vec().into(), + eof: false, + }, + FsReadBlockResponse { + chunk: b"12".to_vec().into(), + eof: false, + }, + FsReadBlockResponse { + chunk: b"89".to_vec().into(), + eof: true, + }, + FsReadBlockResponse { + chunk: b"01".to_vec().into(), + eof: false, + }, + ] + ); + client + .fs_close(FsCloseParams { + handle_id: open.handle_id, + }) + .await?; + drop(client); + server.shutdown().await?; + Ok(()) +} + +#[tokio::test] +async fn open_enforces_the_per_connection_limit_and_close_releases_capacity() -> Result<()> { + let mut server = exec_server().await?; + let client = ExecServerClient::connect_websocket(RemoteExecServerConnectArgs::new( + server.websocket_url().to_string(), + "file-stream-protocol-test".to_string(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )) + .await?; + let tmp = TempDir::new()?; + let path = tmp.path().join("limited.bin"); + std::fs::write(&path, b"limited")?; + let path = PathUri::from_host_native_path(path)?; + let mut handles = Vec::with_capacity(OPEN_FILE_LIMIT); + for _ in 0..OPEN_FILE_LIMIT { + let open = client + .fs_open(FsOpenParams { + handle_id: Uuid::new_v4().simple().to_string(), + path: path.clone(), + sandbox: None, + }) + .await?; + handles.push(open.handle_id); + } + + let error = client + .fs_open(FsOpenParams { + handle_id: Uuid::new_v4().simple().to_string(), + path: path.clone(), + sandbox: None, + }) + .await + .expect_err("opening beyond the limit should fail"); + let ExecServerError::Server { code, message } = error else { + anyhow::bail!("expected server error, got {error:?}"); + }; + assert_eq!( + (code, message), + ( + -32600, + format!("at most {OPEN_FILE_LIMIT} file reads may be open per connection"), + ) + ); + + client + .fs_close(FsCloseParams { + handle_id: handles.remove(0), + }) + .await?; + client + .fs_open(FsOpenParams { + handle_id: Uuid::new_v4().simple().to_string(), + path, + sandbox: None, + }) + .await?; + drop(client); + server.shutdown().await?; + Ok(()) +} + +#[tokio::test] +async fn open_rejects_handle_ids_longer_than_32_bytes() -> Result<()> { + let server = exec_server().await?; + let client = ExecServerClient::connect_websocket(RemoteExecServerConnectArgs::new( + server.websocket_url().to_string(), + "file-stream-protocol-test".to_string(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )) + .await?; + let tmp = TempDir::new()?; + let path = tmp.path().join("handle-id-limit.bin"); + std::fs::write(&path, b"limited")?; + + let error = client + .fs_open(FsOpenParams { + handle_id: "x".repeat(33), + path: PathUri::from_host_native_path(path)?, + sandbox: None, + }) + .await + .expect_err("oversized handle ID should fail"); + + let ExecServerError::Server { code, message } = error else { + anyhow::bail!("expected server error, got {error:?}"); + }; + assert_eq!( + (code, message), + ( + -32600, + "file read handle ID must not exceed 32 bytes".to_string(), + ) + ); + Ok(()) +} + +fn connect_file_system(websocket_url: &str) -> Result> { + let environment = Environment::create_for_tests(Some(websocket_url.to_string()))?; + Ok(environment.get_filesystem()) +} + +// Only the Unix stream tests above need this sandbox builder. +#[cfg(unix)] +fn read_only_sandbox(path: std::path::PathBuf) -> codex_exec_server::FileSystemSandboxContext { + use codex_exec_server::FileSystemSandboxContext; + use codex_protocol::models::PermissionProfile; + use codex_protocol::permissions::FileSystemAccessMode; + use codex_protocol::permissions::FileSystemSandboxEntry; + use codex_protocol::permissions::FileSystemSandboxPolicy; + use codex_protocol::permissions::NetworkSandboxPolicy; + use codex_utils_absolute_path::AbsolutePathBuf; + + let path = AbsolutePathBuf::from_absolute_path(&path) + .unwrap_or_else(|err| panic!("sandbox path should be absolute: {err}")); + FileSystemSandboxContext::from_permission_profile(PermissionProfile::from_runtime_permissions( + &FileSystemSandboxPolicy::restricted(vec![FileSystemSandboxEntry { + path: path.into(), + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }]), + NetworkSandboxPolicy::Restricted, + )) +} diff --git a/codex-rs/exec-server/tests/file_system/shared.rs b/codex-rs/exec-server/tests/file_system/shared.rs new file mode 100644 index 0000000000000000000000000000000000000000..96dfa143492f1ca2e6413711985343a764dea72a --- /dev/null +++ b/codex-rs/exec-server/tests/file_system/shared.rs @@ -0,0 +1,1142 @@ +use anyhow::Context; +use anyhow::Result; +use codex_exec_server::CopyOptions; +use codex_exec_server::CreateDirectoryOptions; +#[cfg(unix)] +use codex_exec_server::ExecServerRuntimePaths; +#[cfg(unix)] +use codex_exec_server::ExecutorFileSystem; +use codex_exec_server::FILE_READ_CHUNK_SIZE; +use codex_exec_server::FileMetadata; +#[cfg(unix)] +use codex_exec_server::LocalFileSystem; +use codex_exec_server::ReadDirectoryEntry; +use codex_exec_server::RemoveOptions; +use codex_exec_server::WalkEntry; +use codex_exec_server::WalkEntryKind; +use codex_exec_server::WalkOptions; +use codex_exec_server::WalkOutcome; +use codex_exec_server::WriteFileOptions; +use codex_file_system::MAX_WALK_DEPTH; +use codex_file_system::MAX_WALK_DIRECTORIES; +use codex_file_system::MAX_WALK_ENTRIES; +use codex_protocol::models::AdditionalPermissionProfile; +use codex_protocol::models::FileSystemPermissions; +use codex_protocol::models::PermissionProfile; +use codex_sandboxing::policy_transforms::effective_file_system_sandbox_policy; +use codex_sandboxing::policy_transforms::effective_network_sandbox_policy; +use codex_utils_path_uri::PathUri; +use futures::TryStreamExt; +use pretty_assertions::assert_eq; +use std::path::Path; +use tempfile::TempDir; +use test_case::test_case; + +use super::support::FileSystemImplementation; +use super::support::absolute_path; +use super::support::create_file_system_context; +#[cfg(windows)] +use super::support::is_unsupported_restricted_token_host; +use super::support::read_only_sandbox; +use super::support::workspace_write_sandbox; + +#[test] +fn sandbox_context_from_profile_preserves_workspace_write_read_only_subpaths() -> Result<()> { + let tmp = TempDir::new()?; + let writable_dir = tmp.path().join("writable"); + let git_dir = writable_dir.join(".git"); + std::fs::create_dir_all(&git_dir)?; + + let sandbox = workspace_write_sandbox(writable_dir.clone()); + let permissions: PermissionProfile = sandbox.permissions.try_into()?; + let policy = permissions.file_system_sandbox_policy(); + let cwd = absolute_path(writable_dir.clone()); + let writable_roots = policy.get_writable_roots_with_cwd(cwd.as_path()); + let writable_dir = absolute_path(std::fs::canonicalize(writable_dir)?); + let git_dir = absolute_path(std::fs::canonicalize(git_dir)?); + let Some(writable_root) = writable_roots + .iter() + .find(|writable_root| writable_root.root == writable_dir) + else { + panic!("writable root should be preserved"); + }; + + assert!(writable_root.read_only_subpaths.contains(&git_dir)); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_get_metadata_reports_files_and_directories( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let file_path = tmp.path().join("note.txt"); + let directory_path = tmp.path().join("notes"); + std::fs::write(&file_path, "hello")?; + std::fs::create_dir(&directory_path)?; + + let file_metadata = file_system + .get_metadata( + &PathUri::from_host_native_path(&file_path)?, + Default::default(), + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + file_metadata, + FileMetadata { + is_directory: false, + is_file: true, + is_symlink: false, + size: 5, + created_at_ms: file_metadata.created_at_ms, + modified_at_ms: file_metadata.modified_at_ms, + } + ); + assert!(file_metadata.modified_at_ms > 0); + + let directory_metadata = file_system + .get_metadata( + &PathUri::from_host_native_path(&directory_path)?, + Default::default(), + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + directory_metadata, + FileMetadata { + is_directory: true, + is_file: false, + is_symlink: false, + size: std::fs::metadata(&directory_path)?.len(), + created_at_ms: directory_metadata.created_at_ms, + modified_at_ms: directory_metadata.modified_at_ms, + } + ); + assert!(directory_metadata.modified_at_ms > 0); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local, true, false ; "local_follow")] +#[test_case(FileSystemImplementation::Local, false, false ; "local_no_follow")] +#[test_case(FileSystemImplementation::Remote, true, false ; "remote_follow")] +#[test_case(FileSystemImplementation::Remote, false, false ; "remote_no_follow")] +#[cfg_attr(any(target_os = "linux", windows), test_case(FileSystemImplementation::Local, false, true ; "local_sandboxed_no_follow"))] +#[cfg_attr(any(target_os = "linux", windows), test_case(FileSystemImplementation::Remote, false, true ; "remote_sandboxed_no_follow"))] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_create_directory_creates_nested_directories( + implementation: FileSystemImplementation, + follow_symlinks: bool, + sandboxed: bool, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let root = tmp.path().canonicalize()?; + let nested_dir = root.join("source").join("nested"); + let sandbox = sandboxed.then(|| workspace_write_sandbox(root)); + + let result = file_system + .create_directory( + &PathUri::from_host_native_path(&nested_dir)?, + CreateDirectoryOptions { + recursive: true, + follow_symlinks, + }, + sandbox.as_ref(), + ) + .await; + #[cfg(windows)] + if is_unsupported_restricted_token_host(&result) { + return Ok(()); + } + result.with_context(|| format!("mode={implementation}, sandboxed={sandboxed}"))?; + assert!(nested_dir.is_dir()); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local, true, false ; "local_follow")] +#[test_case(FileSystemImplementation::Local, false, false ; "local_no_follow")] +#[test_case(FileSystemImplementation::Remote, true, false ; "remote_follow")] +#[test_case(FileSystemImplementation::Remote, false, false ; "remote_no_follow")] +#[cfg_attr(any(target_os = "linux", windows), test_case(FileSystemImplementation::Local, false, true ; "local_sandboxed_no_follow"))] +#[cfg_attr(any(target_os = "linux", windows), test_case(FileSystemImplementation::Remote, false, true ; "remote_sandboxed_no_follow"))] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_write_file_writes_bytes( + implementation: FileSystemImplementation, + follow_symlinks: bool, + sandboxed: bool, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let root = tmp.path().canonicalize()?; + let file_path = root.join("note.txt"); + let sandbox = sandboxed.then(|| workspace_write_sandbox(root.clone())); + let result = file_system + .write_file( + &PathUri::from_host_native_path(&file_path)?, + b"hello from trait".to_vec(), + WriteFileOptions { follow_symlinks }, + sandbox.as_ref(), + ) + .await; + #[cfg(windows)] + if is_unsupported_restricted_token_host(&result) { + return Ok(()); + } + result.with_context(|| format!("mode={implementation}, sandboxed={sandboxed}"))?; + assert_eq!(std::fs::read(file_path)?, b"hello from trait"); + + let file_path = root.join("existing.txt"); + std::fs::write(&file_path, b"before")?; + file_system + .write_file( + &PathUri::from_host_native_path(&file_path)?, + b"after".to_vec(), + WriteFileOptions { follow_symlinks }, + sandbox.as_ref(), + ) + .await + .with_context(|| format!("mode={implementation}, sandboxed={sandboxed}"))?; + assert_eq!(std::fs::read(file_path)?, b"after"); + + Ok(()) +} + +#[test] +fn path_uri_join_and_parent_preserve_lexical_paths() -> Result<()> { + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + let source_dir_uri = PathUri::from_host_native_path(&source_dir)?; + let joined_nested = source_dir_uri.join("nested/note.txt")?; + assert_eq!( + joined_nested, + PathUri::from_host_native_path(source_dir.join("nested").join("note.txt"))? + ); + let joined_parent = joined_nested.parent(); + assert_eq!( + joined_parent, + Some(PathUri::from_host_native_path(source_dir.join("nested"))?) + ); + let joined_parent_traversal = source_dir_uri.join("../outside")?; + assert_eq!( + joined_parent_traversal, + PathUri::from_host_native_path(source_dir.join("../outside"))? + ); + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_read_file_returns_bytes( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let file_path = tmp.path().join("note.txt"); + std::fs::write(&file_path, "hello from trait")?; + + let contents = file_system + .read_file( + &PathUri::from_host_native_path(&file_path)?, + Default::default(), + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!(contents, b"hello from trait"); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_read_file_stream_returns_bounded_chunks( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let file_path = tmp.path().join("blocks.bin"); + let contents = (0..FILE_READ_CHUNK_SIZE * 2 + 17) + .map(|index| (index % 251) as u8) + .collect::>(); + std::fs::write(&file_path, &contents)?; + + let path = PathUri::from_host_native_path(file_path)?; + let sandbox = read_only_sandbox(tmp.path().to_path_buf()); + for sandbox in [None, Some(&sandbox)] { + let chunks = file_system + .read_file_stream(&path, sandbox) + .await + .with_context(|| format!("mode={implementation}"))? + .try_collect::>() + .await?; + + assert!( + chunks + .iter() + .all(|chunk| !chunk.is_empty() && chunk.len() <= FILE_READ_CHUNK_SIZE) + ); + assert_eq!( + chunks + .iter() + .flat_map(|chunk| chunk.iter().copied()) + .collect::>(), + contents + ); + } + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_read_file_text_returns_string( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let file_path = tmp.path().join("note.txt"); + std::fs::write(&file_path, "hello from trait")?; + + let contents = file_system + .read_file_text( + &PathUri::from_host_native_path(&file_path)?, + Default::default(), + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!(contents, "hello from trait"); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_copy_copies_file(implementation: FileSystemImplementation) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_file = tmp.path().join("source.txt"); + let copied_file = tmp.path().join("copy.txt"); + std::fs::write(&source_file, "hello from trait")?; + + file_system + .copy( + &PathUri::from_host_native_path(&source_file)?, + &PathUri::from_host_native_path(&copied_file)?, + CopyOptions { recursive: false }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!(std::fs::read_to_string(copied_file)?, "hello from trait"); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_copy_copies_directory_recursively( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + let nested_dir = source_dir.join("nested"); + let nested_file = nested_dir.join("note.txt"); + let copied_dir = tmp.path().join("copied"); + std::fs::create_dir_all(&nested_dir)?; + std::fs::write(&nested_file, "hello from trait")?; + + file_system + .copy( + &PathUri::from_host_native_path(&source_dir)?, + &PathUri::from_host_native_path(&copied_dir)?, + CopyOptions { recursive: true }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + std::fs::read_to_string(copied_dir.join("nested").join("note.txt"))?, + "hello from trait" + ); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_read_directory_lists_entries( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + std::fs::create_dir_all(source_dir.join("nested"))?; + std::fs::write(source_dir.join("root.txt"), "hello")?; + + let mut entries = file_system + .read_directory( + &PathUri::from_host_native_path(&source_dir)?, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + entries.sort_by(|left, right| left.file_name.cmp(&right.file_name)); + assert_eq!( + entries, + vec![ + ReadDirectoryEntry { + file_name: "nested".to_string(), + is_directory: true, + is_file: false, + }, + ReadDirectoryEntry { + file_name: "root.txt".to_string(), + is_directory: false, + is_file: true, + }, + ] + ); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_walk_returns_a_bounded_tree( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + let nested_dir = source_dir.join("nested"); + std::fs::create_dir_all(&nested_dir)?; + std::fs::write(source_dir.join("root.txt"), "root")?; + std::fs::write(nested_dir.join("note.txt"), "nested")?; + + let source_uri = PathUri::from_host_native_path(&source_dir)?; + let outcome = file_system + .walk( + &source_uri, + WalkOptions { + max_depth: 4, + max_directories: 10, + max_entries: 10, + follow_directory_symlinks: false, + prune_hidden_directories: false, + }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + outcome, + WalkOutcome { + entries: vec![ + WalkEntry { + path: PathUri::from_host_native_path(&nested_dir)?, + kind: WalkEntryKind::Directory, + }, + WalkEntry { + path: PathUri::from_host_native_path(source_dir.join("root.txt"))?, + kind: WalkEntryKind::File, + }, + WalkEntry { + path: PathUri::from_host_native_path(nested_dir.join("note.txt"))?, + kind: WalkEntryKind::File, + }, + ], + errors: Vec::new(), + truncated: false, + } + ); + + let root_entries = vec![ + WalkEntry { + path: PathUri::from_host_native_path(&nested_dir)?, + kind: WalkEntryKind::Directory, + }, + WalkEntry { + path: PathUri::from_host_native_path(source_dir.join("root.txt"))?, + kind: WalkEntryKind::File, + }, + ]; + let shallow = file_system + .walk( + &source_uri, + WalkOptions { + max_depth: 0, + max_directories: 10, + max_entries: 10, + follow_directory_symlinks: false, + prune_hidden_directories: false, + }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + shallow, + WalkOutcome { + entries: root_entries.clone(), + errors: Vec::new(), + truncated: false, + } + ); + + let directory_bounded = file_system + .walk( + &source_uri, + WalkOptions { + max_depth: 4, + max_directories: 1, + max_entries: 10, + follow_directory_symlinks: false, + prune_hidden_directories: false, + }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + directory_bounded, + WalkOutcome { + entries: root_entries, + errors: Vec::new(), + truncated: true, + } + ); + + let bounded = file_system + .walk( + &source_uri, + WalkOptions { + max_depth: 4, + max_directories: 10, + max_entries: 1, + follow_directory_symlinks: false, + prune_hidden_directories: false, + }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + bounded, + WalkOutcome { + entries: vec![WalkEntry { + path: PathUri::from_host_native_path(&nested_dir)?, + kind: WalkEntryKind::Directory, + }], + errors: Vec::new(), + truncated: true, + } + ); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_walk_handles_invalid_roots_and_limits( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + let tmp = TempDir::new()?; + let file_path = tmp.path().join("file.txt"); + std::fs::write(&file_path, "contents")?; + let missing = PathUri::from_host_native_path(tmp.path().join("missing"))?; + let options = WalkOptions { + max_depth: 8, + max_directories: 100, + max_entries: 100, + follow_directory_symlinks: false, + prune_hidden_directories: false, + }; + + let outcome = file_system + .walk( + &PathUri::from_host_native_path(file_path)?, + options, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!(outcome, WalkOutcome::default()); + + let error = file_system + .walk(&missing, options, /*sandbox*/ None) + .await + .expect_err("a missing root must fail"); + assert_eq!(error.kind(), std::io::ErrorKind::NotFound); + + let zero_limit_message = "filesystem walk limits must be greater than zero"; + let excessive_limit_message = format!( + "filesystem walk limits exceed maximums: depth={MAX_WALK_DEPTH}, directories={MAX_WALK_DIRECTORIES}, entries={MAX_WALK_ENTRIES}" + ); + for (max_depth, max_directories, max_entries, message) in [ + (8, 0, 100, zero_limit_message), + (8, 100, 0, zero_limit_message), + ( + MAX_WALK_DEPTH + 1, + 100, + 100, + excessive_limit_message.as_str(), + ), + ( + 8, + MAX_WALK_DIRECTORIES + 1, + 100, + excessive_limit_message.as_str(), + ), + ( + 8, + 100, + MAX_WALK_ENTRIES + 1, + excessive_limit_message.as_str(), + ), + ] { + let options = WalkOptions { + max_depth, + max_directories, + max_entries, + ..options + }; + // Invalid limits take precedence over the missing root. + let error = file_system + .walk(&missing, options, /*sandbox*/ None) + .await + .expect_err("invalid walk limits must fail"); + assert_eq!( + (error.kind(), error.to_string()), + (std::io::ErrorKind::InvalidInput, message.to_owned()), + "mode={implementation}, options={options:?}", + ); + } + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_walk_honors_read_sandbox( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + let file_path = source_dir.join("note.txt"); + std::fs::create_dir_all(&source_dir)?; + std::fs::write(&file_path, "sandboxed")?; + let sandbox = read_only_sandbox(source_dir.clone()); + + let outcome = file_system + .walk( + &PathUri::from_host_native_path(&source_dir)?, + WalkOptions { + max_depth: 1, + max_directories: 2, + max_entries: 2, + follow_directory_symlinks: false, + prune_hidden_directories: false, + }, + Some(&sandbox), + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + outcome, + WalkOutcome { + entries: vec![WalkEntry { + path: PathUri::from_host_native_path(file_path)?, + kind: WalkEntryKind::File, + }], + errors: Vec::new(), + truncated: false, + } + ); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_remove_removes_directory( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let directory_path = tmp.path().join("remove-me"); + std::fs::create_dir_all(directory_path.join("nested"))?; + + file_system + .remove( + &PathUri::from_host_native_path(&directory_path)?, + RemoveOptions { + recursive: true, + force: true, + follow_symlinks: true, + }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert!(!directory_path.exists()); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local, false ; "local")] +#[test_case(FileSystemImplementation::Remote, false ; "remote")] +#[cfg_attr(any(target_os = "linux", windows), test_case(FileSystemImplementation::Local, true ; "local_sandboxed"))] +#[cfg_attr(any(target_os = "linux", windows), test_case(FileSystemImplementation::Remote, true ; "remote_sandboxed"))] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_remove_no_follow_removes_file_and_empty_directory( + implementation: FileSystemImplementation, + sandboxed: bool, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + let tmp = TempDir::new()?; + let root = tmp.path().canonicalize()?; + let sandbox = sandboxed.then(|| workspace_write_sandbox(root.clone())); + let options = RemoveOptions { + recursive: false, + force: false, + follow_symlinks: false, + }; + + let file_path = root.join("remove-me.txt"); + std::fs::write(&file_path, b"remove")?; + let result = file_system + .remove( + &PathUri::from_host_native_path(&file_path)?, + options, + sandbox.as_ref(), + ) + .await; + #[cfg(windows)] + if is_unsupported_restricted_token_host(&result) { + return Ok(()); + } + result.with_context(|| format!("mode={implementation}, sandboxed={sandboxed}"))?; + assert!(!file_path.exists()); + + let directory_path = root.join("remove-me"); + std::fs::create_dir(&directory_path)?; + file_system + .remove( + &PathUri::from_host_native_path(&directory_path)?, + options, + sandbox.as_ref(), + ) + .await + .with_context(|| format!("mode={implementation}, sandboxed={sandboxed}"))?; + assert!(!directory_path.exists()); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_write_file_reports_missing_parent( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let missing_parent_path = tmp.path().join("missing").join("note.txt"); + + let error = match file_system + .write_file( + &PathUri::from_host_native_path(&missing_parent_path)?, + b"hello from trait".to_vec(), + Default::default(), + /*sandbox*/ None, + ) + .await + { + Ok(()) => anyhow::bail!("write should fail when parent directory is absent"), + Err(error) => error, + }; + assert_eq!( + error.kind(), + std::io::ErrorKind::NotFound, + "mode={implementation}" + ); + assert!(!missing_parent_path.exists(), "mode={implementation}"); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_copy_rejects_directory_without_recursive( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + std::fs::create_dir_all(&source_dir)?; + + let error = file_system + .copy( + &PathUri::from_host_native_path(&source_dir)?, + &PathUri::from_host_native_path(tmp.path().join("dest"))?, + CopyOptions { recursive: false }, + /*sandbox*/ None, + ) + .await; + let error = error.expect_err("copying a directory without recursion should fail"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput); + assert_eq!( + error.to_string(), + "fs/copy requires recursive: true when sourcePath is a directory" + ); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_sandboxed_metadata_and_read_allow_readable_root( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let file_path = allowed_dir.join("note.txt"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::write(&file_path, "sandboxed hello")?; + let sandbox = read_only_sandbox(allowed_dir); + + let metadata = file_system + .get_metadata( + &PathUri::from_host_native_path(&file_path)?, + Default::default(), + Some(&sandbox), + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + metadata, + FileMetadata { + is_directory: false, + is_file: true, + is_symlink: false, + size: 15, + created_at_ms: metadata.created_at_ms, + modified_at_ms: metadata.modified_at_ms, + } + ); + + let contents = file_system + .read_file( + &PathUri::from_host_native_path(&file_path)?, + Default::default(), + Some(&sandbox), + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!(contents, b"sandboxed hello"); + + let chunks = file_system + .read_file_stream(&PathUri::from_host_native_path(&file_path)?, Some(&sandbox)) + .await + .with_context(|| format!("stream mode={implementation}"))? + .try_collect::>() + .await?; + assert_eq!(chunks.concat(), b"sandboxed hello"); + + Ok(()) +} + +#[cfg(unix)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn sandboxed_file_operations_cannot_read_helper_siblings() -> Result<()> { + let helper_paths = crate::common::exec_server::test_codex_helper_paths()?; + let root = TempDir::new()?; + let runtime_dir = root.path().join("runtime"); + let workspace = root.path().join("workspace"); + std::fs::create_dir(&runtime_dir)?; + std::fs::create_dir(&workspace)?; + + let helper = runtime_dir.join("codex-test-helper"); + std::fs::hard_link(&helper_paths.codex_exe, &helper) + .or_else(|_| std::fs::copy(&helper_paths.codex_exe, &helper).map(|_| ()))?; + let linux_sandbox = if helper_paths.codex_linux_sandbox_exe.is_some() { + let alias = runtime_dir.join("codex-linux-sandbox"); + std::fs::hard_link(&helper, &alias) + .or_else(|_| std::fs::copy(&helper, &alias).map(|_| ()))?; + Some(alias) + } else { + None + }; + let file_system = + LocalFileSystem::with_runtime_paths(ExecServerRuntimePaths::new(helper, linux_sandbox)?); + + let sibling = runtime_dir.join("credentials.json"); + std::fs::write(&sibling, "secret")?; + let escaping_link = workspace.join("credentials-link"); + std::os::unix::fs::symlink(&sibling, &escaping_link)?; + let sandbox = workspace_write_sandbox(workspace.clone()); + let allowed_file = workspace.join("allowed.txt"); + std::fs::write(&allowed_file, b"allowed")?; + let allowed_contents = file_system + .read_file( + &PathUri::from_host_native_path(&allowed_file)?, + Default::default(), + Some(&sandbox), + ) + .await?; + assert_eq!(allowed_contents, b"allowed"); + + #[cfg(target_os = "macos")] + assert!( + file_system + .read_directory( + &PathUri::from_host_native_path("/Applications")?, + Some(&sandbox) + ) + .await + .is_err(), + "filesystem helpers should not inherit the normal process sandbox's /Applications access" + ); + + let sibling_uri = PathUri::from_host_native_path(&sibling)?; + let destination = PathUri::from_host_native_path(workspace.join("copied.json"))?; + + for path in [sibling, escaping_link] { + let path = PathUri::from_host_native_path(path)?; + assert!( + file_system + .read_file(&path, Default::default(), Some(&sandbox)) + .await + .is_err(), + "sandboxed read unexpectedly accessed helper sibling {path}" + ); + assert!( + file_system + .read_file_stream(&path, Some(&sandbox)) + .await + .is_err(), + "sandboxed streaming unexpectedly accessed helper sibling {path}" + ); + } + assert!( + file_system + .copy( + &sibling_uri, + &destination, + CopyOptions { recursive: false }, + Some(&sandbox), + ) + .await + .is_err(), + "sandboxed copy unexpectedly accessed helper sibling {sibling_uri}" + ); + + Ok(()) +} + +pub(crate) async fn assert_canonicalize_resolves_directory_alias( + implementation: FileSystemImplementation, + create_directory_alias: impl FnOnce(&Path, &Path) -> Result<()>, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + let nested_dir = source_dir.join("nested"); + let file_path = nested_dir.join("note.txt"); + let alias_dir = tmp.path().join("source-alias"); + std::fs::create_dir_all(&nested_dir)?; + std::fs::write(&file_path, "canonical hello")?; + create_directory_alias(&source_dir, &alias_dir)?; + + let requested_path = PathUri::from_host_native_path(alias_dir.join("nested").join("note.txt"))?; + let expected_path = PathUri::from_host_native_path(std::fs::canonicalize(&file_path)?)?; + assert_ne!(requested_path, expected_path); + + let canonical_path = file_system + .canonicalize(&requested_path, /*sandbox*/ None) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!(canonical_path, expected_path); + + Ok(()) +} + +pub(crate) async fn assert_sandboxed_canonicalize_resolves_directory_alias( + implementation: FileSystemImplementation, + create_directory_alias: impl FnOnce(&Path, &Path) -> Result<()>, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + let nested_dir = source_dir.join("nested"); + let file_path = nested_dir.join("note.txt"); + let alias_dir = tmp.path().join("source-alias"); + std::fs::create_dir_all(&nested_dir)?; + std::fs::write(&file_path, "sandboxed canonical hello")?; + create_directory_alias(&source_dir, &alias_dir)?; + let sandbox = read_only_sandbox(tmp.path().to_path_buf()); + + let requested_path = PathUri::from_host_native_path(alias_dir.join("nested").join("note.txt"))?; + let expected_path = PathUri::from_host_native_path(std::fs::canonicalize(&file_path)?)?; + assert_ne!(requested_path, expected_path); + + let canonical_path = file_system + .canonicalize(&requested_path, Some(&sandbox)) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!(canonical_path, expected_path); + + Ok(()) +} + +/// Verifies that effective additional permissions extend a read-only sandbox with a writable root. +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[cfg_attr( + windows, + ignore = "Windows restricted-token sandbox cannot enforce split writable roots" +)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_sandboxed_write_allows_additional_write_root( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let readable_dir = tmp.path().join("readable"); + let writable_dir = tmp.path().join("writable"); + let file_path = writable_dir.join("note.txt"); + std::fs::create_dir_all(&readable_dir)?; + std::fs::create_dir_all(&writable_dir)?; + + let mut sandbox = read_only_sandbox(readable_dir); + let additional_permissions = AdditionalPermissionProfile { + network: None, + file_system: Some(FileSystemPermissions::from_read_write_roots( + /*read*/ None, + Some(vec![absolute_path(writable_dir)]), + )), + }; + let native_permissions: PermissionProfile = sandbox.permissions.clone().try_into()?; + let file_system_policy = effective_file_system_sandbox_policy( + &native_permissions.file_system_sandbox_policy(), + Some(&additional_permissions), + ); + let network_policy = effective_network_sandbox_policy( + native_permissions.network_sandbox_policy(), + Some(&additional_permissions), + ); + sandbox.permissions = PermissionProfile::from_runtime_permissions_with_enforcement( + native_permissions.enforcement(), + &file_system_policy, + network_policy, + ) + .into(); + + file_system + .write_file( + &PathUri::from_host_native_path(&file_path)?, + b"created".to_vec(), + Default::default(), + Some(&sandbox), + ) + .await + .with_context(|| format!("write file through additional root mode={implementation}"))?; + assert_eq!(std::fs::read(&file_path)?, b"created"); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_copy_rejects_copying_directory_into_descendant( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + std::fs::create_dir_all(source_dir.join("nested"))?; + + let error = file_system + .copy( + &PathUri::from_host_native_path(&source_dir)?, + &PathUri::from_host_native_path(source_dir.join("nested").join("copy"))?, + CopyOptions { recursive: true }, + /*sandbox*/ None, + ) + .await; + let error = error.expect_err("copying a directory into itself should fail"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput); + assert_eq!( + error.to_string(), + "fs/copy cannot copy a directory to itself or one of its descendants" + ); + + Ok(()) +} diff --git a/codex-rs/exec-server/tests/file_system/support.rs b/codex-rs/exec-server/tests/file_system/support.rs new file mode 100644 index 0000000000000000000000000000000000000000..34125c24aeed629c8c90db5e16184438be9eb6f6 --- /dev/null +++ b/codex-rs/exec-server/tests/file_system/support.rs @@ -0,0 +1,169 @@ +use std::fmt; +use std::sync::Arc; + +use anyhow::Result; +use codex_exec_server::Environment; +use codex_exec_server::ExecServerRuntimePaths; +use codex_exec_server::ExecutorFileSystem; +use codex_exec_server::FileSystemSandboxContext; +use codex_exec_server::LocalFileSystem; +use codex_exec_server::WindowsSandboxSelection; +use codex_protocol::models::PermissionProfile; +use codex_protocol::permissions::FileSystemAccessMode; +use codex_protocol::permissions::FileSystemPath; +use codex_protocol::permissions::FileSystemSandboxEntry; +use codex_protocol::permissions::FileSystemSandboxPolicy; +use codex_protocol::permissions::FileSystemSpecialPath; +use codex_protocol::permissions::NetworkSandboxPolicy; +use codex_utils_absolute_path::AbsolutePathBuf; +#[cfg(windows)] +use codex_utils_path_uri::PathUri; + +use crate::common::exec_server::ExecServerHarness; +use crate::common::exec_server::TestCodexHelperPaths; +use crate::common::exec_server::exec_server; +use crate::common::exec_server::test_codex_helper_paths; + +pub(crate) struct FileSystemContext { + pub(crate) file_system: Arc, + _helper_paths: Option, + _server: Option, +} + +#[derive(Clone, Copy, Debug)] +pub(crate) enum FileSystemImplementation { + Local, + Remote, +} + +impl fmt::Display for FileSystemImplementation { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Local => formatter.write_str("local"), + Self::Remote => formatter.write_str("remote"), + } + } +} + +pub(crate) async fn create_file_system_context( + implementation: FileSystemImplementation, +) -> Result { + match implementation { + FileSystemImplementation::Local => { + let helper_paths = test_codex_helper_paths()?; + let runtime_paths = ExecServerRuntimePaths::new( + helper_paths.codex_exe.clone(), + helper_paths.codex_linux_sandbox_exe.clone(), + )?; + Ok(FileSystemContext { + file_system: Arc::new(LocalFileSystem::with_runtime_paths(runtime_paths)), + _helper_paths: Some(helper_paths), + _server: None, + }) + } + FileSystemImplementation::Remote => { + let server = exec_server().await?; + let environment = + Environment::create_for_tests(Some(server.websocket_url().to_string()))?; + Ok(FileSystemContext { + file_system: environment.get_filesystem(), + _helper_paths: None, + _server: Some(server), + }) + } + } +} + +#[cfg(windows)] +pub(crate) fn is_unsupported_restricted_token_host(result: &std::io::Result) -> bool { + result + .as_ref() + .err() + .is_some_and(|err| err.to_string().contains("CreateRestrictedToken failed: 87")) +} + +pub(crate) fn absolute_path(path: std::path::PathBuf) -> AbsolutePathBuf { + assert!( + path.is_absolute(), + "path must be absolute: {}", + path.display() + ); + AbsolutePathBuf::try_from(path).expect("path should be absolute") +} + +pub(crate) fn read_only_sandbox(readable_root: std::path::PathBuf) -> FileSystemSandboxContext { + let readable_root = absolute_path(readable_root); + sandbox_context(vec![FileSystemSandboxEntry { + path: FileSystemPath::Path { + path: readable_root.into(), + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }]) +} + +#[cfg(not(windows))] +pub(crate) fn workspace_write_sandbox( + writable_root: std::path::PathBuf, +) -> FileSystemSandboxContext { + let writable_root = absolute_path(writable_root); + sandbox_context(vec![FileSystemSandboxEntry { + path: FileSystemPath::Path { + path: writable_root.into(), + }, + access: FileSystemAccessMode::Write, + missing_path_behavior: None, + }]) +} + +#[cfg(windows)] +pub(crate) fn workspace_write_sandbox( + writable_root: std::path::PathBuf, +) -> FileSystemSandboxContext { + let writable_root = absolute_path(writable_root); + // Keep the runtime policy aligned with the legacy workspace-write projection used by the + // unelevated restricted-token preflight. + let policy = FileSystemSandboxPolicy::restricted(vec![ + FileSystemSandboxEntry::new( + FileSystemPath::Special { + value: FileSystemSpecialPath::Root, + }, + FileSystemAccessMode::Read, + ), + FileSystemSandboxEntry::new( + FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + FileSystemAccessMode::Write, + ), + ]); + let mut sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + PathUri::from_abs_path(&writable_root), + ); + sandbox.windows_sandbox_selection = WindowsSandboxSelection::RestrictedToken; + sandbox +} + +fn sandbox_context(mut entries: Vec) -> FileSystemSandboxContext { + if cfg!(windows) { + // Restricted-token sandboxing cannot enforce read restrictions, so leave the root + // readable while exercising the requested write restrictions. + entries.push(FileSystemSandboxEntry::new( + FileSystemPath::Special { + value: FileSystemSpecialPath::Root, + }, + FileSystemAccessMode::Read, + )); + } + let mut sandbox = FileSystemSandboxContext::from_permission_profile( + PermissionProfile::from_runtime_permissions( + &FileSystemSandboxPolicy::restricted(entries), + NetworkSandboxPolicy::Restricted, + ), + ); + if cfg!(windows) { + sandbox.windows_sandbox_selection = WindowsSandboxSelection::RestrictedToken; + } + sandbox +} diff --git a/codex-rs/exec-server/tests/file_system_unix.rs b/codex-rs/exec-server/tests/file_system_unix.rs new file mode 100644 index 0000000000000000000000000000000000000000..4b5615611f3ef666985f8b4f8d2ed19211928497 --- /dev/null +++ b/codex-rs/exec-server/tests/file_system_unix.rs @@ -0,0 +1,1573 @@ +#![cfg(unix)] +#![allow(clippy::expect_used)] + +mod common; +#[cfg(target_os = "linux")] +#[path = "common/fake_bwrap.rs"] +mod fake_bwrap; + +#[path = "file_system/shared.rs"] +mod shared; +#[path = "file_system/support.rs"] +mod support; + +use std::ffi::CString; +use std::os::unix::ffi::OsStrExt; +use std::os::unix::fs::FileTypeExt; +use std::os::unix::fs::MetadataExt; +use std::os::unix::fs::PermissionsExt; +use std::os::unix::fs::symlink; +use std::os::unix::net::UnixListener; +use std::path::Path; +use std::path::PathBuf; +use std::process::Command; +use std::sync::Arc; +use std::time::Duration; +#[cfg(target_os = "linux")] +use std::time::UNIX_EPOCH; + +use anyhow::Context; +use anyhow::Result; +use codex_exec_server::CopyOptions; +use codex_exec_server::CreateDirectoryOptions; +#[cfg(target_os = "linux")] +use codex_exec_server::Environment; +use codex_exec_server::FileMetadata; +use codex_exec_server::FileSystemSandboxContext; +use codex_exec_server::GetMetadataOptions; +use codex_exec_server::ReadDirectoryEntry; +use codex_exec_server::ReadFileOptions; +use codex_exec_server::RemoveOptions; +use codex_exec_server::WalkEntry; +use codex_exec_server::WalkEntryKind; +use codex_exec_server::WalkOptions; +use codex_exec_server::WalkOutcome; +use codex_exec_server::WriteFileOptions; +use codex_protocol::models::PermissionProfile; +use codex_protocol::permissions::FileSystemAccessMode; +use codex_protocol::permissions::FileSystemPath; +use codex_protocol::permissions::FileSystemSandboxEntry; +use codex_protocol::permissions::FileSystemSandboxPolicy; +use codex_protocol::permissions::FileSystemSpecialPath; +use codex_protocol::permissions::NetworkSandboxPolicy; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +use tempfile::TempDir; +use test_case::test_case; +use tokio::time::timeout; + +#[cfg(target_os = "linux")] +use crate::common::exec_server::exec_server_with_env; +#[cfg(target_os = "linux")] +use crate::fake_bwrap::write_fake_bwrap; + +use crate::support::FileSystemImplementation; +use crate::support::create_file_system_context; +use crate::support::read_only_sandbox; +use crate::support::workspace_write_sandbox; + +fn assert_sandbox_denied(error: &std::io::Error) { + match error.kind() { + std::io::ErrorKind::InvalidInput | std::io::ErrorKind::PermissionDenied => { + let message = error.to_string(); + assert!( + message.contains("is not permitted") + || message.contains("Operation not permitted") + || message.contains("Permission denied"), + "unexpected sandbox error message: {message}", + ); + } + std::io::ErrorKind::NotFound => assert!( + error.to_string().contains("No such file or directory"), + "unexpected sandbox not-found message: {error}", + ), + std::io::ErrorKind::Other => assert!( + error.to_string().contains("Read-only file system"), + "unexpected sandbox other error message: {error}", + ), + other => panic!("unexpected sandbox error kind: {other:?}: {error:?}"), + } +} + +fn assert_normalized_path_rejected(error: &std::io::Error) { + match error.kind() { + std::io::ErrorKind::NotFound => assert!( + error.to_string().contains("No such file or directory"), + "unexpected not-found message: {error}", + ), + std::io::ErrorKind::InvalidInput | std::io::ErrorKind::PermissionDenied => { + let message = error.to_string(); + assert!( + message.contains("is not permitted") + || message.contains("Operation not permitted") + || message.contains("Permission denied"), + "unexpected rejection message: {message}", + ); + } + other => panic!("unexpected normalized-path error kind: {other:?}: {error:?}"), + } +} + +fn alias_root_candidate() -> Result> { + for root in [Path::new("/tmp").to_path_buf(), std::env::temp_dir()] { + if root.is_dir() && root.canonicalize().is_ok_and(|canonical| canonical != root) { + return Ok(Some(root)); + } + } + Ok(None) +} + +fn create_directory_symlink(target: &Path, alias: &Path) -> Result<()> { + symlink(target, alias)?; + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_canonicalize_resolves_directory_symlink( + implementation: FileSystemImplementation, +) -> Result<()> { + shared::assert_canonicalize_resolves_directory_alias(implementation, create_directory_symlink) + .await +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_operations_can_reject_symlinks_in_any_path_component( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let tmp = TempDir::new()?; + let tmp_path = tmp.path().canonicalize()?; + let real = tmp_path.join("real"); + std::fs::create_dir(&real)?; + let existing = real.join("existing.txt"); + std::fs::write(&existing, "unchanged")?; + let removable = real.join("removable.txt"); + std::fs::write(&removable, "keep")?; + let directory_link = tmp_path.join("directory-link"); + symlink(&real, &directory_link)?; + let file_link = tmp_path.join("file-link"); + symlink(&existing, &file_link)?; + + let no_follow_read = ReadFileOptions { + follow_symlinks: false, + }; + let no_follow_write = WriteFileOptions { + follow_symlinks: false, + }; + let no_follow_metadata = GetMetadataOptions { + follow_symlinks: false, + }; + let no_follow_create = CreateDirectoryOptions { + recursive: true, + follow_symlinks: false, + }; + let no_follow_remove = RemoveOptions { + recursive: false, + force: false, + follow_symlinks: false, + }; + let uri = |path: &Path| PathUri::from_host_native_path(path); + #[cfg(target_os = "linux")] + let strict_sandbox = workspace_write_sandbox(tmp_path); + #[cfg(target_os = "linux")] + let sandboxes = [None, Some(&strict_sandbox)]; + #[cfg(not(target_os = "linux"))] + let sandboxes: [Option<&FileSystemSandboxContext>; 1] = [None]; + + for sandbox in sandboxes { + assert!( + context + .file_system + .read_file(&uri(&file_link)?, no_follow_read, sandbox) + .await + .is_err() + ); + assert!( + context + .file_system + .read_file( + &uri(&directory_link.join("existing.txt"))?, + no_follow_read, + sandbox, + ) + .await + .is_err() + ); + assert!( + context + .file_system + .write_file( + &uri(&file_link)?, + b"changed".to_vec(), + no_follow_write, + sandbox, + ) + .await + .is_err() + ); + assert_eq!(std::fs::read_to_string(&existing)?, "unchanged"); + assert!( + context + .file_system + .write_file( + &uri(&directory_link.join("existing.txt"))?, + b"changed".to_vec(), + no_follow_write, + sandbox, + ) + .await + .is_err() + ); + assert_eq!(std::fs::read_to_string(&existing)?, "unchanged"); + assert!( + context + .file_system + .get_metadata(&uri(&file_link)?, no_follow_metadata, sandbox) + .await + .is_err() + ); + let directory_metadata = context + .file_system + .get_metadata(&uri(&real)?, no_follow_metadata, sandbox) + .await?; + assert!(directory_metadata.is_directory); + assert!( + context + .file_system + .create_directory( + &uri(&directory_link.join("created"))?, + no_follow_create, + sandbox, + ) + .await + .is_err() + ); + assert!(!real.join("created").exists()); + assert!( + context + .file_system + .remove( + &uri(&directory_link.join("removable.txt"))?, + no_follow_remove, + sandbox, + ) + .await + .is_err() + ); + assert!(removable.exists()); + assert!( + context + .file_system + .remove(&uri(&file_link)?, no_follow_remove, sandbox) + .await + .is_err() + ); + assert!(file_link.symlink_metadata()?.file_type().is_symlink()); + } + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test] +async fn file_system_no_follow_non_recursive_root_creation_fails( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let result = context + .file_system + .create_directory( + &PathUri::from_host_native_path(Path::new("/"))?, + CreateDirectoryOptions { + recursive: false, + follow_symlinks: false, + }, + /*sandbox*/ None, + ) + .await; + + assert!(result.is_err()); + Ok(()) +} + +#[cfg(target_os = "linux")] +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test] +async fn file_system_no_follow_metadata_preserves_linux_birthtime( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let tmp = TempDir::new()?; + let file = tmp.path().join("created.txt"); + std::fs::write(&file, "created")?; + let expected = match std::fs::metadata(&file)?.created() { + Ok(created) => Some(i64::try_from( + created.duration_since(UNIX_EPOCH)?.as_millis(), + )?), + Err(error) if error.kind() == std::io::ErrorKind::Unsupported => None, + Err(error) => return Err(error.into()), + }; + + let metadata = context + .file_system + .get_metadata( + &PathUri::from_host_native_path(&file)?, + GetMetadataOptions { + follow_symlinks: false, + }, + /*sandbox*/ None, + ) + .await?; + + if let Some(expected) = expected { + assert!(expected > 0); + assert_eq!(metadata.created_at_ms, expected); + } else { + assert_eq!(metadata.created_at_ms, 0); + } + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_no_follow_operations_support_search_only_ancestors( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let tmp = TempDir::new()?; + let root = tmp.path().canonicalize()?; + let search_only = root.join("search-only"); + std::fs::create_dir(&search_only)?; + let existing = search_only.join("existing.txt"); + std::fs::write(&existing, "before")?; + let unreadable = search_only.join("unreadable.txt"); + std::fs::write(&unreadable, "metadata only")?; + std::fs::set_permissions(&unreadable, std::fs::Permissions::from_mode(0o000))?; + let socket_path = search_only.join("socket"); + let _socket = UnixListener::bind(&socket_path)?; + let removable = search_only.join("removable.txt"); + std::fs::write(&removable, "remove")?; + std::fs::set_permissions(&search_only, std::fs::Permissions::from_mode(0o300))?; + + let uri = |path: &Path| PathUri::from_host_native_path(path); + let result: Result<()> = async { + let root_metadata = context + .file_system + .get_metadata( + &uri(Path::new("/"))?, + GetMetadataOptions { + follow_symlinks: false, + }, + /*sandbox*/ None, + ) + .await?; + assert!(root_metadata.is_directory); + + assert_eq!( + context + .file_system + .read_file( + &uri(&existing)?, + ReadFileOptions { + follow_symlinks: false, + }, + /*sandbox*/ None, + ) + .await?, + b"before" + ); + context + .file_system + .write_file( + &uri(&existing)?, + b"after".to_vec(), + WriteFileOptions { + follow_symlinks: false, + }, + /*sandbox*/ None, + ) + .await?; + assert_eq!(std::fs::read_to_string(&existing)?, "after"); + + for metadata_path in [&unreadable, &socket_path] { + context + .file_system + .get_metadata( + &uri(metadata_path)?, + GetMetadataOptions { + follow_symlinks: false, + }, + /*sandbox*/ None, + ) + .await?; + } + + let nested = search_only.join("created").join("nested"); + context + .file_system + .create_directory( + &uri(&nested)?, + CreateDirectoryOptions { + recursive: true, + follow_symlinks: false, + }, + /*sandbox*/ None, + ) + .await?; + assert!(nested.is_dir()); + + context + .file_system + .remove( + &uri(&removable)?, + RemoveOptions { + recursive: false, + force: false, + follow_symlinks: false, + }, + /*sandbox*/ None, + ) + .await?; + assert!(!removable.exists()); + Ok(()) + } + .await; + + std::fs::set_permissions(&search_only, std::fs::Permissions::from_mode(0o700))?; + std::fs::set_permissions(&unreadable, std::fs::Permissions::from_mode(0o600))?; + result +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_no_follow_write_rejects_fifo_without_blocking( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let tmp = TempDir::new()?; + let fifo = tmp.path().canonicalize()?.join("fifo"); + let fifo_c = CString::new(fifo.as_os_str().as_bytes())?; + if unsafe { libc::mkfifo(fifo_c.as_ptr(), 0o600) } != 0 { + return Err(std::io::Error::last_os_error().into()); + } + + let result = timeout( + Duration::from_secs(1), + context.file_system.write_file( + &PathUri::from_host_native_path(&fifo)?, + b"must not be written".to_vec(), + WriteFileOptions { + follow_symlinks: false, + }, + /*sandbox*/ None, + ), + ) + .await + .context("strict FIFO write must not block")?; + assert!(result.is_err()); + assert!(fifo.symlink_metadata()?.file_type().is_fifo()); + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn file_system_no_follow_recursive_mkdir_handles_concurrent_creators( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let tmp = TempDir::new()?; + let path = tmp.path().canonicalize()?.join("shared").join("nested"); + let path_uri = PathUri::from_host_native_path(&path)?; + let barrier = Arc::new(tokio::sync::Barrier::new(16)); + let mut tasks = Vec::new(); + for _ in 0..16 { + let file_system = Arc::clone(&context.file_system); + let path_uri = path_uri.clone(); + let barrier = Arc::clone(&barrier); + tasks.push(tokio::spawn(async move { + barrier.wait().await; + file_system + .create_directory( + &path_uri, + CreateDirectoryOptions { + recursive: true, + follow_symlinks: false, + }, + /*sandbox*/ None, + ) + .await + })); + } + for task in tasks { + task.await??; + } + assert!(path.is_dir()); + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_sandboxed_canonicalize_resolves_directory_symlink( + implementation: FileSystemImplementation, +) -> Result<()> { + shared::assert_sandboxed_canonicalize_resolves_directory_alias( + implementation, + create_directory_symlink, + ) + .await +} + +#[cfg(target_os = "linux")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn sandboxed_file_system_helper_finds_bwrap_on_preserved_path() -> Result<()> { + let tmp = TempDir::new()?; + let fake_bin_dir = tmp.path().join("bin"); + let fake_bwrap = write_fake_bwrap(&fake_bin_dir)?; + let mut path_entries = vec![fake_bin_dir]; + if let Some(path) = std::env::var_os("PATH") { + path_entries.extend(std::env::split_paths(&path)); + } + let helper_path = std::env::join_paths(path_entries)?; + + let server = exec_server_with_env([("PATH", helper_path.as_os_str())], &[]).await?; + let environment = Environment::create_for_tests(Some(server.websocket_url().to_string()))?; + let file_system = environment.get_filesystem(); + let workspace = tmp.path().join("workspace"); + std::fs::create_dir_all(&workspace)?; + let file_path = workspace.join("created.txt"); + let sandbox = workspace_write_sandbox(workspace); + + file_system + .write_file( + &PathUri::from_host_native_path(&file_path)?, + b"written through fs helper".to_vec(), + Default::default(), + Some(&sandbox), + ) + .await?; + + assert_eq!(std::fs::read(&file_path)?, b"written through fs helper"); + + let bwrap_log = fake_bwrap.with_file_name("bwrap.log"); + let log = std::fs::read_to_string(&bwrap_log) + .with_context(|| format!("expected fake bwrap log at {}", bwrap_log.display()))?; + assert!( + log.contains("--argv0"), + "expected fs helper sandbox path to invoke PATH bwrap with --argv0, got: {log}" + ); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn remote_read_file_materializes_environment_workspace_roots() -> Result<()> { + let context = create_file_system_context(FileSystemImplementation::Remote).await?; + let file_system = context.file_system; + let tmp = TempDir::new()?; + let workspace = tmp.path().join("workspace"); + let workspace_file = workspace.join("included.txt"); + let excluded_file = tmp.path().join("excluded.txt"); + std::fs::create_dir(&workspace)?; + std::fs::write(&workspace_file, b"included")?; + std::fs::write(&excluded_file, b"excluded")?; + + let policy = FileSystemSandboxPolicy::restricted(vec![FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }]); + let mut sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + PathUri::from_host_native_path(tmp.path())?, + ); + sandbox.workspace_roots = vec![PathUri::from_host_native_path(&workspace)?]; + + assert_eq!( + file_system + .read_file( + &PathUri::from_host_native_path(&workspace_file)?, + Default::default(), + Some(&sandbox), + ) + .await?, + b"included" + ); + let error = file_system + .read_file( + &PathUri::from_host_native_path(&excluded_file)?, + Default::default(), + Some(&sandbox), + ) + .await + .expect_err("read outside environment workspace roots should fail"); + assert_sandbox_denied(&error); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn remote_read_file_preserves_empty_workspace_roots() -> Result<()> { + let context = create_file_system_context(FileSystemImplementation::Remote).await?; + let file_system = context.file_system; + let tmp = TempDir::new()?; + let file = tmp.path().join("excluded.txt"); + std::fs::write(&file, b"excluded")?; + + let policy = FileSystemSandboxPolicy::restricted(vec![FileSystemSandboxEntry { + path: FileSystemPath::Special { + value: FileSystemSpecialPath::project_roots(/*subpath*/ None), + }, + access: FileSystemAccessMode::Read, + missing_path_behavior: None, + }]); + let mut sandbox = FileSystemSandboxContext::from_permission_profile_with_cwd( + PermissionProfile::from_runtime_permissions(&policy, NetworkSandboxPolicy::Restricted), + PathUri::from_host_native_path(tmp.path())?, + ); + sandbox.workspace_roots.clear(); + + let error = file_system + .read_file( + &PathUri::from_host_native_path(&file)?, + Default::default(), + Some(&sandbox), + ) + .await + .expect_err("empty workspace roots should not grant cwd access"); + assert_sandbox_denied(&error); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_metadata_and_directory_listing_follow_symlinks( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let file_path = tmp.path().join("note.txt"); + std::fs::write(&file_path, "hello")?; + let symlink_path = tmp.path().join("note-link.txt"); + symlink(&file_path, &symlink_path)?; + let symlink_metadata = file_system + .get_metadata( + &PathUri::from_host_native_path(&symlink_path)?, + Default::default(), + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + symlink_metadata, + FileMetadata { + is_directory: false, + is_file: true, + is_symlink: true, + size: 5, + created_at_ms: symlink_metadata.created_at_ms, + modified_at_ms: symlink_metadata.modified_at_ms, + } + ); + assert!(symlink_metadata.modified_at_ms > 0); + + let dir_path = tmp.path().join("notes"); + std::fs::create_dir(&dir_path)?; + let dir_symlink_path = tmp.path().join("notes-link"); + symlink(&dir_path, &dir_symlink_path)?; + let dir_symlink_metadata = file_system + .get_metadata( + &PathUri::from_host_native_path(&dir_symlink_path)?, + Default::default(), + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + dir_symlink_metadata, + FileMetadata { + is_directory: true, + is_file: false, + is_symlink: true, + size: std::fs::metadata(&dir_path)?.len(), + created_at_ms: dir_symlink_metadata.created_at_ms, + modified_at_ms: dir_symlink_metadata.modified_at_ms, + } + ); + + let dangling_symlink_path = tmp.path().join("dangling-link"); + symlink(tmp.path().join("missing"), &dangling_symlink_path)?; + let error = file_system + .get_metadata( + &PathUri::from_host_native_path(&dangling_symlink_path)?, + Default::default(), + /*sandbox*/ None, + ) + .await + .expect_err("dangling symlink should not resolve"); + assert_eq!(error.kind(), std::io::ErrorKind::NotFound); + + let mut entries = file_system + .read_directory( + &PathUri::from_host_native_path(tmp.path())?, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + entries.retain(|entry| entry.file_name.contains("link")); + entries.sort_by(|left, right| left.file_name.cmp(&right.file_name)); + + assert_eq!( + entries, + vec![ + ReadDirectoryEntry { + file_name: "note-link.txt".to_string(), + is_directory: false, + is_file: true, + }, + ReadDirectoryEntry { + file_name: "notes-link".to_string(), + is_directory: true, + is_file: false, + }, + ] + ); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_walk_handles_directory_symlinks( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let root = tmp.path().join("root"); + let target = tmp.path().join("target"); + let target_file = target.join("note.txt"); + let target_link = root.join("target-link"); + let root_link = target.join("root-link"); + std::fs::create_dir_all(&root)?; + std::fs::create_dir_all(&target)?; + std::fs::write(&target_file, "target")?; + symlink(&target, &target_link)?; + symlink(&root, &root_link)?; + symlink(&target_file, root.join("file-link"))?; + symlink(root.join("missing"), root.join("broken-link"))?; + + for root in [&root, &root_link] { + let target_link = root.join("target-link"); + + let outcome = file_system + .walk( + &PathUri::from_host_native_path(root)?, + WalkOptions { + max_depth: 2, + max_directories: 4, + max_entries: 8, + follow_directory_symlinks: false, + prune_hidden_directories: false, + }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + outcome, + WalkOutcome { + entries: Vec::new(), + errors: Vec::new(), + truncated: false, + } + ); + + let outcome = file_system + .walk( + &PathUri::from_host_native_path(root)?, + WalkOptions { + max_depth: 2, + max_directories: 4, + max_entries: 8, + follow_directory_symlinks: true, + prune_hidden_directories: false, + }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + outcome, + WalkOutcome { + entries: vec![ + WalkEntry { + path: PathUri::from_host_native_path(&target_link)?, + kind: WalkEntryKind::Directory, + }, + WalkEntry { + path: PathUri::from_host_native_path(target_link.join("note.txt"))?, + kind: WalkEntryKind::File, + }, + WalkEntry { + path: PathUri::from_host_native_path(target_link.join("root-link"))?, + kind: WalkEntryKind::Directory, + }, + ], + errors: Vec::new(), + truncated: false, + } + ); + } + + Ok(()) +} + +#[cfg(target_os = "linux")] +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_walk_reports_non_utf8_names( + implementation: FileSystemImplementation, +) -> Result<()> { + use std::ffi::OsString; + use std::os::unix::ffi::OsStringExt; + + use codex_exec_server::WalkError; + + let context = create_file_system_context(implementation).await?; + let tmp = TempDir::new()?; + std::fs::write( + tmp.path() + .join(OsString::from_vec(b"invalid-\xff".to_vec())), + "contents", + )?; + let lossy_path = tmp.path().join("invalid-\u{fffd}"); + let error = + std::fs::symlink_metadata(&lossy_path).expect_err("the lossy filename must not exist"); + let outcome = context + .file_system + .walk( + &PathUri::from_host_native_path(tmp.path())?, + WalkOptions { + max_depth: 0, + max_directories: 1, + max_entries: 1, + follow_directory_symlinks: false, + prune_hidden_directories: false, + }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + assert_eq!( + outcome, + WalkOutcome { + entries: Vec::new(), + errors: vec![WalkError { + path: PathUri::from_host_native_path(lossy_path)?, + message: error.to_string(), + }], + truncated: false, + } + ); + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_walk_prunes_hidden_directories_without_claiming_visible_aliases( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let root = tmp.path().join("root"); + let hidden = root.join(".hidden"); + let hidden_nested = hidden.join("nested"); + let visible = root.join("visible"); + std::fs::create_dir_all(&hidden_nested)?; + std::fs::write(hidden_nested.join("note.txt"), "visible through alias")?; + symlink(&hidden, &visible)?; + + let outcome = file_system + .walk( + &PathUri::from_host_native_path(&root)?, + WalkOptions { + max_depth: 3, + max_directories: 3, + max_entries: 6, + follow_directory_symlinks: true, + prune_hidden_directories: true, + }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + + assert_eq!( + outcome, + WalkOutcome { + entries: vec![ + WalkEntry { + path: PathUri::from_host_native_path(hidden)?, + kind: WalkEntryKind::Directory, + }, + WalkEntry { + path: PathUri::from_host_native_path(&visible)?, + kind: WalkEntryKind::Directory, + }, + WalkEntry { + path: PathUri::from_host_native_path(visible.join("nested"))?, + kind: WalkEntryKind::Directory, + }, + WalkEntry { + path: PathUri::from_host_native_path(visible.join("nested/note.txt"))?, + kind: WalkEntryKind::File, + }, + ], + errors: Vec::new(), + truncated: false, + } + ); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_sandboxed_write_rejects_unwritable_path( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let blocked_path = tmp.path().join("blocked.txt"); + + let sandbox = read_only_sandbox(tmp.path().to_path_buf()); + let error = match file_system + .write_file( + &PathUri::from_host_native_path(&blocked_path)?, + b"nope".to_vec(), + Default::default(), + Some(&sandbox), + ) + .await + { + Ok(()) => anyhow::bail!("write should be blocked"), + Err(error) => error, + }; + assert_sandbox_denied(&error); + assert!(!blocked_path.exists()); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_sandboxed_write_allows_explicit_alias_roots( + implementation: FileSystemImplementation, +) -> Result<()> { + let Some(alias_root) = alias_root_candidate()? else { + return Ok(()); + }; + + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = tempfile::Builder::new() + .prefix("codex-fs-sandbox-alias-") + .tempdir_in(&alias_root)?; + let file_path = tmp.path().join("note.txt"); + let sandbox = workspace_write_sandbox(alias_root.clone()); + + file_system + .write_file( + &PathUri::from_host_native_path(&file_path)?, + b"created".to_vec(), + Default::default(), + Some(&sandbox), + ) + .await + .with_context(|| format!("write file through alias root mode={implementation}"))?; + assert_eq!(std::fs::read(&file_path)?, b"created"); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_sandboxed_read_rejects_symlink_escape( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + std::fs::write(outside_dir.join("secret.txt"), "nope")?; + symlink(&outside_dir, allowed_dir.join("link"))?; + + let requested_path = allowed_dir.join("link").join("secret.txt"); + let sandbox = read_only_sandbox(allowed_dir); + let error = match file_system + .read_file( + &PathUri::from_host_native_path(&requested_path)?, + Default::default(), + Some(&sandbox), + ) + .await + { + Ok(_) => anyhow::bail!("read should be blocked"), + Err(error) => error, + }; + assert_sandbox_denied(&error); + + let error = file_system + .read_file_stream( + &PathUri::from_host_native_path(&requested_path)?, + Some(&sandbox), + ) + .await + .err() + .context("streaming read should be blocked")?; + assert_sandbox_denied(&error); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_sandboxed_read_rejects_symlink_parent_dotdot_escape( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + let secret_path = tmp.path().join("secret.txt"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + std::fs::write(&secret_path, "nope")?; + symlink(&outside_dir, allowed_dir.join("link"))?; + + let requested_path = + PathUri::from_host_native_path(allowed_dir.join("link").join("..").join("secret.txt"))?; + let sandbox = read_only_sandbox(allowed_dir); + let error = match file_system + .read_file(&requested_path, Default::default(), Some(&sandbox)) + .await + { + Ok(_) => anyhow::bail!("read should fail after path normalization"), + Err(error) => error, + }; + // PathUri's native path constructor normalizes `link/../secret.txt` to + // `allowed/secret.txt` before the request reaches the filesystem layer. + // Depending on whether the platform/runtime resolves that normalized path + // through a top-level symlink alias, the request can surface as either + // "missing file" or an upfront sandbox rejection. + assert_normalized_path_rejected(&error); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_sandboxed_write_rejects_symlink_escape( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + symlink(&outside_dir, allowed_dir.join("link"))?; + + let requested_path = allowed_dir.join("link").join("blocked.txt"); + let sandbox = workspace_write_sandbox(allowed_dir); + let error = match file_system + .write_file( + &PathUri::from_host_native_path(&requested_path)?, + b"nope".to_vec(), + Default::default(), + Some(&sandbox), + ) + .await + { + Ok(()) => anyhow::bail!("write should be blocked"), + Err(error) => error, + }; + assert_sandbox_denied(&error); + assert!(!outside_dir.join("blocked.txt").exists()); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_sandboxed_write_preserves_existing_hard_link( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + + let outside_file = outside_dir.join("outside.txt"); + let hard_link = allowed_dir.join("hard-link.txt"); + std::fs::write(&outside_file, "outside\n")?; + std::fs::hard_link(&outside_file, &hard_link)?; + + let sandbox = workspace_write_sandbox(allowed_dir); + file_system + .write_file( + &PathUri::from_host_native_path(&hard_link)?, + b"updated through existing hard link\n".to_vec(), + Default::default(), + Some(&sandbox), + ) + .await + .with_context(|| format!("mode={implementation}"))?; + + assert_eq!( + std::fs::read_to_string(&outside_file)?, + "updated through existing hard link\n" + ); + assert_eq!( + std::fs::read_to_string(&hard_link)?, + "updated through existing hard link\n" + ); + + let outside_metadata = std::fs::metadata(&outside_file)?; + let link_metadata = std::fs::metadata(&hard_link)?; + assert_eq!( + (link_metadata.dev(), link_metadata.ino()), + (outside_metadata.dev(), outside_metadata.ino()) + ); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_create_directory_rejects_symlink_escape( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + symlink(&outside_dir, allowed_dir.join("link"))?; + + let requested_path = allowed_dir.join("link").join("created"); + let sandbox = workspace_write_sandbox(allowed_dir); + let error = match file_system + .create_directory( + &PathUri::from_host_native_path(&requested_path)?, + CreateDirectoryOptions { + recursive: false, + follow_symlinks: true, + }, + Some(&sandbox), + ) + .await + { + Ok(()) => anyhow::bail!("create_directory should be blocked"), + Err(error) => error, + }; + assert_sandbox_denied(&error); + assert!(!outside_dir.join("created").exists()); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_read_directory_rejects_symlink_escape( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + std::fs::write(outside_dir.join("secret.txt"), "nope")?; + symlink(&outside_dir, allowed_dir.join("link"))?; + + let requested_path = allowed_dir.join("link"); + let sandbox = read_only_sandbox(allowed_dir); + let error = match file_system + .read_directory( + &PathUri::from_host_native_path(&requested_path)?, + Some(&sandbox), + ) + .await + { + Ok(_) => anyhow::bail!("read_directory should be blocked"), + Err(error) => error, + }; + assert_sandbox_denied(&error); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_copy_rejects_symlink_escape_destination( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + std::fs::write(allowed_dir.join("source.txt"), "hello")?; + symlink(&outside_dir, allowed_dir.join("link"))?; + + let requested_destination = allowed_dir.join("link").join("copied.txt"); + let sandbox = workspace_write_sandbox(allowed_dir.clone()); + let error = match file_system + .copy( + &PathUri::from_host_native_path(allowed_dir.join("source.txt"))?, + &PathUri::from_host_native_path(&requested_destination)?, + CopyOptions { recursive: false }, + Some(&sandbox), + ) + .await + { + Ok(()) => anyhow::bail!("copy should be blocked"), + Err(error) => error, + }; + assert_sandbox_denied(&error); + assert!(!outside_dir.join("copied.txt").exists()); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_remove_removes_symlink_not_target( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + let outside_file = outside_dir.join("keep.txt"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + std::fs::write(&outside_file, "outside")?; + let symlink_path = allowed_dir.join("link"); + symlink(&outside_file, &symlink_path)?; + + let sandbox = workspace_write_sandbox(allowed_dir); + file_system + .remove( + &PathUri::from_host_native_path(&symlink_path)?, + RemoveOptions { + recursive: false, + force: false, + follow_symlinks: true, + }, + Some(&sandbox), + ) + .await + .with_context(|| format!("mode={implementation}"))?; + + assert!(!symlink_path.exists()); + assert!(outside_file.exists()); + assert_eq!(std::fs::read_to_string(outside_file)?, "outside"); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_copy_preserves_symlink_source( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + let outside_file = outside_dir.join("outside.txt"); + let source_symlink = allowed_dir.join("link"); + let copied_symlink = allowed_dir.join("copied-link"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + std::fs::write(&outside_file, "outside")?; + symlink(&outside_file, &source_symlink)?; + + let sandbox = workspace_write_sandbox(allowed_dir.clone()); + file_system + .copy( + &PathUri::from_host_native_path(&source_symlink)?, + &PathUri::from_host_native_path(&copied_symlink)?, + CopyOptions { recursive: false }, + Some(&sandbox), + ) + .await + .with_context(|| format!("mode={implementation}"))?; + + let copied_metadata = std::fs::symlink_metadata(&copied_symlink)?; + assert!(copied_metadata.file_type().is_symlink()); + assert_eq!(std::fs::read_link(copied_symlink)?, outside_file); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_remove_rejects_symlink_escape( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + let outside_file = outside_dir.join("secret.txt"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + std::fs::write(&outside_file, "outside")?; + symlink(&outside_dir, allowed_dir.join("link"))?; + + let requested_path = allowed_dir.join("link").join("secret.txt"); + let sandbox = workspace_write_sandbox(allowed_dir); + let error = match file_system + .remove( + &PathUri::from_host_native_path(&requested_path)?, + RemoveOptions { + recursive: false, + force: false, + follow_symlinks: true, + }, + Some(&sandbox), + ) + .await + { + Ok(()) => anyhow::bail!("remove should be blocked"), + Err(error) => error, + }; + assert_sandbox_denied(&error); + assert_eq!(std::fs::read_to_string(outside_file)?, "outside"); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_copy_rejects_symlink_escape_source( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let allowed_dir = tmp.path().join("allowed"); + let outside_dir = tmp.path().join("outside"); + let outside_file = outside_dir.join("secret.txt"); + let requested_destination = allowed_dir.join("copied.txt"); + std::fs::create_dir_all(&allowed_dir)?; + std::fs::create_dir_all(&outside_dir)?; + std::fs::write(&outside_file, "outside")?; + symlink(&outside_dir, allowed_dir.join("link"))?; + + let requested_source = allowed_dir.join("link").join("secret.txt"); + let sandbox = workspace_write_sandbox(allowed_dir); + let error = match file_system + .copy( + &PathUri::from_host_native_path(&requested_source)?, + &PathUri::from_host_native_path(&requested_destination)?, + CopyOptions { recursive: false }, + Some(&sandbox), + ) + .await + { + Ok(()) => anyhow::bail!("copy should be blocked"), + Err(error) => error, + }; + assert_sandbox_denied(&error); + assert!(!requested_destination.exists()); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_copy_preserves_symlinks_in_recursive_copy( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + let nested_dir = source_dir.join("nested"); + let copied_dir = tmp.path().join("copied"); + std::fs::create_dir_all(&nested_dir)?; + symlink("nested", source_dir.join("nested-link"))?; + + file_system + .copy( + &PathUri::from_host_native_path(&source_dir)?, + &PathUri::from_host_native_path(&copied_dir)?, + CopyOptions { recursive: true }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + + let copied_link = copied_dir.join("nested-link"); + let metadata = std::fs::symlink_metadata(&copied_link)?; + assert!(metadata.file_type().is_symlink()); + assert_eq!( + std::fs::read_link(copied_link)?, + std::path::PathBuf::from("nested") + ); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_copy_ignores_unknown_special_files_in_recursive_copy( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let source_dir = tmp.path().join("source"); + let copied_dir = tmp.path().join("copied"); + std::fs::create_dir_all(&source_dir)?; + std::fs::write(source_dir.join("note.txt"), "hello")?; + + let fifo_path = source_dir.join("named-pipe"); + let output = Command::new("mkfifo").arg(&fifo_path).output()?; + if !output.status.success() { + anyhow::bail!( + "mkfifo failed: stdout={} stderr={}", + String::from_utf8_lossy(&output.stdout).trim(), + String::from_utf8_lossy(&output.stderr).trim() + ); + } + + file_system + .copy( + &PathUri::from_host_native_path(&source_dir)?, + &PathUri::from_host_native_path(&copied_dir)?, + CopyOptions { recursive: true }, + /*sandbox*/ None, + ) + .await + .with_context(|| format!("mode={implementation}"))?; + + assert_eq!( + std::fs::read_to_string(copied_dir.join("note.txt"))?, + "hello" + ); + assert!(!copied_dir.join("named-pipe").exists()); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_copy_rejects_standalone_fifo_source( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let file_system = context.file_system; + + let tmp = TempDir::new()?; + let fifo_path = tmp.path().join("named-pipe"); + let output = Command::new("mkfifo").arg(&fifo_path).output()?; + if !output.status.success() { + anyhow::bail!( + "mkfifo failed: stdout={} stderr={}", + String::from_utf8_lossy(&output.stdout).trim(), + String::from_utf8_lossy(&output.stderr).trim() + ); + } + + let error = file_system + .copy( + &PathUri::from_host_native_path(&fifo_path)?, + &PathUri::from_host_native_path(tmp.path().join("copied"))?, + CopyOptions { recursive: false }, + /*sandbox*/ None, + ) + .await; + let error = error.expect_err("copying a FIFO should fail"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput); + assert_eq!( + error.to_string(), + "fs/copy only supports regular files, directories, and symlinks" + ); + + Ok(()) +} diff --git a/codex-rs/exec-server/tests/file_system_windows.rs b/codex-rs/exec-server/tests/file_system_windows.rs new file mode 100644 index 0000000000000000000000000000000000000000..7fe9df0f8ec938514724a64f0304e82b8f345109 --- /dev/null +++ b/codex-rs/exec-server/tests/file_system_windows.rs @@ -0,0 +1,601 @@ +#![cfg(windows)] +#![allow(clippy::expect_used)] + +mod common; + +#[path = "file_system/shared.rs"] +mod shared; +#[path = "file_system/support.rs"] +mod support; + +use std::collections::BTreeSet; +use std::ffi::c_void; +use std::path::Path; +use std::process::Command; +use std::time::Duration; + +use anyhow::Result; +use codex_exec_server::CreateDirectoryOptions; +use codex_exec_server::FileSystemSandboxContext; +use codex_exec_server::GetMetadataOptions; +use codex_exec_server::ReadFileOptions; +use codex_exec_server::RemoveOptions; +use codex_exec_server::WindowsSandboxSelection; +use codex_exec_server::WriteFileOptions; +use codex_protocol::protocol::SandboxPolicy; +use codex_sandboxing::SandboxType; +use codex_utils_path_uri::PathUri; +use futures::TryStreamExt; +use pretty_assertions::assert_eq; +use test_case::test_case; +use tokio::net::windows::named_pipe::ServerOptions; +use tokio::time::timeout; +use uuid::Uuid; +use windows_sys::Win32::System::Threading::GetCurrentProcess; + +use crate::support::FileSystemImplementation; +use crate::support::create_file_system_context; +use crate::support::is_unsupported_restricted_token_host; +use crate::support::workspace_write_sandbox; + +fn create_directory_junction(target: &Path, alias: &Path) -> Result<()> { + let output = Command::new("cmd") + .args(["/C", "mklink", "/J"]) + .arg(alias) + .arg(target) + .output()?; + if !output.status.success() { + anyhow::bail!( + "mklink /J failed: stdout={} stderr={}", + String::from_utf8_lossy(&output.stdout).trim(), + String::from_utf8_lossy(&output.stderr).trim() + ); + } + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_canonicalize_resolves_directory_junction( + implementation: FileSystemImplementation, +) -> Result<()> { + shared::assert_canonicalize_resolves_directory_alias(implementation, create_directory_junction) + .await +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_sandboxed_canonicalize_resolves_directory_junction( + implementation: FileSystemImplementation, +) -> Result<()> { + shared::assert_sandboxed_canonicalize_resolves_directory_alias( + implementation, + create_directory_junction, + ) + .await +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_operations_can_reject_junctions_in_any_path_component( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let tmp = tempfile::TempDir::new()?; + let real = tmp.path().join("real"); + std::fs::create_dir(&real)?; + let existing = real.join("existing.txt"); + std::fs::write(&existing, "unchanged")?; + let removable = real.join("removable.txt"); + std::fs::write(&removable, "keep")?; + let directory_junction = tmp.path().join("directory-junction"); + create_directory_junction(&real, &directory_junction)?; + + let no_follow_read = ReadFileOptions { + follow_symlinks: false, + }; + let no_follow_write = WriteFileOptions { + follow_symlinks: false, + }; + let no_follow_metadata = GetMetadataOptions { + follow_symlinks: false, + }; + let no_follow_create = CreateDirectoryOptions { + recursive: true, + follow_symlinks: false, + }; + let no_follow_remove = RemoveOptions { + recursive: false, + force: false, + follow_symlinks: false, + }; + let uri = |path: &Path| PathUri::from_host_native_path(path); + + assert_eq!( + context + .file_system + .read_file(&uri(&existing)?, no_follow_read, /*sandbox*/ None,) + .await?, + b"unchanged" + ); + + let file_link_target = real.join("file-link-target.txt"); + std::fs::write(&file_link_target, "target")?; + let file_link = tmp.path().join("file-link.txt"); + if std::os::windows::fs::symlink_file(&file_link_target, &file_link).is_ok() { + assert!( + context + .file_system + .write_file( + &uri(&file_link)?, + b"changed".to_vec(), + no_follow_write, + /*sandbox*/ None, + ) + .await + .is_err() + ); + assert_eq!(std::fs::read_to_string(&file_link_target)?, "target"); + } + + assert!( + context + .file_system + .read_file( + &uri(&directory_junction.join("existing.txt"))?, + no_follow_read, + /*sandbox*/ None, + ) + .await + .is_err() + ); + assert!( + context + .file_system + .write_file( + &uri(&directory_junction.join("existing.txt"))?, + b"changed".to_vec(), + no_follow_write, + /*sandbox*/ None, + ) + .await + .is_err() + ); + assert_eq!(std::fs::read_to_string(&existing)?, "unchanged"); + assert!( + context + .file_system + .get_metadata( + &uri(&directory_junction)?, + no_follow_metadata, + /*sandbox*/ None, + ) + .await + .is_err() + ); + let directory_metadata = context + .file_system + .get_metadata(&uri(&real)?, no_follow_metadata, /*sandbox*/ None) + .await?; + assert!(directory_metadata.is_directory); + assert!( + context + .file_system + .create_directory( + &uri(&directory_junction.join("created"))?, + no_follow_create, + /*sandbox*/ None, + ) + .await + .is_err() + ); + assert!(!real.join("created").exists()); + assert!( + context + .file_system + .remove( + &uri(&directory_junction.join("removable.txt"))?, + no_follow_remove, + /*sandbox*/ None, + ) + .await + .is_err() + ); + assert!(removable.exists()); + assert!( + context + .file_system + .remove( + &uri(&directory_junction)?, + no_follow_remove, + /*sandbox*/ None, + ) + .await + .is_err() + ); + assert!( + directory_junction + .symlink_metadata()? + .file_type() + .is_symlink() + ); + + let sandbox = workspace_write_sandbox(tmp.path().to_path_buf()); + let read_result = context + .file_system + .read_file(&uri(&existing)?, no_follow_read, Some(&sandbox)) + .await; + if is_unsupported_restricted_token_host(&read_result) { + return Ok(()); + } + assert_eq!(read_result?, b"unchanged"); + assert!( + context + .file_system + .read_file( + &uri(&directory_junction.join("existing.txt"))?, + no_follow_read, + Some(&sandbox), + ) + .await + .is_err() + ); + assert!( + context + .file_system + .write_file( + &uri(&directory_junction.join("existing.txt"))?, + b"changed".to_vec(), + no_follow_write, + Some(&sandbox), + ) + .await + .is_err() + ); + assert_eq!(std::fs::read_to_string(&existing)?, "unchanged"); + assert!( + context + .file_system + .get_metadata( + &uri(&directory_junction)?, + no_follow_metadata, + Some(&sandbox), + ) + .await + .is_err() + ); + assert!( + context + .file_system + .create_directory( + &uri(&directory_junction.join("sandbox-created"))?, + no_follow_create, + Some(&sandbox), + ) + .await + .is_err() + ); + assert!(!real.join("sandbox-created").exists()); + assert!( + context + .file_system + .remove( + &uri(&directory_junction.join("removable.txt"))?, + no_follow_remove, + Some(&sandbox), + ) + .await + .is_err() + ); + assert!(removable.exists()); + + Ok(()) +} + +#[test_case(FileSystemImplementation::Local ; "local")] +#[test_case(FileSystemImplementation::Remote ; "remote")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_no_follow_operations_reject_named_pipes( + implementation: FileSystemImplementation, +) -> Result<()> { + let context = create_file_system_context(implementation).await?; + let pipe_name = format!("codex-fs-no-follow-{}", Uuid::new_v4()); + let server_path = format!(r"\\.\pipe\{pipe_name}"); + let client_path = format!(r"\\localhost\pipe\{pipe_name}"); + let _pipe = ServerOptions::new() + .first_pipe_instance(true) + .create(&server_path)?; + + let error = timeout( + Duration::from_secs(1), + context.file_system.read_file( + &PathUri::from_host_native_path(Path::new(&client_path))?, + ReadFileOptions { + follow_symlinks: false, + }, + /*sandbox*/ None, + ), + ) + .await + .expect("strict named-pipe read must not hang") + .expect_err("strict named-pipe read must be rejected"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput); + + let pipe_name = format!("codex-fs-no-follow-write-{}", Uuid::new_v4()); + let server_path = format!(r"\\.\pipe\{pipe_name}"); + let client_path = format!(r"\\localhost\pipe\{pipe_name}"); + let _pipe = ServerOptions::new() + .first_pipe_instance(true) + .create(&server_path)?; + timeout( + Duration::from_secs(1), + context.file_system.write_file( + &PathUri::from_host_native_path(Path::new(&client_path))?, + b"must not be written".to_vec(), + WriteFileOptions { + follow_symlinks: false, + }, + /*sandbox*/ None, + ), + ) + .await + .expect("strict named-pipe write must not hang") + .expect_err("strict named-pipe write must be rejected"); + Ok(()) +} + +#[test_case(SandboxType::WindowsRestrictedToken; "restricted_token")] +#[test_case(SandboxType::WindowsMxc; "mxc")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_remote_fs_helper_respects_windows_sandbox_write_policy( + sandbox_type: SandboxType, +) -> Result<()> { + let context = create_file_system_context(FileSystemImplementation::Remote).await?; + let file_system = context.file_system; + let tmp = tempfile::TempDir::new()?; + let readonly_dir = tmp.path().join("readonly"); + std::fs::create_dir_all(&readonly_dir)?; + + let mut sandbox = read_only_sandbox_for_cwd(readonly_dir.clone())?; + match sandbox_type { + SandboxType::WindowsRestrictedToken => { + sandbox.windows_sandbox_selection = WindowsSandboxSelection::RestrictedToken; + } + SandboxType::WindowsMxc => { + sandbox.windows_sandbox_selection = WindowsSandboxSelection::Mxc; + } + SandboxType::None | SandboxType::MacosSeatbelt | SandboxType::LinuxSeccomp => { + anyhow::bail!("expected a Windows sandbox type") + } + } + + let blocked_file = readonly_dir.join("blocked.txt"); + if sandbox_type == SandboxType::WindowsMxc && !codex_sandboxing::windows_mxc_available() { + let error = file_system + .write_file( + &PathUri::from_host_native_path(&blocked_file)?, + b"blocked".to_vec(), + WriteFileOptions::default(), + Some(&sandbox), + ) + .await + .expect_err("unavailable MXC must fail closed"); + assert_eq!( + (error.kind(), error.to_string()), + ( + std::io::ErrorKind::InvalidInput, + "failed to prepare fs sandbox: failed to prepare MXC sandbox: native MXC is unavailable on this executor".to_owned(), + ) + ); + assert!(!blocked_file.exists()); + return Ok(()); + } + + let readable_file = readonly_dir.join("readable.txt"); + std::fs::write(&readable_file, b"readable")?; + let read_result = file_system + .read_file( + &PathUri::from_host_native_path(&readable_file)?, + ReadFileOptions::default(), + Some(&sandbox), + ) + .await; + // Some local Windows hosts cannot create restricted tokens. Reaching that + // error still proves the remote fs helper went through the Windows sandbox + // launcher; before the wrapper fix this read would have run unsandboxed. + if sandbox_type == SandboxType::WindowsRestrictedToken + && is_unsupported_restricted_token_host(&read_result) + { + return Ok(()); + } + assert_eq!(read_result?, b"readable"); + + let error = file_system + .write_file( + &PathUri::from_host_native_path(&blocked_file)?, + b"blocked".to_vec(), + WriteFileOptions::default(), + Some(&sandbox), + ) + .await + .expect_err("write outside the sandbox should fail"); + assert!( + !blocked_file.exists(), + "sandboxed fs helper must not create blocked file after error: {error}" + ); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn file_system_private_desktop_survives_helper_exits_and_separates_permissions() -> Result<()> +{ + let context = create_file_system_context(FileSystemImplementation::Local).await?; + let file_system = context.file_system; + let tmp = tempfile::TempDir::new()?; + let path = tmp.path().join("contents.txt"); + std::fs::write(&path, b"initial")?; + let uri = PathUri::from_host_native_path(&path)?; + let mut sandbox = workspace_write_sandbox(tmp.path().to_path_buf()); + sandbox.windows_sandbox_private_desktop = true; + let before = process_private_desktops()?; + let read = file_system + .read_file(&uri, ReadFileOptions::default(), Some(&sandbox)) + .await; + if is_unsupported_restricted_token_host(&read) { + eprintln!("Skipping private desktop reuse: this host cannot create restricted tokens"); + return Ok(()); + } + assert_eq!(read?, b"initial"); + + // Ownership must outlive each helper; a desktop held only by the helper disappears here. + let warmed = process_private_desktops()?; + assert_eq!(warmed.difference(&before).count(), 1); + for contents in ["updated", "updated again"] { + file_system + .write_file( + &uri, + contents.as_bytes().to_vec(), + WriteFileOptions::default(), + Some(&sandbox), + ) + .await?; + assert_eq!(process_private_desktops()?, warmed); + assert_eq!( + file_system + .read_file(&uri, ReadFileOptions::default(), Some(&sandbox)) + .await?, + contents.as_bytes() + ); + assert_eq!(process_private_desktops()?, warmed); + assert!( + file_system + .get_metadata(&uri, GetMetadataOptions::default(), Some(&sandbox)) + .await? + .is_file + ); + assert_eq!(process_private_desktops()?, warmed); + let chunks = file_system + .read_file_stream(&uri, Some(&sandbox)) + .await? + .try_collect::>() + .await?; + assert_eq!(chunks.concat(), contents.as_bytes()); + assert_eq!(process_private_desktops()?, warmed); + } + + let mut readonly = read_only_sandbox_for_cwd(tmp.path().to_path_buf())?; + readonly.windows_sandbox_selection = WindowsSandboxSelection::RestrictedToken; + readonly.windows_sandbox_private_desktop = true; + assert_eq!( + file_system + .read_file(&uri, ReadFileOptions::default(), Some(&readonly)) + .await?, + b"updated again" + ); + let separated = process_private_desktops()?; + assert!(warmed.is_subset(&separated)); + assert_eq!(separated.difference(&warmed).count(), 1); + file_system + .write_file( + &uri, + b"blocked".to_vec(), + WriteFileOptions::default(), + Some(&readonly), + ) + .await + .expect_err("read-only filesystem requests must reject writes"); + assert_eq!(std::fs::read(&path)?, b"updated again"); + assert_eq!(process_private_desktops()?, separated); + Ok(()) +} + +fn process_private_desktops() -> Result> { + // Query this process so other tests' private desktops cannot affect the assertions. + // Native layout: https://github.com/winsiderss/phnt/blob/master/ntpsapi.h + #[repr(C)] + struct HandleEntry { + handle: isize, + _handle_count: usize, + _pointer_count: usize, + _granted_access: u32, + _object_type_index: u32, + _handle_attributes: u32, + _reserved: u32, + } + #[link(name = "ntdll")] + unsafe extern "system" { + fn NtQueryInformationProcess( + process: isize, + class: u32, + information: *mut c_void, + length: u32, + return_length: *mut u32, + ) -> i32; + } + #[link(name = "user32")] + unsafe extern "system" { + fn GetUserObjectInformationW( + object: isize, + index: i32, + information: *mut c_void, + length: u32, + length_needed: *mut u32, + ) -> i32; + } + let mut snapshot = vec![0usize; 8192]; + let bytes = std::mem::size_of_val(snapshot.as_slice()); + let status = unsafe { + NtQueryInformationProcess( + GetCurrentProcess(), + /*class*/ 51, + snapshot.as_mut_ptr().cast(), + bytes as u32, + std::ptr::null_mut(), + ) + }; + anyhow::ensure!(status >= 0, "process handle query failed: {status:#x}"); + let count = snapshot[0]; + anyhow::ensure!( + count <= (bytes - 2 * size_of::()) / size_of::(), + "process handle snapshot exceeds its buffer" + ); + let entries = unsafe { + std::slice::from_raw_parts(snapshot.as_ptr().add(2).cast::(), count) + }; + let mut desktops = BTreeSet::new(); + for entry in entries { + let mut name = [0u16; 64]; + let mut length_needed = 0; + if unsafe { + GetUserObjectInformationW( + entry.handle, + /*index*/ 2, + name.as_mut_ptr().cast(), + std::mem::size_of_val(&name) as u32, + &mut length_needed, + ) + } != 0 + { + let end = name + .iter() + .position(|&unit| unit == 0) + .unwrap_or(name.len()); + let name = String::from_utf16(&name[..end])?; + if name.starts_with("CodexSandboxDesktop-") { + desktops.insert(name); + } + } + } + Ok(desktops) +} + +fn read_only_sandbox_for_cwd(cwd: std::path::PathBuf) -> Result { + Ok(FileSystemSandboxContext::from_legacy_sandbox_policy( + SandboxPolicy::new_read_only_policy(), + PathUri::from_host_native_path(cwd)?, + )?) +} diff --git a/codex-rs/exec-server/tests/forward.rs b/codex-rs/exec-server/tests/forward.rs new file mode 100644 index 0000000000000000000000000000000000000000..526bd12119f1ace21197c0c5edc1358f54ca7dd5 --- /dev/null +++ b/codex-rs/exec-server/tests/forward.rs @@ -0,0 +1,141 @@ +mod common; + +#[path = "common/relay.rs"] +mod relay_support; + +use std::collections::HashMap; + +use anyhow::Result; +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD; +use codex_exec_server::ExecOutputStream; +use codex_exec_server::ExecParams; +use codex_exec_server::ExecServerError; +use codex_exec_server::FsReadFileParams; +use codex_exec_server::FsWriteFileParams; +use codex_exec_server::ProcessId; +use codex_exec_server::ReadParams; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +use relay_support::RelayTest; +use relay_support::TEST_TIMEOUT; +use tempfile::TempDir; +use tokio::time::timeout; +use tokio_util::task::AbortOnDropHandle; + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn forwarder_runs_commands_and_transfers_files() -> Result<()> { + let relay = RelayTest::new().await?; + let mut destination = common::exec_server::exec_server().await?; + let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel(); + let forwarder = AbortOnDropHandle::new(tokio::spawn( + codex_exec_server::run_remote_environment_forward_until_shutdown( + relay.config()?, + destination.websocket_url().to_string(), + async move { + let _ = shutdown_rx.await; + }, + ), + )); + let connection = relay.connect().await?; + let client = &connection.client; + + for (process_id, expected_output, expected_exit) in [ + ("forward-success", "forwarded output", 0), + ("forward-failure", "nonzero output", 7), + ] { + let process_id = ProcessId::from(process_id); + let argv = if cfg!(windows) { + vec![ + "cmd.exe".to_string(), + "/D".to_string(), + "/C".to_string(), + format!("echo {expected_output}& exit /B {expected_exit}"), + ] + } else { + vec![ + "/bin/sh".to_string(), + "-c".to_string(), + format!("printf '%s\\n' '{expected_output}'; exit {expected_exit}"), + ] + }; + timeout( + TEST_TIMEOUT, + client.exec(ExecParams { + metadata: Default::default(), + process_id: process_id.clone(), + argv, + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }), + ) + .await??; + let (output, exit_code) = timeout(TEST_TIMEOUT, async { + let mut output = Vec::new(); + let mut after_seq = None; + loop { + let response = client + .read(ReadParams { + process_id: process_id.clone(), + after_seq, + max_bytes: None, + wait_ms: Some(1_000), + }) + .await?; + assert_eq!(response.failure, None); + for chunk in response.chunks { + assert_eq!(chunk.stream, ExecOutputStream::Stdout); + output.extend(chunk.chunk.0); + } + if response.closed { + break Ok::<_, ExecServerError>((output, response.exit_code)); + } + after_seq = response.next_seq.checked_sub(1); + } + }) + .await??; + assert_eq!( + (String::from_utf8(output)?.replace("\r\n", "\n"), exit_code), + (format!("{expected_output}\n"), Some(expected_exit)) + ); + } + + // Base64 exceeds the destination's 16 MiB frame limit. + let temp_dir = TempDir::new()?; + let path = temp_dir.path().join("large-request.bin"); + let contents = vec![0xa5; 13 * 1024 * 1024]; + timeout( + TEST_TIMEOUT, + client.fs_write_file(FsWriteFileParams { + path: PathUri::from_host_native_path(&path)?, + follow_symlinks: None, + data_base64: STANDARD.encode(&contents), + sandbox: None, + }), + ) + .await??; + // A following request must still arrive as its own complete message. + let read_response = client + .fs_read_file(FsReadFileParams { + path: PathUri::from_host_native_path(path)?, + follow_symlinks: None, + sandbox: None, + }) + .await?; + assert_eq!(STANDARD.decode(read_response.data_base64)?, contents); + connection.assert_encrypted()?; + connection.close().await; + let _ = shutdown_tx.send(()); + timeout(TEST_TIMEOUT, forwarder).await???; + destination.shutdown().await?; + Ok(()) +} diff --git a/codex-rs/exec-server/tests/health.rs b/codex-rs/exec-server/tests/health.rs new file mode 100644 index 0000000000000000000000000000000000000000..e1b484494d22e3570f7176e029748a26389a3948 --- /dev/null +++ b/codex-rs/exec-server/tests/health.rs @@ -0,0 +1,92 @@ +#![cfg(unix)] + +mod common; + +use codex_exec_server::Environment; +use codex_exec_server::EnvironmentStatus; +use codex_exec_server::EnvironmentStatusKind; +use codex_exec_server::InitializeParams; +use codex_exec_server::InitializeResponse; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCResponse; +use common::exec_server::exec_server; +use pretty_assertions::assert_eq; + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_serves_readyz_alongside_websocket_endpoint() -> anyhow::Result<()> { + let mut server = exec_server().await?; + let http_base_url = server + .websocket_url() + .strip_prefix("ws://") + .expect("websocket URL should use ws://"); + + let client = codex_http_client::HttpClientBuilder::new().build_direct()?; + let response = client + .get(format!("http://{http_base_url}/readyz")) + .send() + .await?; + assert_eq!(response.status(), http::StatusCode::OK); + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn remote_environment_fetches_info_from_exec_server() -> anyhow::Result<()> { + let mut server = exec_server().await?; + let environment = Environment::create_for_tests(Some(server.websocket_url().to_string()))?; + assert!(environment.is_remote()); + + let remote_info = environment.info().await?; + let mut local_info = Environment::default_for_tests().info().await?; + // Only the remote executor advertises its optional build identity. + local_info.provider_id = remote_info.provider_id.clone(); + assert_eq!(remote_info, local_info); + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_reports_environment_status_over_websocket() -> anyhow::Result<()> { + let mut server = exec_server().await?; + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-health-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { + id, + result: initialize_result, + }) = server.next_event().await? + else { + panic!("expected initialize response"); + }; + assert_eq!(id, initialize_id); + let _: InitializeResponse = serde_json::from_value(initialize_result)?; + server + .send_notification("initialized", serde_json::json!({})) + .await?; + + let status_id = server + .send_request("environment/status", serde_json::json!({})) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = server.next_event().await? + else { + panic!("expected environment status response"); + }; + assert_eq!(id, status_id); + assert_eq!( + serde_json::from_value::(result)?, + EnvironmentStatus { + status: EnvironmentStatusKind::Ready, + } + ); + + server.shutdown().await?; + Ok(()) +} diff --git a/codex-rs/exec-server/tests/http_client.rs b/codex-rs/exec-server/tests/http_client.rs new file mode 100644 index 0000000000000000000000000000000000000000..ebe8471dee5c6abb55ada84858601baf2ce176fd --- /dev/null +++ b/codex-rs/exec-server/tests/http_client.rs @@ -0,0 +1,1547 @@ +use std::future::Future; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use codex_exec_server::ExecServerClient; +use codex_exec_server::HttpHeader; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestBodyDeltaNotification; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::InitializeParams; +use codex_exec_server::InitializeResponse; +use codex_exec_server::RemoteExecServerConnectArgs; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCNotification; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::MAX_HTTP_BODY_DELTA_BYTES; +use codex_exec_server_protocol::RequestId; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use futures::SinkExt; +use futures::StreamExt; +use pretty_assertions::assert_eq; +use serde::Serialize; +use serde::de::DeserializeOwned; +use serde_json::from_slice; +use serde_json::from_str; +use serde_json::from_value; +use serde_json::to_string; +use serde_json::to_value; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::sync::oneshot; +use tokio::task::JoinHandle; +use tokio::time::timeout; +use tokio_tungstenite::WebSocketStream; +use tokio_tungstenite::accept_async; +use tokio_tungstenite::tungstenite::Message; + +const CLIENT_NAME: &str = "test-exec-server-client"; +const HTTP_REQUEST_METHOD: &str = "http/request"; +const HTTP_REQUEST_BODY_DELTA_METHOD: &str = "http/request/bodyDelta"; +const INITIALIZE_METHOD: &str = "initialize"; +const INITIALIZED_METHOD: &str = "initialized"; +const TEST_TIMEOUT: Duration = Duration::from_secs(5); +const BYTE_BUDGET_TEST_TIMEOUT: Duration = Duration::from_secs(30); +const HTTP_BODY_DELTA_CHANNEL_CAPACITY: u64 = 256; +const HTTP_BODY_DELTA_BYTE_BUDGET: usize = 16 * 1024 * 1024; +const OVERFLOWING_BODY_DELTA_FRAMES: u64 = 1_024; + +/// What this tests: the buffered HTTP helper always sends a buffered +/// `http/request`, even when a caller accidentally provides streaming flags. +#[tokio::test] +async fn http_request_forces_buffered_request_params() -> Result<()> { + // Phase 1: start a fake WebSocket exec-server so the test covers the + // public client connection path without depending on the HTTP runner. + let server = spawn_scripted_exec_server(|mut peer| async move { + // Phase 2: verify the buffered helper forces buffered mode before it + // sends the JSON-RPC call. + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/buffered".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "ignored-stream-id".to_string(), + stream_response: false, + } + ); + + peer.write_response( + request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: b"buffered".to_vec().into(), + }, + ) + .await + }) + .await?; + let client = server.connect_client().await?; + + // Phase 3: call the buffered helper with streaming-only fields populated + // and assert callers still receive the buffered response body. + let response = timeout( + TEST_TIMEOUT, + client.http_request(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/buffered".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "ignored-stream-id".to_string(), + stream_response: true, + }), + ) + .await + .context("buffered http/request should complete")??; + assert_eq!( + response, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: b"buffered".to_vec().into(), + } + ); + + drop(client); + server.finish().await?; + Ok(()) +} + +/// What this tests: streamed executor HTTP response frames are routed by the +/// client's generated request id, delivered in sequence, and concatenated by +/// the caller. +#[tokio::test] +async fn http_response_body_stream_uses_generated_ids_and_receives_ordered_deltas() -> Result<()> { + // Phase 1: script two requests. The caller supplies reusable ids, but the + // client replaces them with connection-local ids on the wire. + let server = spawn_scripted_exec_server(|mut peer| async move { + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp".to_string(), + headers: vec![HttpHeader { + name: "accept".to_string(), + value: "text/event-stream".to_string(), + value_env_var: None, + }], + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-1".to_string(), + stream_response: true, + } + ); + + // Phase 2: return headers first, then body notifications in the order + // the public body stream should expose them. + peer.write_response( + request_id, + HttpRequestResponse { + status: 200, + headers: vec![HttpHeader { + name: "content-type".to_string(), + value: "text/event-stream".to_string(), + value_env_var: None, + }], + body: Vec::new().into(), + }, + ) + .await?; + for delta in [ + HttpRequestBodyDeltaNotification { + request_id: "http-1".to_string(), + seq: 1, + delta: b"hello ".to_vec().into(), + done: false, + error: None, + }, + HttpRequestBodyDeltaNotification { + request_id: "http-1".to_string(), + seq: 2, + delta: b"world".to_vec().into(), + done: false, + error: None, + }, + HttpRequestBodyDeltaNotification { + request_id: "http-1".to_string(), + seq: 3, + delta: b"!".to_vec().into(), + done: true, + error: None, + }, + ] { + peer.write_body_delta(delta).await?; + } + + // Phase 3: accept the next generated request id after EOF. + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/reuse".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-2".to_string(), + stream_response: true, + } + ); + peer.write_response( + request_id, + HttpRequestResponse { + status: 204, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await + }) + .await?; + let client = server.connect_client().await?; + + // Phase 4: start a streaming HTTP request through the public client API. + let (response, mut body_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp".to_string(), + headers: vec![HttpHeader { + name: "accept".to_string(), + value: "text/event-stream".to_string(), + value_env_var: None, + }], + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }), + ) + .await + .context("streamed http/request should return headers")??; + assert_eq!( + response, + HttpRequestResponse { + status: 200, + headers: vec![HttpHeader { + name: "content-type".to_string(), + value: "text/event-stream".to_string(), + value_env_var: None, + }], + body: Vec::new().into(), + } + ); + + // Phase 5: drain the body stream and verify the caller-visible byte order. + let mut body = Vec::new(); + while let Some(chunk) = timeout(TEST_TIMEOUT, body_stream.recv()) + .await + .context("http response body delta should arrive")?? + { + body.extend_from_slice(&chunk); + } + assert_eq!(body, b"hello world!".to_vec()); + + // Phase 6: start another stream through the public API to validate cleanup + // after EOF without reaching into the client routing table. + let (reuse_response, _reuse_body_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/reuse".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }), + ) + .await + .context("second streamed http/request should return headers")??; + assert_eq!( + reuse_response, + HttpRequestResponse { + status: 204, + headers: Vec::new(), + body: Vec::new().into(), + } + ); + + drop(client); + server.finish().await?; + Ok(()) +} + +/// What this tests: dropping a body stream with a queued terminal frame removes +/// the old route while the next stream gets a fresh generated id. +#[tokio::test] +async fn http_response_body_stream_drops_queued_terminal_before_next_generated_id() -> Result<()> { + // Phase 1: send terminal EOF before the header response so the public body + // stream starts with EOF already queued but unread. + let server = spawn_scripted_exec_server(|mut peer| async move { + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/queued-terminal".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-1".to_string(), + stream_response: true, + } + ); + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-1".to_string(), + seq: 1, + delta: Vec::new().into(), + done: true, + error: None, + }) + .await?; + peer.write_response( + request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await?; + + // Phase 2: accept another stream after the client drops the unread + // body. The second request receives a distinct generated id. + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/retry-queued-terminal".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-2".to_string(), + stream_response: true, + } + ); + peer.write_response( + request_id, + HttpRequestResponse { + status: 204, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await + }) + .await?; + let client = server.connect_client().await?; + + // Phase 3: drop the body stream without reading the queued EOF frame. + let (response, body_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/queued-terminal".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }), + ) + .await + .context("streamed http/request should return headers")??; + assert_eq!( + response, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + } + ); + drop(body_stream); + + // Phase 4: start another stream through the public API. The caller-provided + // id is ignored, so the request uses the next generated route id. + let params = HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/retry-queued-terminal".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }; + let (reuse_response, _reuse_body_stream) = + timeout(TEST_TIMEOUT, client.http_request_stream(params)) + .await + .context("second streamed http/request should return headers")??; + assert_eq!( + reuse_response, + HttpRequestResponse { + status: 204, + headers: Vec::new(), + body: Vec::new().into(), + } + ); + + drop(client); + server.finish().await?; + Ok(()) +} + +/// What this tests: cancelling a streaming HTTP request while it is waiting for +/// headers drops its route, and a later stream gets a fresh generated id. +#[tokio::test] +async fn http_response_body_stream_ignores_late_deltas_after_cancelled_request() -> Result<()> { + // Phase 1: coordinate cancellation after the fake server observes the + // first request but before it returns headers. The server later sends a + // stale delta for the cancelled id before serving the fresh stream. + let (request_seen_tx, request_seen_rx) = oneshot::channel(); + let server = spawn_scripted_exec_server(|mut peer| async move { + let (_request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/cancel".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-1".to_string(), + stream_response: true, + } + ); + request_seen_tx + .send(()) + .expect("test should wait for the first request"); + + // Phase 2: the next stream uses a new generated id. A late body delta + // for the cancelled id is ignored by the client-side router. + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/retry-cancelled".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-2".to_string(), + stream_response: true, + } + ); + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-1".to_string(), + seq: 1, + delta: b"stale".to_vec().into(), + done: false, + error: None, + }) + .await?; + peer.write_response( + request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await?; + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-2".to_string(), + seq: 1, + delta: b"fresh".to_vec().into(), + done: true, + error: None, + }) + .await + }) + .await?; + let client = server.connect_client().await?; + + // Phase 3: start a streaming request and abort the caller future while it + // is blocked waiting for response headers. + let client_for_request = client.clone(); + let stream_task = tokio::spawn(async move { + let _ = client_for_request + .http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/cancel".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }) + .await; + }); + request_seen_rx + .await + .expect("server should observe the first http/request"); + stream_task.abort(); + let _ = stream_task.await; + + // Phase 4: start a new stream immediately. It receives only the fresh body + // bytes for its generated id. + let (response, mut body_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/retry-cancelled".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }), + ) + .await + .context("second streamed http/request should return headers")??; + assert_eq!( + response, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + } + ); + let mut body = Vec::new(); + while let Some(chunk) = timeout(TEST_TIMEOUT, body_stream.recv()) + .await + .context("fresh http response body delta should arrive")?? + { + body.extend_from_slice(&chunk); + } + assert_eq!(body, b"fresh".to_vec()); + + drop(client); + server.finish().await?; + Ok(()) +} + +/// What this tests: dropping a returned body stream before EOF removes its +/// route and prevents stale body deltas from reaching the next stream. +#[tokio::test] +async fn http_response_body_stream_ignores_late_deltas_after_drop() -> Result<()> { + // Phase 1: script two requests. The first returns only headers; after the + // client drops its body receiver, the server sends a stale body delta. + let (body_dropped_tx, body_dropped_rx) = oneshot::channel(); + let (stale_delta_sent_tx, stale_delta_sent_rx) = oneshot::channel(); + let server = spawn_scripted_exec_server(|mut peer| async move { + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/drop".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-1".to_string(), + stream_response: true, + } + ); + peer.write_response( + request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await?; + body_dropped_rx + .await + .expect("test should drop the first body stream"); + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-1".to_string(), + seq: 1, + delta: b"stale".to_vec().into(), + done: false, + error: None, + }) + .await?; + stale_delta_sent_tx + .send(()) + .expect("test should wait for the stale delta"); + + // Phase 2: accept the next request with a new generated id. The new + // stream must receive only fresh body bytes. + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/retry-dropped".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-2".to_string(), + stream_response: true, + } + ); + peer.write_response( + request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await?; + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-2".to_string(), + seq: 1, + delta: b"fresh".to_vec().into(), + done: true, + error: None, + }) + .await + }) + .await?; + let client = server.connect_client().await?; + + // Phase 3: receive headers for the first stream, then drop the body stream + // without reading any body frames. + let (response, body_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/drop".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }), + ) + .await + .context("streamed http/request should return headers")??; + assert_eq!( + response, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + } + ); + drop(body_stream); + body_dropped_tx + .send(()) + .expect("server should wait for the body stream drop"); + stale_delta_sent_rx + .await + .expect("server should send one stale nonterminal delta"); + + // Phase 4: start the next stream immediately. The caller-provided id is + // ignored, and the fresh generated id isolates it from stale bytes. + let (reuse_response, mut reuse_body_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/retry-dropped".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }), + ) + .await + .context("second streamed http/request should return headers")??; + assert_eq!( + reuse_response, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + } + ); + let mut body = Vec::new(); + while let Some(chunk) = timeout(TEST_TIMEOUT, reuse_body_stream.recv()) + .await + .context("fresh http response body delta should arrive")?? + { + body.extend_from_slice(&chunk); + } + assert_eq!(body, b"fresh".to_vec()); + + drop(client); + server.finish().await?; + Ok(()) +} + +/// What this tests: an in-flight streamed HTTP body is failed when the shared +/// JSON-RPC transport disconnects before a terminal body frame. +#[tokio::test] +async fn http_response_body_stream_fails_when_transport_disconnects() -> Result<()> { + // Phase 1: return response headers for a streaming request, then drop the + // fake server transport without sending EOF. + let server = spawn_scripted_exec_server(|mut peer| async move { + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/disconnect".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-1".to_string(), + stream_response: true, + } + ); + peer.write_response( + request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await + }) + .await?; + let client = server.connect_client().await?; + + // Phase 2: start a streaming HTTP request and receive headers. + let (_response, mut body_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/disconnect".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }), + ) + .await + .context("streamed http/request should return headers")??; + + // Phase 3: assert transport disconnect wakes the body stream with a + // terminal error instead of hanging. + let error = timeout(TEST_TIMEOUT, body_stream.recv()) + .await + .context("disconnect should wake http body stream")? + .expect_err("disconnect should fail the http body stream"); + let error_message = error.to_string(); + assert_eq!( + error_message.starts_with( + "exec-server protocol error: http response stream `http-1` failed: exec-server transport disconnected" + ), + true + ); + + drop(client); + server.finish().await?; + Ok(()) +} + +/// What this tests: an executor cannot make the orchestrator decode and retain +/// a body frame larger than the response-stream wire contract allows. +#[tokio::test] +async fn http_response_body_stream_rejects_oversized_delta() -> Result<()> { + let (finish_tx, finish_rx) = oneshot::channel(); + let server = spawn_scripted_exec_server(|mut peer| async move { + let (_request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/oversized-delta".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-1".to_string(), + stream_response: true, + } + ); + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-1".to_string(), + seq: 1, + delta: vec![0; MAX_HTTP_BODY_DELTA_BYTES + 1].into(), + done: false, + error: None, + }) + .await?; + finish_rx.await.expect("test should finish server task"); + Ok(()) + }) + .await?; + let client = server.connect_client().await?; + + let request = HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/oversized-delta".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }; + let result = timeout(TEST_TIMEOUT, client.http_request_stream(request)) + .await + .context("oversized body delta should close the executor transport")?; + let error = match result { + Ok(_) => bail!("oversized body delta should fail the request"), + Err(error) => error, + }; + let error = error.to_string(); + assert_eq!(error, "exec-server transport disconnected"); + + finish_tx.send(()).expect("server task should stay active"); + drop(client); + server.finish().await?; + Ok(()) +} + +/// What this tests: frame-count backpressure cannot hide an unbounded amount +/// of executor-controlled body bytes across the orchestrator's stream queues. +#[tokio::test] +async fn http_response_body_stream_enforces_queued_byte_budget() -> Result<()> { + let (finish_tx, finish_rx) = oneshot::channel(); + let server = spawn_scripted_exec_server(|mut peer| async move { + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/byte-budget".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-1".to_string(), + stream_response: true, + } + ); + peer.write_response( + request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await?; + + let frame_count = HTTP_BODY_DELTA_BYTE_BUDGET / MAX_HTTP_BODY_DELTA_BYTES + 1; + for seq in 1..=frame_count as u64 { + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-1".to_string(), + seq, + delta: vec![0; MAX_HTTP_BODY_DELTA_BYTES].into(), + done: false, + error: None, + }) + .await?; + } + + let (barrier_request_id, barrier_params) = peer.read_http_request().await?; + assert_eq!( + barrier_params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/byte-budget-barrier".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-2".to_string(), + stream_response: true, + } + ); + peer.write_response( + barrier_request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await?; + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-2".to_string(), + seq: 1, + delta: Vec::new().into(), + done: true, + error: None, + }) + .await?; + finish_rx.await.expect("test should finish server task"); + Ok(()) + }) + .await?; + let client = server.connect_client().await?; + + let (_response, mut body_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/byte-budget".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }), + ) + .await + .context("streamed http/request should return headers")??; + + // Receiving this terminal notification proves the earlier byte-budget + // notifications have all passed through the ordered notification handler. + let (_response, mut barrier_stream) = timeout( + BYTE_BUDGET_TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/byte-budget-barrier".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-barrier-id".to_string(), + stream_response: false, + }), + ) + .await + .context("barrier http/request should return headers")??; + assert_eq!( + timeout(TEST_TIMEOUT, barrier_stream.recv()) + .await + .context("barrier body stream should finish")??, + None + ); + + let mut delivered_bytes = 0; + let error = loop { + match timeout(TEST_TIMEOUT, body_stream.recv()) + .await + .context("queued body stream should finish")? + { + Ok(Some(chunk)) => delivered_bytes += chunk.len(), + Ok(None) => bail!("byte-budget exhaustion should not look like clean EOF"), + Err(error) => break error, + } + }; + assert_eq!(delivered_bytes, HTTP_BODY_DELTA_BYTE_BUDGET); + assert!( + error + .to_string() + .contains("queued body deltas exceed 16777216 bytes") + ); + + finish_tx.send(()).expect("server task should stay active"); + drop(client); + server.finish().await?; + Ok(()) +} + +/// What this tests: every response stream on one executor connection shares +/// the same queued-body byte budget. +#[tokio::test] +async fn http_response_body_streams_share_queued_byte_budget() -> Result<()> { + let (finish_tx, finish_rx) = oneshot::channel(); + let server = spawn_scripted_exec_server(|mut peer| async move { + for (request_id, url) in [ + ("http-1", "https://example.test/mcp/shared-budget-one"), + ("http-2", "https://example.test/mcp/shared-budget-two"), + ] { + let (rpc_request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: url.to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: request_id.to_string(), + stream_response: true, + } + ); + peer.write_response( + rpc_request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await?; + } + + let frames_per_stream = HTTP_BODY_DELTA_BYTE_BUDGET / MAX_HTTP_BODY_DELTA_BYTES / 2; + for request_id in ["http-1", "http-2"] { + for seq in 1..=frames_per_stream as u64 { + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: request_id.to_string(), + seq, + delta: vec![0; MAX_HTTP_BODY_DELTA_BYTES].into(), + done: false, + error: None, + }) + .await?; + } + } + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-2".to_string(), + seq: frames_per_stream as u64 + 1, + delta: vec![0; MAX_HTTP_BODY_DELTA_BYTES].into(), + done: false, + error: None, + }) + .await?; + + let (barrier_request_id, barrier_params) = peer.read_http_request().await?; + assert_eq!( + barrier_params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/shared-budget-barrier".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-3".to_string(), + stream_response: true, + } + ); + peer.write_response( + barrier_request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await?; + for (request_id, seq) in [ + ("http-3", 1), + ("http-1", frames_per_stream as u64 + 1), + ("http-2", frames_per_stream as u64 + 2), + ] { + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: request_id.to_string(), + seq, + delta: Vec::new().into(), + done: true, + error: None, + }) + .await?; + } + finish_rx.await.expect("test should finish server task"); + Ok(()) + }) + .await?; + let client = server.connect_client().await?; + + let (_response, mut first_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/shared-budget-one".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-one".to_string(), + stream_response: false, + }), + ) + .await + .context("first streamed http/request should return headers")??; + let (_response, mut second_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/shared-budget-two".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-two".to_string(), + stream_response: false, + }), + ) + .await + .context("second streamed http/request should return headers")??; + + // This terminal notification is ordered after both streams contend for the + // budget, so neither stream is drained before the overflow is observed. + let (_response, mut barrier_stream) = timeout( + BYTE_BUDGET_TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/shared-budget-barrier".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-barrier-id".to_string(), + stream_response: false, + }), + ) + .await + .context("barrier http/request should return headers")??; + assert_eq!( + timeout(TEST_TIMEOUT, barrier_stream.recv()) + .await + .context("barrier body stream should finish")??, + None + ); + + let mut failed_stream_bytes = 0; + let error = loop { + match timeout(TEST_TIMEOUT, second_stream.recv()) + .await + .context("second body stream should finish")? + { + Ok(Some(chunk)) => failed_stream_bytes += chunk.len(), + Ok(None) => bail!("shared byte-budget exhaustion should not look like clean EOF"), + Err(error) => break error, + } + }; + assert_eq!( + (failed_stream_bytes, error.to_string()), + ( + HTTP_BODY_DELTA_BYTE_BUDGET / 2, + "exec-server protocol error: http response stream `http-2` failed: queued body deltas exceed 16777216 bytes".to_string(), + ) + ); + + let mut surviving_stream_bytes = 0; + while let Some(chunk) = timeout(TEST_TIMEOUT, first_stream.recv()) + .await + .context("first body stream should finish")?? + { + surviving_stream_bytes += chunk.len(); + } + assert_eq!(surviving_stream_bytes, HTTP_BODY_DELTA_BYTE_BUDGET / 2); + + finish_tx.send(()).expect("server task should stay active"); + drop(client); + server.finish().await?; + Ok(()) +} + +/// What this tests: transport disconnect still records a terminal stream +/// failure even when the client-side body-delta queue is already full. +#[tokio::test] +async fn http_response_body_stream_reports_disconnect_when_queue_is_full() -> Result<()> { + // Phase 1: fill the queued body-delta route exactly to capacity before the + // response headers arrive, then drop the transport without sending EOF. + let server = spawn_scripted_exec_server(|mut peer| async move { + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/disconnect-full-queue".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-1".to_string(), + stream_response: true, + } + ); + for seq in 1..=HTTP_BODY_DELTA_CHANNEL_CAPACITY { + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-1".to_string(), + seq, + delta: b"x".to_vec().into(), + done: false, + error: None, + }) + .await?; + } + peer.write_response( + request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await + }) + .await?; + let client = server.connect_client().await?; + + // Phase 2: start the streaming request and receive headers while the + // queue is already full. + let (_response, mut body_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/disconnect-full-queue".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }), + ) + .await + .context("streamed http/request should return headers")??; + + // Phase 3: drain the queued chunks and assert the transport disconnect is + // still reported as an error rather than a clean EOF. + let mut chunks = 0; + let error = loop { + match timeout(TEST_TIMEOUT, body_stream.recv()) + .await + .context("disconnect should wake the full queued body stream")? + { + Ok(Some(_chunk)) => { + chunks += 1; + } + Ok(None) => bail!("disconnect with a full queue should not look like clean EOF"), + Err(error) => break error, + } + }; + assert_eq!( + ( + chunks, + error + .to_string() + .starts_with( + "exec-server protocol error: http response stream `http-1` failed: exec-server transport disconnected", + ), + ), + (HTTP_BODY_DELTA_CHANNEL_CAPACITY as usize, true) + ); + + drop(client); + server.finish().await?; + Ok(()) +} + +/// What this tests: body-delta backpressure closes the public body stream as +/// an error rather than letting callers accept a truncated body as clean EOF. +#[tokio::test] +async fn http_response_body_stream_reports_backpressure_truncation() -> Result<()> { + // Phase 1: send enough body frames before headers to overflow the bounded + // client-side route while the public request future is still pending. + let (finish_tx, finish_rx) = oneshot::channel(); + let server = spawn_scripted_exec_server(|mut peer| async move { + let (request_id, params) = peer.read_http_request().await?; + assert_eq!( + params, + HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/backpressure".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "http-1".to_string(), + stream_response: true, + } + ); + for seq in 1..=OVERFLOWING_BODY_DELTA_FRAMES { + peer.write_body_delta(HttpRequestBodyDeltaNotification { + request_id: "http-1".to_string(), + seq, + delta: b"x".to_vec().into(), + done: false, + error: None, + }) + .await?; + } + peer.write_response( + request_id, + HttpRequestResponse { + status: 200, + headers: Vec::new(), + body: Vec::new().into(), + }, + ) + .await?; + + // Phase 2: keep the transport connected so the body stream reports the + // backpressure failure rather than a disconnect. + finish_rx.await.expect("test should finish server task"); + Ok(()) + }) + .await?; + let client = server.connect_client().await?; + + // Phase 3: start the streaming request; the server overfills the route + // before returning the body stream to this consumer. + let (_response, mut body_stream) = timeout( + TEST_TIMEOUT, + client.http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: "https://example.test/mcp/backpressure".to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "caller-stream-id".to_string(), + stream_response: false, + }), + ) + .await + .context("streamed http/request should return headers")??; + + // Phase 4: drain queued chunks and assert the truncated stream ends in an + // explicit error, not a clean EOF. + let mut chunks = 0; + let error = loop { + match timeout(TEST_TIMEOUT, body_stream.recv()) + .await + .context("backpressure should close http body stream")? + { + Ok(Some(_chunk)) => { + chunks += 1; + } + Ok(None) => bail!("backpressure truncation should not look like clean EOF"), + Err(error) => break error, + } + }; + assert_eq!( + ( + chunks < OVERFLOWING_BODY_DELTA_FRAMES as usize, + error.to_string(), + ), + ( + true, + "exec-server protocol error: http response stream `http-1` failed: body delta channel filled before delivery".to_string(), + ) + ); + + finish_tx + .send(()) + .expect("server task should wait for test completion"); + drop(client); + server.finish().await?; + Ok(()) +} + +/// Fake WebSocket exec-server used by the integration tests. +/// +/// The helper exercises `ExecServerClient::connect_websocket`, including the +/// initialize handshake, while each test controls the exact JSON-RPC traffic +/// that follows. +struct ScriptedExecServer { + websocket_url: String, + task: JoinHandle>, +} + +impl ScriptedExecServer { + /// Connects the public exec-server client to this fake WebSocket endpoint. + async fn connect_client(&self) -> Result { + ExecServerClient::connect_websocket(RemoteExecServerConnectArgs::new( + self.websocket_url.clone(), + CLIENT_NAME.to_string(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )) + .await + .context("client should connect to fake exec-server") + } + + /// Waits for the scripted fake server to finish. + async fn finish(self) -> Result<()> { + self.task + .await + .context("fake exec-server task should join")??; + Ok(()) + } +} + +/// Starts a fake exec-server that accepts one WebSocket client. +async fn spawn_scripted_exec_server(script: F) -> Result +where + F: FnOnce(JsonRpcPeer) -> Fut + Send + 'static, + Fut: Future> + Send + 'static, +{ + let listener = TcpListener::bind("127.0.0.1:0") + .await + .context("fake exec-server should bind")?; + let websocket_url = format!("ws://{}", listener.local_addr()?); + let task = tokio::spawn(async move { + let (stream, _) = timeout(TEST_TIMEOUT, listener.accept()) + .await + .context("fake exec-server should accept a client")??; + let websocket = accept_async(stream) + .await + .context("fake exec-server websocket handshake should complete")?; + let mut peer = JsonRpcPeer { websocket }; + peer.complete_initialize().await?; + script(peer).await + }); + Ok(ScriptedExecServer { + websocket_url, + task, + }) +} + +/// JSON-RPC peer for the fake exec-server WebSocket. +struct JsonRpcPeer { + websocket: WebSocketStream, +} + +impl JsonRpcPeer { + /// Completes and validates the client initialize handshake. + async fn complete_initialize(&mut self) -> Result<()> { + let request = self.read_request(INITIALIZE_METHOD).await?; + let params: InitializeParams = decode_request_params(&request)?; + assert_eq!( + params, + InitializeParams { + client_name: CLIENT_NAME.to_string(), + resume_session_id: None, + } + ); + self.write_response( + request.id, + InitializeResponse { + session_id: "session-1".to_string(), + environment_info: None, + }, + ) + .await?; + self.read_notification(INITIALIZED_METHOD).await?; + Ok(()) + } + + /// Reads one typed `http/request` call from the client. + async fn read_http_request(&mut self) -> Result<(RequestId, HttpRequestParams)> { + let request = self.read_request(HTTP_REQUEST_METHOD).await?; + let params = decode_request_params(&request)?; + Ok((request.id, params)) + } + + /// Reads a JSON-RPC request and validates its method. + async fn read_request(&mut self, expected_method: &str) -> Result { + let message = self.read_message().await?; + let JSONRPCMessage::Request(request) = message else { + bail!("expected JSON-RPC request `{expected_method}`, got {message:?}"); + }; + if request.method != expected_method { + bail!( + "expected JSON-RPC request `{expected_method}`, got `{}`", + request.method + ); + } + Ok(request) + } + + /// Reads a JSON-RPC notification and validates its method. + async fn read_notification(&mut self, expected_method: &str) -> Result { + let message = self.read_message().await?; + let JSONRPCMessage::Notification(notification) = message else { + bail!("expected JSON-RPC notification `{expected_method}`, got {message:?}"); + }; + if notification.method != expected_method { + bail!( + "expected JSON-RPC notification `{expected_method}`, got `{}`", + notification.method + ); + } + Ok(notification) + } + + /// Sends a successful JSON-RPC response. + async fn write_response(&mut self, id: RequestId, result: T) -> Result<()> + where + T: Serialize, + { + self.write_message(JSONRPCMessage::Response(JSONRPCResponse { + id, + result: to_value(result)?, + })) + .await + } + + /// Sends one streamed HTTP body notification. + async fn write_body_delta(&mut self, delta: HttpRequestBodyDeltaNotification) -> Result<()> { + self.write_message(JSONRPCMessage::Notification(JSONRPCNotification { + method: HTTP_REQUEST_BODY_DELTA_METHOD.to_string(), + params: Some(to_value(delta)?), + })) + .await + } + + /// Reads one WebSocket JSON-RPC message. + async fn read_message(&mut self) -> Result { + let message = timeout(TEST_TIMEOUT, self.websocket.next()) + .await + .context("timed out waiting for JSON-RPC message")? + .context("client websocket closed before JSON-RPC message arrived")? + .context("failed to read websocket message")?; + match message { + Message::Text(text) => from_str(text.as_ref()).context("text JSON-RPC"), + Message::Binary(bytes) => from_slice(bytes.as_ref()).context("binary JSON-RPC"), + Message::Close(frame) => bail!("client websocket closed: {frame:?}"), + other => bail!("expected text or binary JSON-RPC message, got {other:?}"), + } + } + + /// Writes one WebSocket JSON-RPC message. + async fn write_message(&mut self, message: JSONRPCMessage) -> Result<()> { + let encoded = to_string(&message)?; + timeout( + TEST_TIMEOUT, + self.websocket.send(Message::Text(encoded.into())), + ) + .await + .context("timed out writing JSON-RPC message")? + .context("failed to write JSON-RPC message") + } +} + +/// Decodes a request params object into its typed protocol payload. +fn decode_request_params(request: &JSONRPCRequest) -> Result +where + T: DeserializeOwned, +{ + let params = request + .params + .clone() + .context("JSON-RPC request should include params")?; + from_value(params).context("JSON-RPC request params should decode") +} diff --git a/codex-rs/exec-server/tests/http_request.rs b/codex-rs/exec-server/tests/http_request.rs new file mode 100644 index 0000000000000000000000000000000000000000..55144de11eca27bf9f4e6654162dba43c6c429b1 --- /dev/null +++ b/codex-rs/exec-server/tests/http_request.rs @@ -0,0 +1,953 @@ +#![cfg(unix)] + +mod common; + +use std::collections::BTreeMap; +use std::io::ErrorKind; +use std::time::Duration; + +use codex_exec_server::HttpHeader; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestBodyDeltaNotification; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::InitializeParams; +use codex_exec_server_protocol::JSONRPCError; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCNotification; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::RequestId; +use common::SYSTEM_PROXY_REQUEST_URL_ENV; +use common::SYSTEM_PROXY_URL_ENV; +use common::exec_server::ExecServerHarness; +use common::exec_server::exec_server; +use common::exec_server::exec_server_with_env; +use pretty_assertions::assert_eq; +use serde::de::DeserializeOwned; +use serde_json::Value; +use tokio::io::AsyncBufReadExt; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::io::BufReader; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::sync::oneshot; +use tokio::time::timeout; + +/// HTTP request captured by the ad-hoc TCP server in these integration tests. +#[derive(Debug)] +struct CapturedHttpRequest { + stream: TcpStream, + request_line: String, + headers: BTreeMap, + body: Vec, +} + +/// What this tests: a real exec-server websocket `http/request` performs one +/// HTTP request through the runner and returns the complete response body in +/// the JSON-RPC response. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_http_request_buffers_response_body() -> anyhow::Result<()> { + // Phase 1: start exec-server and complete the JSON-RPC handshake. + let mut server = exec_server_with_env( + [ + ("NODE_REPL_AUTH_TOKEN", "executor-token"), + ("LINEAR_API_KEY", "linear-token"), + ], + &[], + ) + .await?; + initialize_exec_server(&mut server).await?; + + // Phase 2: start a local HTTP peer and ask exec-server to POST to it. + let listener = TcpListener::bind("127.0.0.1:0").await?; + let url = format!("http://{}/mcp?case=buffered", listener.local_addr()?); + let http_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "POST".to_string(), + url, + headers: vec![ + HttpHeader { + name: "x-codex-test".to_string(), + value: "buffered".to_string(), + value_env_var: None, + }, + HttpHeader { + name: "authorization".to_string(), + value: "Bearer ".to_string(), + value_env_var: Some("NODE_REPL_AUTH_TOKEN".to_string()), + }, + HttpHeader { + name: "x-linear-authorization".to_string(), + value: "Bearer ".to_string(), + value_env_var: Some("LINEAR_API_KEY".to_string()), + }, + ], + body: Some(b"request-body".to_vec().into()), + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "buffered-request".to_string(), + stream_response: false, + })?, + ) + .await?; + + // Phase 3: assert the HTTP peer observes the expected method, path, + // headers, and body before returning a fixed-length response. + let captured = accept_http_request(&listener).await?; + assert_eq!( + ( + captured.request_line.as_str(), + captured.headers.get("x-codex-test").map(String::as_str), + captured.headers.get("authorization").map(String::as_str), + captured + .headers + .get("x-linear-authorization") + .map(String::as_str), + captured.body.as_slice(), + ), + ( + "POST /mcp?case=buffered HTTP/1.1", + Some("buffered"), + Some("Bearer executor-token"), + Some("Bearer linear-token"), + b"request-body".as_slice(), + ) + ); + respond_with_status_and_headers( + captured.stream, + "201 Created", + &[("x-mcp-test", "buffered")], + b"response-body", + ) + .await?; + + // Phase 4: assert exec-server returns status, response headers, and the + // full response body in the JSON-RPC result. + let response: HttpRequestResponse = wait_for_response(&mut server, http_request_id).await?; + assert_eq!( + ( + response.status, + response_header(&response.headers, "x-mcp-test"), + response.body.into_inner(), + ), + (201, Some("buffered".to_string()), b"response-body".to_vec(),) + ); + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_http_request_rejects_protected_environment_headers() -> anyhow::Result<()> { + let mut server = exec_server_with_env( + [( + "CODEX_EXEC_SERVER_NOISE_AUTH_TOKEN", + "executor-internal-token", + )], + &[], + ) + .await?; + initialize_exec_server(&mut server).await?; + + let listener = TcpListener::bind("127.0.0.1:0").await?; + for (index, env_var) in [ + "CODEX_EXEC_SERVER_NOISE_AUTH_TOKEN", + "codex_exec_server_noise_auth_token", + "OPENAI_API_KEY", + "CODEX_ACCESS_TOKEN", + "CODEX_CONNECTORS_TOKEN", + "AWS_SECRET_ACCESS_KEY", + "AZURE_FEDERATED_TOKEN_FILE", + "OPENAI_IDENTITY_TOKEN_FILE", + ] + .into_iter() + .enumerate() + { + let request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url: format!("http://{}/mcp", listener.local_addr()?), + headers: vec![HttpHeader { + name: "authorization".to_string(), + value: "Bearer ".to_string(), + value_env_var: Some(env_var.to_string()), + }], + body: None, + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: format!("protected-header-request-{index}"), + stream_response: false, + })?, + ) + .await?; + let error = wait_for_error_response(&mut server, request_id).await?; + assert_eq!(error.code, -32602); + assert_eq!( + error.message, + format!( + "http/request header authorization cannot use executor environment variable {env_var}" + ) + ); + } + assert!( + timeout(Duration::from_millis(50), listener.accept()) + .await + .is_err() + ); + + server.shutdown().await?; + Ok(()) +} + +/// What this tests: delegated HTTP accepts URLs with client-side fragments and +/// does not include those fragments in the request sent over the network. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_http_request_omits_url_fragment() -> anyhow::Result<()> { + let mut server = exec_server().await?; + initialize_exec_server(&mut server).await?; + + let listener = TcpListener::bind("127.0.0.1:0").await?; + let url = format!( + "http://{}/mcp?case=fragment#client-section", + listener.local_addr()? + ); + let http_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url, + headers: Vec::new(), + body: None, + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "fragment-request".to_string(), + stream_response: false, + })?, + ) + .await?; + + let captured = accept_http_request(&listener).await?; + assert_eq!(captured.request_line, "GET /mcp?case=fragment HTTP/1.1"); + respond_with_status_and_headers(captured.stream, "200 OK", &[], b"fragment-response").await?; + + let response: HttpRequestResponse = wait_for_response(&mut server, http_request_id).await?; + assert_eq!( + (response.status, response.body.into_inner()), + (200, b"fragment-response".to_vec()) + ); + + server.shutdown().await?; + Ok(()) +} + +/// What this tests: a configured system-proxy factory survives the complete +/// executor transport, processor, and handler chain for delegated HTTP. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_http_request_uses_configured_system_proxy() -> anyhow::Result<()> { + let proxy_listener = TcpListener::bind("127.0.0.1:0").await?; + let proxy_url = format!("http://{}", proxy_listener.local_addr()?); + let request_url = "http://exec-server-system-proxy.invalid/delegated?route=system"; + let mut server = exec_server_with_env( + [ + (SYSTEM_PROXY_REQUEST_URL_ENV, request_url), + (SYSTEM_PROXY_URL_ENV, proxy_url.as_str()), + ("HTTP_PROXY", ""), + ("http_proxy", ""), + ("HTTPS_PROXY", ""), + ("https_proxy", ""), + ("ALL_PROXY", ""), + ("all_proxy", ""), + ("NO_PROXY", ""), + ("no_proxy", ""), + ], + &[], + ) + .await?; + initialize_exec_server(&mut server).await?; + + let http_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url: request_url.to_string(), + headers: Vec::new(), + body: None, + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "system-proxy-request".to_string(), + stream_response: false, + })?, + ) + .await?; + + let captured = accept_http_request(&proxy_listener).await?; + assert_eq!( + captured.request_line, + "GET http://exec-server-system-proxy.invalid/delegated?route=system HTTP/1.1" + ); + respond_with_status_and_headers(captured.stream, "200 OK", &[], b"proxied-response").await?; + + let response: HttpRequestResponse = wait_for_response(&mut server, http_request_id).await?; + assert_eq!( + (response.status, response.body.into_inner()), + (200, b"proxied-response".to_vec()) + ); + + server.shutdown().await?; + Ok(()) +} + +/// What this tests: delegated HTTP preserves WHATWG IDNA normalization before +/// selecting a configured system-proxy route. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_http_request_normalizes_unicode_hostname() -> anyhow::Result<()> { + let proxy_listener = TcpListener::bind("127.0.0.1:0").await?; + let proxy_url = format!("http://{}", proxy_listener.local_addr()?); + let request_url = "http://münich.invalid/mcp?route=unicode"; + let normalized_url = "http://xn--mnich-kva.invalid/mcp?route=unicode"; + let mut server = exec_server_with_env( + [ + (SYSTEM_PROXY_REQUEST_URL_ENV, normalized_url), + (SYSTEM_PROXY_URL_ENV, proxy_url.as_str()), + ("HTTP_PROXY", ""), + ("http_proxy", ""), + ("HTTPS_PROXY", ""), + ("https_proxy", ""), + ("ALL_PROXY", ""), + ("all_proxy", ""), + ("NO_PROXY", ""), + ("no_proxy", ""), + ], + &[], + ) + .await?; + initialize_exec_server(&mut server).await?; + + let http_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url: request_url.to_string(), + headers: Vec::new(), + body: None, + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "unicode-hostname-request".to_string(), + stream_response: false, + })?, + ) + .await?; + + let captured = accept_http_request(&proxy_listener).await?; + assert_eq!( + ( + captured.request_line.as_str(), + captured.headers.get("host").map(String::as_str), + ), + ( + "GET http://xn--mnich-kva.invalid/mcp?route=unicode HTTP/1.1", + Some("xn--mnich-kva.invalid"), + ) + ); + respond_with_status_and_headers(captured.stream, "200 OK", &[], b"unicode-response").await?; + + let response: HttpRequestResponse = wait_for_response(&mut server, http_request_id).await?; + assert_eq!( + (response.status, response.body.into_inner()), + (200, b"unicode-response".to_vec()) + ); + + server.shutdown().await?; + Ok(()) +} + +/// What this tests: OAuth callers can inspect redirect responses without the +/// executor following the Location header. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_http_request_can_stop_at_redirects() -> anyhow::Result<()> { + let mut server = exec_server().await?; + initialize_exec_server(&mut server).await?; + + let listener = TcpListener::bind("127.0.0.1:0").await?; + let base_url = format!("http://{}", listener.local_addr()?); + let http_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url: format!("{base_url}/redirect"), + headers: Vec::new(), + body: None, + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Stop, + request_id: "redirect-request".to_string(), + stream_response: false, + })?, + ) + .await?; + + let captured = accept_http_request(&listener).await?; + assert_eq!(captured.request_line, "GET /redirect HTTP/1.1"); + respond_with_status_and_headers( + captured.stream, + "302 Found", + &[("location", &format!("{base_url}/final"))], + b"redirect", + ) + .await?; + + let response: HttpRequestResponse = wait_for_response(&mut server, http_request_id).await?; + assert_eq!( + ( + response.status, + response_header(&response.headers, "location"), + response.body.into_inner(), + ), + (302, Some(format!("{base_url}/final")), b"redirect".to_vec(),) + ); + assert!( + timeout(Duration::from_millis(100), listener.accept()) + .await + .is_err(), + "redirect target should not be requested" + ); + + server.shutdown().await?; + Ok(()) +} + +/// What this tests: the executor follows redirects when the HTTP request +/// explicitly selects the redirect-following client pool. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_http_request_can_follow_redirects() -> anyhow::Result<()> { + let mut server = exec_server().await?; + initialize_exec_server(&mut server).await?; + + let listener = TcpListener::bind("127.0.0.1:0").await?; + let base_url = format!("http://{}", listener.local_addr()?); + let http_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url: format!("{base_url}/redirect"), + headers: Vec::new(), + body: None, + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "follow-redirect-request".to_string(), + stream_response: false, + })?, + ) + .await?; + + let redirect_request = accept_http_request(&listener).await?; + assert_eq!(redirect_request.request_line, "GET /redirect HTTP/1.1"); + respond_with_status_and_headers( + redirect_request.stream, + "302 Found", + &[("location", &format!("{base_url}/final"))], + b"redirect", + ) + .await?; + + let final_request = accept_http_request(&listener).await?; + assert_eq!(final_request.request_line, "GET /final HTTP/1.1"); + respond_with_status_and_headers( + final_request.stream, + "200 OK", + &[("x-mcp-test", "redirected")], + b"final-response-body", + ) + .await?; + + let response: HttpRequestResponse = wait_for_response(&mut server, http_request_id).await?; + assert_eq!( + ( + response.status, + response_header(&response.headers, "x-mcp-test"), + response.body.into_inner(), + ), + ( + 200, + Some("redirected".to_string()), + b"final-response-body".to_vec(), + ) + ); + + server.shutdown().await?; + Ok(()) +} + +/// What this tests: a real exec-server websocket `http/request` can return +/// response headers immediately and stream the response body as ordered +/// `http/request/bodyDelta` notifications. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_http_request_streams_response_body_notifications() -> anyhow::Result<()> { + // Phase 1: start exec-server and complete the JSON-RPC handshake. + let mut server = exec_server().await?; + initialize_exec_server(&mut server).await?; + + // Phase 2: start a local HTTP peer and ask exec-server for a streamed GET. + let listener = TcpListener::bind("127.0.0.1:0").await?; + let url = format!("http://{}/mcp?case=streaming", listener.local_addr()?); + let http_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url, + headers: vec![HttpHeader { + name: "accept".to_string(), + value: "text/event-stream".to_string(), + value_env_var: None, + }], + body: None, + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "stream-1".to_string(), + stream_response: true, + })?, + ) + .await?; + + // Phase 3: assert the HTTP peer observes the expected request and then + // respond with chunked transfer encoding to exercise streaming. + let captured = accept_http_request(&listener).await?; + assert_eq!( + ( + captured.request_line.as_str(), + captured.headers.get("accept").map(String::as_str), + captured.body, + ), + ( + "GET /mcp?case=streaming HTTP/1.1", + Some("text/event-stream"), + Vec::new(), + ) + ); + respond_with_chunked_body( + captured.stream, + &[("x-mcp-test", "streaming")], + &[b"hello ".as_slice(), b"world".as_slice()], + ) + .await?; + + // Phase 4: assert the JSON-RPC response reaches the wire before any body + // delta notifications, and that it contains status and headers but no + // buffered body when streaming is requested. + let first_event = server.next_event().await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = first_event else { + anyhow::bail!("expected http/request response before body deltas, got {first_event:?}"); + }; + assert_eq!(id, http_request_id); + let response: HttpRequestResponse = serde_json::from_value(result)?; + assert_eq!( + ( + response.status, + response_header(&response.headers, "x-mcp-test"), + response.body.into_inner(), + ), + (200, Some("streaming".to_string()), Vec::new()) + ); + + // Phase 5: assert the body notifications are contiguous, ordered, and end + // with a clean terminal frame. + let deltas = collect_response_body_deltas(&mut server, "stream-1").await?; + let seqs = deltas.iter().map(|delta| delta.seq).collect::>(); + let body = deltas + .iter() + .flat_map(|delta| delta.delta.clone().into_inner()) + .collect::>(); + let terminal = deltas.last().map(|delta| (delta.done, delta.error.clone())); + let expected_seqs = (1..=deltas.len() as u64).collect::>(); + assert_eq!( + (seqs, body, terminal), + (expected_seqs, b"hello world".to_vec(), Some((true, None))) + ); + + server.shutdown().await?; + Ok(()) +} + +/// What this tests: streamed `requestId`s stay reserved until the body stream +/// finishes, so a second in-flight request cannot reuse the same id. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_http_request_rejects_duplicate_stream_request_ids() -> anyhow::Result<()> { + let mut server = exec_server().await?; + initialize_exec_server(&mut server).await?; + + let listener = TcpListener::bind("127.0.0.1:0").await?; + let url = format!( + "http://{}/mcp?case=duplicate-stream-id", + listener.local_addr()? + ); + let first_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url: url.clone(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "stream-dup".to_string(), + stream_response: true, + })?, + ) + .await?; + + let captured = accept_http_request(&listener).await?; + let (finish_tx, finish_rx) = oneshot::channel(); + let response_task = tokio::spawn(async move { + respond_with_chunked_body_until_finish(captured.stream, &[], &[b"hello"], finish_rx).await + }); + + let _: HttpRequestResponse = wait_for_response(&mut server, first_request_id).await?; + + let duplicate_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url, + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "stream-dup".to_string(), + stream_response: true, + })?, + ) + .await?; + + let duplicate_response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Error(JSONRPCError { id, .. }) if id == &duplicate_request_id + ) + }) + .await?; + let JSONRPCMessage::Error(JSONRPCError { error, .. }) = duplicate_response else { + anyhow::bail!("expected duplicate requestId error response"); + }; + assert_eq!(error.code, -32602); + assert_eq!( + error.message, + "http/request streamResponse requestId `stream-dup` is already active" + ); + + finish_tx + .send(()) + .expect("response task should still be waiting"); + response_task.await??; + let _ = collect_response_body_deltas(&mut server, "stream-dup").await?; + + server.shutdown().await?; + Ok(()) +} + +/// What this tests: omitting `timeoutMs` leaves the request unbounded, while +/// an explicit short timeout still fails the same delayed response. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_http_request_honors_optional_timeout() -> anyhow::Result<()> { + let mut server = exec_server().await?; + initialize_exec_server(&mut server).await?; + + let listener = TcpListener::bind("127.0.0.1:0").await?; + let delayed_url = format!( + "http://{}/mcp?case=optional-timeout", + listener.local_addr()? + ); + let no_timeout_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url: delayed_url.clone(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "buffered-request".to_string(), + stream_response: false, + })?, + ) + .await?; + + let captured = accept_http_request(&listener).await?; + let delayed_response = tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(100)).await; + respond_with_status_and_headers(captured.stream, "200 OK", &[], b"slow-success").await + }); + let response: HttpRequestResponse = + wait_for_response(&mut server, no_timeout_request_id).await?; + assert_eq!(response.body.into_inner(), b"slow-success".to_vec()); + delayed_response.await??; + + let timeout_request_id = server + .send_request( + "http/request", + serde_json::to_value(HttpRequestParams { + method: "GET".to_string(), + url: delayed_url, + headers: Vec::new(), + body: None, + timeout_ms: Some(10), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "buffered-request".to_string(), + stream_response: false, + })?, + ) + .await?; + + let captured = accept_http_request(&listener).await?; + let delayed_timeout_response = tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(100)).await; + respond_with_status_and_headers(captured.stream, "200 OK", &[], b"too-late").await + }); + let error = wait_for_error_response(&mut server, timeout_request_id).await?; + assert_eq!(error.code, -32603); + assert!( + error.message.starts_with("http/request failed: "), + "unexpected timeout error: {}", + error.message + ); + match delayed_timeout_response.await? { + Ok(()) => {} + Err(err) if is_expected_peer_disconnect(&err) => {} + Err(err) => return Err(err), + } + + server.shutdown().await?; + Ok(()) +} + +/// Performs the JSON-RPC initialize handshake required before executor methods. +async fn initialize_exec_server(server: &mut ExecServerHarness) -> anyhow::Result<()> { + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-http-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + let _: Value = wait_for_response(server, initialize_id).await?; + server + .send_notification("initialized", serde_json::json!({})) + .await?; + Ok(()) +} + +/// Waits for a typed JSON-RPC response with the requested id. +async fn wait_for_response( + server: &mut ExecServerHarness, + request_id: RequestId, +) -> anyhow::Result +where + T: DeserializeOwned, +{ + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &request_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { result, .. }) = response else { + anyhow::bail!("expected JSON-RPC response for {request_id:?}"); + }; + Ok(serde_json::from_value(result)?) +} + +/// Waits for a JSON-RPC error with the requested id. +async fn wait_for_error_response( + server: &mut ExecServerHarness, + request_id: RequestId, +) -> anyhow::Result { + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Error(JSONRPCError { id, .. }) if id == &request_id + ) + }) + .await?; + let JSONRPCMessage::Error(JSONRPCError { error, .. }) = response else { + anyhow::bail!("expected JSON-RPC error for {request_id:?}"); + }; + Ok(error) +} + +/// Accepts one HTTP/1.1 request and captures its wire-visible fields. +async fn accept_http_request(listener: &TcpListener) -> anyhow::Result { + let (stream, _) = timeout(Duration::from_secs(5), listener.accept()).await??; + let mut reader = BufReader::new(stream); + + let mut request_line = String::new(); + reader.read_line(&mut request_line).await?; + let request_line = request_line.trim_end_matches("\r\n").to_string(); + + let mut headers = BTreeMap::new(); + loop { + let mut line = String::new(); + reader.read_line(&mut line).await?; + if line == "\r\n" { + break; + } + let line = line.trim_end_matches("\r\n"); + let (name, value) = line + .split_once(':') + .ok_or_else(|| anyhow::anyhow!("HTTP header should contain colon: {line}"))?; + headers.insert(name.to_ascii_lowercase(), value.trim().to_string()); + } + + let content_length = headers + .get("content-length") + .and_then(|value| value.parse::().ok()) + .unwrap_or(0); + let mut body = vec![0; content_length]; + reader.read_exact(&mut body).await?; + + Ok(CapturedHttpRequest { + stream: reader.into_inner(), + request_line, + headers, + body, + }) +} + +/// Writes a fixed-length HTTP response to the captured request stream. +async fn respond_with_status_and_headers( + mut stream: TcpStream, + status: &str, + headers: &[(&str, &str)], + body: &[u8], +) -> anyhow::Result<()> { + let extra_headers = headers + .iter() + .map(|(name, value)| format!("{name}: {value}\r\n")) + .collect::(); + let response = format!( + "HTTP/1.1 {status}\r\ncontent-type: text/plain\r\ncontent-length: {}\r\nconnection: close\r\n{extra_headers}\r\n", + body.len(), + ); + stream.write_all(response.as_bytes()).await?; + stream.write_all(body).await?; + stream.flush().await?; + Ok(()) +} + +fn is_expected_peer_disconnect(err: &anyhow::Error) -> bool { + err.chain().any(|cause| { + cause + .downcast_ref::() + .is_some_and(|io_err| { + matches!( + io_err.kind(), + ErrorKind::BrokenPipe | ErrorKind::ConnectionReset | ErrorKind::UnexpectedEof + ) + }) + }) +} + +/// Writes a chunked HTTP response so the shared client must drive the streaming path. +async fn respond_with_chunked_body( + mut stream: TcpStream, + headers: &[(&str, &str)], + chunks: &[&[u8]], +) -> anyhow::Result<()> { + let extra_headers = headers + .iter() + .map(|(name, value)| format!("{name}: {value}\r\n")) + .collect::(); + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: text/plain\r\ntransfer-encoding: chunked\r\nconnection: close\r\n{extra_headers}\r\n", + ); + stream.write_all(response.as_bytes()).await?; + for chunk in chunks { + stream + .write_all(format!("{:x}\r\n", chunk.len()).as_bytes()) + .await?; + stream.write_all(chunk).await?; + stream.write_all(b"\r\n").await?; + stream.flush().await?; + } + stream.write_all(b"0\r\n\r\n").await?; + stream.flush().await?; + Ok(()) +} + +/// Writes a chunked response and keeps the stream open until the test allows EOF. +async fn respond_with_chunked_body_until_finish( + mut stream: TcpStream, + headers: &[(&str, &str)], + chunks: &[&[u8]], + finish_rx: oneshot::Receiver<()>, +) -> anyhow::Result<()> { + let extra_headers = headers + .iter() + .map(|(name, value)| format!("{name}: {value}\r\n")) + .collect::(); + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: text/plain\r\ntransfer-encoding: chunked\r\nconnection: close\r\n{extra_headers}\r\n", + ); + stream.write_all(response.as_bytes()).await?; + for chunk in chunks { + stream + .write_all(format!("{:x}\r\n", chunk.len()).as_bytes()) + .await?; + stream.write_all(chunk).await?; + stream.write_all(b"\r\n").await?; + stream.flush().await?; + } + finish_rx.await?; + stream.write_all(b"0\r\n\r\n").await?; + stream.flush().await?; + Ok(()) +} + +/// Collects streamed response-body notifications until the terminal frame. +async fn collect_response_body_deltas( + server: &mut ExecServerHarness, + request_id: &str, +) -> anyhow::Result> { + let mut deltas = Vec::new(); + loop { + let event = server.next_event().await?; + let JSONRPCMessage::Notification(JSONRPCNotification { method, params }) = event else { + anyhow::bail!("expected http/request body delta notification, got {event:?}"); + }; + assert_eq!(method, "http/request/bodyDelta"); + let delta: HttpRequestBodyDeltaNotification = + serde_json::from_value(params.unwrap_or(Value::Null))?; + assert_eq!(delta.request_id, request_id); + + let done = delta.done; + deltas.push(delta); + if done { + return Ok(deltas); + } + } +} + +/// Returns a response header value without depending on header-name casing. +fn response_header(headers: &[HttpHeader], name: &str) -> Option { + headers + .iter() + .find(|header| header.name.eq_ignore_ascii_case(name)) + .map(|header| header.value.clone()) +} diff --git a/codex-rs/exec-server/tests/http_request_logging.rs b/codex-rs/exec-server/tests/http_request_logging.rs new file mode 100644 index 0000000000000000000000000000000000000000..45e26499575fe93e993d55e020f49a0e9141dca2 --- /dev/null +++ b/codex-rs/exec-server/tests/http_request_logging.rs @@ -0,0 +1,174 @@ +use std::io::Write; +use std::sync::Arc; +use std::sync::Mutex; + +use codex_exec_server::HttpClient; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::RouteAwareHttpClient; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use pretty_assertions::assert_eq; +use tokio::io::AsyncBufReadExt; +use tokio::io::AsyncWriteExt; +use tokio::io::BufReader; +use tokio::net::TcpListener; +use tracing_subscriber::Layer; +use tracing_subscriber::layer::SubscriberExt; + +#[tokio::test(flavor = "current_thread")] +async fn delegated_http_success_logs_do_not_expose_sensitive_request_or_response_data() +-> anyhow::Result<()> { + let log_buffer = Arc::new(Mutex::new(Vec::new())); + let writer_buffer = Arc::clone(&log_buffer); + let subscriber = tracing_subscriber::registry().with( + tracing_subscriber::fmt::layer() + .with_ansi(false) + .with_writer(move || TestLogWriter(Arc::clone(&writer_buffer))) + .with_filter( + tracing_subscriber::filter::Targets::new() + .with_target("codex_http_client", tracing::Level::TRACE) + .with_target("codex_exec_server", tracing::Level::TRACE), + ), + ); + let _guard = tracing::subscriber::set_default(subscriber); + tracing::debug!(target: "codex_exec_server", "log capture sentinel"); + let client = + RouteAwareHttpClient::new(HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault)); + + for (redirect_policy, status, query_secret, cookie_secret, location_secret) in [ + ( + HttpRedirectPolicy::Follow, + "200 OK", + "follow-query-secret", + "follow-cookie-secret", + "follow-location-secret", + ), + ( + HttpRedirectPolicy::Stop, + "302 Found", + "stop-query-secret", + "stop-cookie-secret", + "stop-location-secret", + ), + ] { + let listener = TcpListener::bind(("127.0.0.1", 0)).await?; + let address = listener.local_addr()?; + let response = format!( + "HTTP/1.1 {status}\r\nSet-Cookie: session={cookie_secret}\r\nLocation: http://127.0.0.1/private?token={location_secret}\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok" + ); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await?; + let mut reader = BufReader::new(stream); + loop { + let mut line = String::new(); + if reader.read_line(&mut line).await? == 0 { + anyhow::bail!("HTTP client disconnected before completing request headers"); + } + if line == "\r\n" { + break; + } + } + reader.get_mut().write_all(response.as_bytes()).await?; + anyhow::Ok(()) + }); + + let response = client + .http_request(HttpRequestParams { + method: "GET".to_string(), + url: format!("http://{address}/delegated?token={query_secret}"), + headers: Vec::new(), + body: None, + timeout_ms: Some(5_000), + redirect_policy, + request_id: "sensitive-request".to_string(), + stream_response: false, + }) + .await?; + let expected_status = match redirect_policy { + HttpRedirectPolicy::Follow => 200, + HttpRedirectPolicy::Stop => 302, + }; + assert_eq!(response.status, expected_status); + server.await??; + } + + let logs = String::from_utf8(log_buffer.lock().expect("log buffer lock").clone())?; + assert!(logs.contains("log capture sentinel")); + for secret in [ + "follow-query-secret", + "follow-cookie-secret", + "follow-location-secret", + "stop-query-secret", + "stop-cookie-secret", + "stop-location-secret", + ] { + assert!(!logs.contains(secret), "logs exposed {secret}:\n{logs}"); + } + + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn delegated_http_failure_warning_redacts_request_url() -> anyhow::Result<()> { + let log_buffer = Arc::new(Mutex::new(Vec::new())); + let writer_buffer = Arc::clone(&log_buffer); + let subscriber = tracing_subscriber::registry().with( + tracing_subscriber::fmt::layer() + .with_ansi(false) + .with_writer(move || TestLogWriter(Arc::clone(&writer_buffer))) + .with_filter( + tracing_subscriber::filter::Targets::new() + .with_target("codex_http_client", tracing::Level::TRACE) + .with_target("codex_exec_server", tracing::Level::TRACE), + ), + ); + let _guard = tracing::subscriber::set_default(subscriber); + let unavailable_server = std::net::TcpListener::bind(("127.0.0.1", 0))?; + let unavailable_address = unavailable_server.local_addr()?; + drop(unavailable_server); + let client = + RouteAwareHttpClient::new(HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault)); + + let error = client + .http_request(HttpRequestParams { + method: "GET".to_string(), + url: format!( + "http://{unavailable_address}/private-path-secret?token=failure-query-secret" + ), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "failed-sensitive-request".to_string(), + stream_response: false, + }) + .await; + assert!(error.is_err(), "request to a closed port should fail"); + + let logs = String::from_utf8(log_buffer.lock().expect("log buffer lock").clone())?; + assert!(logs.contains("http/request send failed")); + assert!(logs.contains("error_is_connect=true")); + for secret in ["private-path-secret", "failure-query-secret"] { + assert!(!logs.contains(secret), "logs exposed {secret}:\n{logs}"); + } + + Ok(()) +} + +#[derive(Clone)] +struct TestLogWriter(Arc>>); + +impl Write for TestLogWriter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.0 + .lock() + .map_err(|_| std::io::Error::other("log buffer lock"))? + .extend_from_slice(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} diff --git a/codex-rs/exec-server/tests/initialize.rs b/codex-rs/exec-server/tests/initialize.rs new file mode 100644 index 0000000000000000000000000000000000000000..8300f621641642a7ae1a377305c62953fb49c012 --- /dev/null +++ b/codex-rs/exec-server/tests/initialize.rs @@ -0,0 +1,163 @@ +mod common; + +use anyhow::Context; +use codex_build_info::BuildInfo; +use codex_build_info::build_id; +use codex_exec_server::EnvironmentInfo; +use codex_exec_server::InitializeParams; +use codex_exec_server::InitializeResponse; +use codex_exec_server_protocol::JSONRPCError; +use codex_exec_server_protocol::JSONRPCErrorError; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCResponse; +use common::TEST_BUILD_COMMIT; +use common::exec_server::ExecServerHarness; +use common::exec_server::exec_server_with_env; +use pretty_assertions::assert_eq; +use tempfile::TempDir; +use tokio::process::Command; +use uuid::Uuid; + +#[test_case::test_case(Some("1.2.3-alpha.4"); "packaged")] +#[test_case::test_case(None; "without_manifest")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_accepts_initialize(version: Option<&str>) -> anyhow::Result<()> { + let package = TempDir::new()?; + let bin_dir = package.path().join("bin"); + std::fs::create_dir(&bin_dir)?; + let executable = bin_dir.join(format!("codex{}", std::env::consts::EXE_SUFFIX)); + std::fs::copy(std::env::current_exe()?, &executable)?; + let manifest = package.path().join("codex-package.json"); + if let Some(version) = version { + std::fs::write( + &manifest, + serde_json::to_vec(&serde_json::json!({ "version": version }))?, + )?; + } + + let mut command = Command::new(&executable); + command.args(["exec-server", "--listen", "ws://127.0.0.1:0"]); + // Runtime environment variables cannot replace the executable's build stamp. + command.envs([ + ( + "STABLE_GIT_COMMIT", + "ffffffffffffffffffffffffffffffffffffffff", + ), + ("GITHUB_SHA", "ffffffffffffffffffffffffffffffffffffffff"), + ("CODEX_BUILD_TARGET", "runtime-override"), + ]); + let mut server = ExecServerHarness::start(command).await?; + + // Updates after startup cannot change the advertised release version. + std::fs::write(&manifest, r#"{"version":"9.9.9"}"#)?; + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + + let response = server.next_event().await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = response else { + panic!("expected initialize response"); + }; + assert_eq!(id, initialize_id); + let initialize_response: InitializeResponse = serde_json::from_value(result)?; + Uuid::parse_str(&initialize_response.session_id)?; + let mut expected_environment = EnvironmentInfo::local(); + expected_environment.executor_version = version.unwrap_or("0.0.0").to_string(); + let build_info = BuildInfo::get(); + let target = build_info + .target() + .context("the test binary has a compiled target")?; + expected_environment.provider_id = build_id(TEST_BUILD_COMMIT, target); + assert!(expected_environment.provider_id.is_some()); + assert_eq!( + initialize_response.environment_info, + Some(expected_environment.clone()) + ); + + server + .send_notification("initialized", serde_json::json!({})) + .await?; + std::fs::remove_file(&manifest)?; + let environment_id = server + .send_request("environment/info", serde_json::json!({})) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = server.next_event().await? + else { + panic!("expected environment info response"); + }; + assert_eq!(id, environment_id); + assert_eq!( + serde_json::from_value::(result)?, + expected_environment + ); + + server.shutdown().await?; + Ok(()) +} + +/// Requests retain their wire-order initialization errors even when later handshake messages are pipelined. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_rejects_pipelined_requests_before_initialized() -> anyhow::Result<()> { + let mut server = exec_server_with_env( + std::iter::empty::<(&str, &str)>(), + &["--concurrent-requests", "32"], + ) + .await?; + let before_initialize_id = server + .send_request("environment/info", serde_json::json!({})) + .await?; + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + + assert_eq!( + server.next_event().await?, + JSONRPCMessage::Error(JSONRPCError { + id: before_initialize_id, + error: JSONRPCErrorError { + code: -32600, + data: None, + message: "client must call initialize before using environment info methods" + .to_string(), + }, + }) + ); + let JSONRPCMessage::Response(JSONRPCResponse { id, .. }) = server.next_event().await? else { + panic!("expected initialize response"); + }; + assert_eq!(id, initialize_id); + + let before_initialized_id = server + .send_request("environment/info", serde_json::json!({})) + .await?; + server + .send_notification("initialized", serde_json::json!({})) + .await?; + assert_eq!( + server.next_event().await?, + JSONRPCMessage::Error(JSONRPCError { + id: before_initialized_id, + error: JSONRPCErrorError { + code: -32600, + data: None, + message: "client must send initialized before using environment info methods" + .to_string(), + }, + }) + ); + + server.shutdown().await?; + Ok(()) +} diff --git a/codex-rs/exec-server/tests/process.rs b/codex-rs/exec-server/tests/process.rs new file mode 100644 index 0000000000000000000000000000000000000000..a82bcbc6264f59a66e29a24b14a889977e63c6e8 --- /dev/null +++ b/codex-rs/exec-server/tests/process.rs @@ -0,0 +1,883 @@ +mod common; + +use std::collections::HashMap; +use std::time::Duration; + +use anyhow::Context; +use codex_exec_server::EnvironmentInfo; +use codex_exec_server::EnvironmentStatus; +use codex_exec_server::EnvironmentStatusKind; +use codex_exec_server::ExecResponse; +use codex_exec_server::InitializeParams; +use codex_exec_server::InitializeResponse; +use codex_exec_server::ProcessId; +use codex_exec_server::ReadResponse; +use codex_exec_server::TerminateResponse; +use codex_exec_server::WriteResponse; +use codex_exec_server::WriteStatus; +use codex_exec_server_protocol::JSONRPCError; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::ProcessSandboxType; +use codex_utils_path_uri::PathUri; +use common::exec_server::exec_server; +use common::exec_server::exec_server_with_env; +use pretty_assertions::assert_eq; +use tokio::time::sleep; +use tokio::time::timeout; + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_starts_process_over_websocket() -> anyhow::Result<()> { + let mut server = exec_server().await?; + let process_argv = if cfg!(windows) { + vec!["cmd.exe", "/D", "/C", "exit 0"] + } else { + vec!["true"] + }; + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + let _ = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &initialize_id + ) + }) + .await?; + + server + .send_notification("initialized", serde_json::json!({})) + .await?; + + let process_start_id = server + .send_request( + "process/start", + serde_json::json!({ + "processId": "proc-1", + "argv": process_argv, + "cwd": PathUri::from_host_native_path(std::env::current_dir()?)?, + "env": {}, + "tty": false, + "pipeStdin": false, + "arg0": null + }), + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &process_start_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = response else { + panic!("expected process/start response"); + }; + assert_eq!(id, process_start_id); + let process_start_response: ExecResponse = serde_json::from_value(result)?; + assert_eq!( + process_start_response, + ExecResponse { + process_id: ProcessId::from("proc-1"), + sandbox_type: Some(ProcessSandboxType::None), + } + ); + + server.shutdown().await?; + Ok(()) +} + +/// Ordinary requests run one at a time when concurrent processing is not enabled. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_runs_ordinary_requests_serially_by_default() -> anyhow::Result<()> { + let temporary_directory = tempfile::tempdir()?; + let temporary_directory_env_vars: &[&str] = if cfg!(windows) { + &["TEMP", "TMP"] + } else { + &["TMPDIR"] + }; + let mut server = exec_server_with_env( + temporary_directory_env_vars + .iter() + .map(|name| (*name, temporary_directory.path())), + &[], + ) + .await?; + let process_argv = if cfg!(windows) { + vec!["cmd.exe", "/D", "/C", "ping -n 601 127.0.0.1 >NUL"] + } else { + vec![ + "/bin/sh", + "-c", + "parent=$PPID; while kill -0 \"$parent\" 2>/dev/null; do sleep 1; done", + ] + }; + let process_env = if cfg!(windows) { + serde_json::json!({ "PATH": std::env::var("PATH")? }) + } else { + serde_json::json!({}) + }; + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &initialize_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { result, .. }) = response else { + panic!("expected initialize response"); + }; + let initialization: InitializeResponse = serde_json::from_value(result)?; + server + .send_notification("initialized", serde_json::json!({})) + .await?; + + let process_start_id = server + .send_request( + "process/start", + serde_json::json!({ + "processId": "proc-serial-read", + "argv": process_argv, + "cwd": PathUri::from_host_native_path(std::env::current_dir()?)?, + "env": process_env, + "tty": false, + "pipeStdin": false, + "arg0": null + }), + ) + .await?; + let _ = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &process_start_id + ) + }) + .await?; + + let read_id = server + .send_request( + "process/read", + serde_json::json!({ + "processId": "proc-serial-read", + "afterSeq": null, + "maxBytes": null, + "waitMs": 250 + }), + ) + .await?; + let queued_environment_info_id = server + .send_request("environment/info", serde_json::json!({})) + .await?; + let queued_start_id = server + .send_request( + "process/start", + serde_json::json!({ + "processId": "proc-serial-queued", + "argv": process_argv, + "cwd": PathUri::from_host_native_path(std::env::current_dir()?)?, + "env": process_env, + "tty": false, + "pipeStdin": false, + "arg0": null + }), + ) + .await?; + + let response = server + .wait_for_event(|event| matches!(event, JSONRPCMessage::Response(_))) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, .. }) = response else { + panic!("expected the blocked process/read to finish before the queued process/start"); + }; + assert_eq!(id, read_id); + let response = server + .wait_for_event(|event| matches!(event, JSONRPCMessage::Response(_))) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = response else { + panic!("expected the queued environment/info response after process/read"); + }; + assert_eq!(id, queued_environment_info_id); + let mut expected_environment_info = EnvironmentInfo::local(); + expected_environment_info.provider_id = initialization + .environment_info + .and_then(|info| info.provider_id); + expected_environment_info.temporary_directories = Some(vec![PathUri::from_host_native_path( + temporary_directory.path(), + )?]); + expected_environment_info.temp_dir = + Some(PathUri::from_host_native_path(temporary_directory.path())?); + assert_eq!( + serde_json::from_value::(result)?, + expected_environment_info + ); + let response = server + .wait_for_event(|event| matches!(event, JSONRPCMessage::Response(_))) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = response else { + panic!("expected the queued process/start response after process/read"); + }; + assert_eq!(id, queued_start_id); + assert_eq!( + serde_json::from_value::(result)?, + ExecResponse { + process_id: ProcessId::from("proc-serial-queued"), + sandbox_type: Some(ProcessSandboxType::None), + } + ); + + for process_id in ["proc-serial-read", "proc-serial-queued"] { + let terminate_id = server + .send_request( + "process/terminate", + serde_json::json!({ "processId": process_id }), + ) + .await?; + server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &terminate_id + ) + }) + .await?; + } + + server.shutdown().await?; + Ok(()) +} + +/// A long read cannot block health checks; saturated ordinary work cannot block health checks or cleanup. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_keeps_control_requests_live_during_long_reads_and_queued_requests() +-> anyhow::Result<()> { + let mut server = exec_server_with_env( + std::iter::empty::<(&str, &str)>(), + &["--concurrent-requests", "32"], + ) + .await?; + let process_argv = if cfg!(windows) { + vec!["cmd.exe", "/D", "/C", "ping -n 601 127.0.0.1 >NUL"] + } else { + vec![ + "/bin/sh", + "-c", + "parent=$PPID; while kill -0 \"$parent\" 2>/dev/null; do sleep 1; done", + ] + }; + let process_env = if cfg!(windows) { + serde_json::json!({ "PATH": std::env::var("PATH")? }) + } else { + serde_json::json!({}) + }; + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + assert!(matches!( + server.next_event().await?, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == initialize_id + )); + server + .send_notification("initialized", serde_json::json!({})) + .await?; + + let process_start_id = server + .send_request( + "process/start", + serde_json::json!({ + "processId": "proc-capacity", + "argv": process_argv, + "cwd": PathUri::from_host_native_path(std::env::current_dir()?)?, + "env": process_env, + "tty": false, + "pipeStdin": false, + "arg0": null + }), + ) + .await?; + assert!(matches!( + server.next_event().await?, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == process_start_id + )); + + let read_params = serde_json::json!({ + "processId": "proc-capacity", + "afterSeq": null, + "maxBytes": null, + "waitMs": 600_000 + }); + server + .send_request("process/read", read_params.clone()) + .await?; + let concurrent_start_id = server + .send_request( + "process/start", + serde_json::json!({ + "processId": "proc-concurrent", + "argv": process_argv, + "cwd": PathUri::from_host_native_path(std::env::current_dir()?)?, + "env": process_env, + "tty": false, + "pipeStdin": false, + "arg0": null + }), + ) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = server.next_event().await? + else { + panic!("expected process/start to finish before the pending process/read"); + }; + assert_eq!(id, concurrent_start_id); + assert_eq!( + serde_json::from_value::(result)?, + ExecResponse { + process_id: ProcessId::from("proc-concurrent"), + sandbox_type: Some(ProcessSandboxType::None), + } + ); + for _ in 1..32 { + server + .send_request("process/read", read_params.clone()) + .await?; + } + + let queued_read_id = server.send_request("process/read", read_params).await?; + + let environment_info_id = server + .send_request("environment/info", serde_json::json!({})) + .await?; + let response = server + .wait_for_event(|event| matches!(event, JSONRPCMessage::Response(_))) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = response else { + panic!("expected environment/info response at regular request capacity"); + }; + assert_eq!(id, environment_info_id); + let _: EnvironmentInfo = serde_json::from_value(result)?; + + let environment_status_id = server + .send_request("environment/status", serde_json::json!({})) + .await?; + let response = server + .wait_for_event(|event| matches!(event, JSONRPCMessage::Response(_))) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = response else { + panic!("expected environment/status response at regular request capacity"); + }; + assert_eq!(id, environment_status_id); + assert_eq!( + serde_json::from_value::(result)?, + EnvironmentStatus { + status: EnvironmentStatusKind::Ready, + } + ); + + let terminate_id = server + .send_request( + "process/terminate", + serde_json::json!({ "processId": "proc-capacity" }), + ) + .await?; + let mut terminate_response = None; + let mut queued_read_completed = false; + while terminate_response.is_none() || !queued_read_completed { + match server.next_event().await? { + JSONRPCMessage::Response(JSONRPCResponse { id, result }) if id == terminate_id => { + terminate_response = Some(serde_json::from_value::(result)?); + } + JSONRPCMessage::Response(JSONRPCResponse { id, result }) if id == queued_read_id => { + let _: ReadResponse = serde_json::from_value(result)?; + queued_read_completed = true; + } + JSONRPCMessage::Error(error) => { + anyhow::bail!("unexpected error while waiting for queued requests: {error:?}"); + } + JSONRPCMessage::Request(_) + | JSONRPCMessage::Response(_) + | JSONRPCMessage::Notification(_) => {} + } + } + assert_eq!( + terminate_response, + Some(TerminateResponse { running: true }) + ); + + let terminate_id = server + .send_request( + "process/terminate", + serde_json::json!({ "processId": "proc-concurrent" }), + ) + .await?; + server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &terminate_id + ) + }) + .await?; + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_defaults_omitted_pipe_stdin_to_closed_stdin() -> anyhow::Result<()> { + let mut server = exec_server().await?; + let process_argv = if cfg!(windows) { + vec!["cmd.exe", "/D", "/C", "ping -n 2 127.0.0.1 >NUL"] + } else { + vec![ + "/bin/sh", + "-c", + "sleep 0.3; if IFS= read -r line; then printf 'read:%s\\n' \"$line\"; else printf 'eof\\n'; fi", + ] + }; + let process_env = if cfg!(windows) { + serde_json::json!({ "PATH": std::env::var("PATH")? }) + } else { + serde_json::json!({}) + }; + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + let _ = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &initialize_id + ) + }) + .await?; + + server + .send_notification("initialized", serde_json::json!({})) + .await?; + + let process_start_id = server + .send_request( + "process/start", + serde_json::json!({ + "processId": "proc-default-stdin", + "argv": process_argv, + "cwd": PathUri::from_host_native_path(std::env::current_dir()?)?, + "env": process_env, + "tty": false, + "arg0": null + }), + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &process_start_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { result, .. }) = response else { + panic!("expected process/start response"); + }; + let process_start_response: ExecResponse = serde_json::from_value(result)?; + assert_eq!( + process_start_response, + ExecResponse { + process_id: ProcessId::from("proc-default-stdin"), + sandbox_type: Some(ProcessSandboxType::None), + } + ); + + let write_id = server + .send_request( + "process/write", + serde_json::json!({ + "processId": "proc-default-stdin", + "chunk": "aWdub3JlZAo=", + "writeId": "write-default-stdin" + }), + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &write_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { result, .. }) = response else { + panic!("expected process/write response"); + }; + let write_response: WriteResponse = serde_json::from_value(result)?; + assert_eq!( + write_response, + WriteResponse { + status: WriteStatus::StdinClosed + } + ); + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_dedupes_retried_process_write_ids() -> anyhow::Result<()> { + let mut server = exec_server().await?; + let process_argv = if cfg!(windows) { + vec![ + "powershell.exe", + "-NoProfile", + "-NonInteractive", + "-Command", + "[Console]::Out.WriteLine('line:' + [Console]::In.ReadLine()); [Console]::Out.WriteLine('line:' + [Console]::In.ReadLine())", + ] + } else { + vec![ + "/bin/sh", + "-c", + "IFS= read -r first; printf 'line:%s\\n' \"$first\"; IFS= read -r second; printf 'line:%s\\n' \"$second\"", + ] + }; + let process_env = if cfg!(windows) { + serde_json::to_value(std::env::vars().collect::>())? + } else { + serde_json::json!({}) + }; + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + let _ = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &initialize_id + ) + }) + .await?; + + server + .send_notification("initialized", serde_json::json!({})) + .await?; + + let process_start_id = server + .send_request( + "process/start", + serde_json::json!({ + "processId": "proc-write-id", + "argv": process_argv, + "cwd": PathUri::from_host_native_path(std::env::current_dir()?)?, + "env": process_env, + "tty": false, + "pipeStdin": true, + "arg0": null + }), + ) + .await?; + let _ = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &process_start_id + ) + }) + .await?; + + for (write_id, chunk) in [ + ("write-1", "Zmlyc3QK"), + ("write-1", "Zmlyc3QK"), + ("write-2", "c2Vjb25kCg=="), + ] { + let request_id = server + .send_request( + "process/write", + serde_json::json!({ + "processId": "proc-write-id", + "chunk": chunk, + "writeId": write_id + }), + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &request_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { result, .. }) = response else { + panic!("expected process/write response"); + }; + let write_response: WriteResponse = serde_json::from_value(result)?; + assert_eq!( + write_response, + WriteResponse { + status: WriteStatus::Accepted + } + ); + } + + let mut after_seq = None; + let mut output = Vec::new(); + for _ in 0..5 { + let read_id = server + .send_request( + "process/read", + serde_json::json!({ + "processId": "proc-write-id", + "afterSeq": after_seq, + "maxBytes": null, + "waitMs": 1000 + }), + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &read_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { result, .. }) = response else { + panic!("expected process/read response"); + }; + let read_response: ReadResponse = serde_json::from_value(result)?; + output.extend( + read_response + .chunks + .into_iter() + .flat_map(|chunk| chunk.chunk.into_inner()), + ); + after_seq = Some(read_response.next_seq.saturating_sub(1)); + if read_response.closed + || output.ends_with(b"line:second\n") + || output.ends_with(b"line:second\r\n") + { + break; + } + } + + assert_eq!( + String::from_utf8(output)?.replace("\r\n", "\n"), + "line:first\nline:second\n".to_string() + ); + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_resumes_detached_session_without_killing_processes() -> anyhow::Result<()> { + const SESSION_ALREADY_ATTACHED_ERROR_CODE: i64 = -32010; + + let mut server = exec_server().await?; + // Keep the process alive until the test explicitly terminates it. + let process_argv = if cfg!(windows) { + vec!["cmd.exe", "/D", "/C", "set /p line="] + } else { + vec!["/bin/sh", "-c", "IFS= read -r line"] + }; + let process_env = if cfg!(windows) { + serde_json::json!({ "PATH": std::env::var("PATH")? }) + } else { + serde_json::json!({}) + }; + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) + | JSONRPCMessage::Error(JSONRPCError { id, .. }) + if id == &initialize_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { result, .. }) = response else { + anyhow::bail!("expected initialize response, got {response:?}"); + }; + let initialize_response: InitializeResponse = serde_json::from_value(result)?; + + server + .send_notification("initialized", serde_json::json!({})) + .await?; + + let process_start_id = server + .send_request( + "process/start", + serde_json::json!({ + "processId": "proc-resume", + "argv": process_argv, + "cwd": PathUri::from_host_native_path(std::env::current_dir()?)?, + "env": process_env, + "tty": false, + "pipeStdin": true, + "arg0": null + }), + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) + | JSONRPCMessage::Error(JSONRPCError { id, .. }) + if id == &process_start_id + ) + }) + .await?; + let JSONRPCMessage::Response(_) = response else { + anyhow::bail!("expected process/start response, got {response:?}"); + }; + + server.disconnect_websocket().await?; + server.reconnect_websocket().await?; + + // Closing the old socket does not wait for the server to detach its session. + let result = timeout(Duration::from_secs(5), async { + loop { + let resume_initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: Some(initialize_response.session_id.clone()), + })?, + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) + | JSONRPCMessage::Error(JSONRPCError { id, .. }) + if id == &resume_initialize_id + ) + }) + .await?; + match response { + JSONRPCMessage::Response(JSONRPCResponse { result, .. }) => break Ok(result), + JSONRPCMessage::Error(JSONRPCError { error, .. }) + if error.code == SESSION_ALREADY_ATTACHED_ERROR_CODE => + { + sleep(Duration::from_millis(25)).await; + } + JSONRPCMessage::Error(error) => { + anyhow::bail!("resume initialize failed: {error:?}"); + } + JSONRPCMessage::Request(_) | JSONRPCMessage::Notification(_) => { + unreachable!("wait_for_event only returns the matching response or error"); + } + } + } + }) + .await + .context("timed out resuming exec-server session after disconnect")??; + let resumed_response: InitializeResponse = serde_json::from_value(result)?; + assert_eq!(resumed_response, initialize_response); + + server + .send_notification("initialized", serde_json::json!({})) + .await?; + + let process_read_id = server + .send_request( + "process/read", + serde_json::json!({ + "processId": "proc-resume", + "afterSeq": null, + "maxBytes": null, + "waitMs": 0 + }), + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) + | JSONRPCMessage::Error(JSONRPCError { id, .. }) + if id == &process_read_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { result, .. }) = response else { + anyhow::bail!("expected process/read response, got {response:?}"); + }; + let process_read_response: ReadResponse = serde_json::from_value(result)?; + assert!(process_read_response.failure.is_none()); + assert!(!process_read_response.exited); + assert!(!process_read_response.closed); + + let terminate_id = server + .send_request( + "process/terminate", + serde_json::json!({ + "processId": "proc-resume" + }), + ) + .await?; + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) + | JSONRPCMessage::Error(JSONRPCError { id, .. }) + if id == &terminate_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { result, .. }) = response else { + anyhow::bail!("expected process/terminate response, got {response:?}"); + }; + let terminate_response: TerminateResponse = serde_json::from_value(result)?; + assert_eq!(terminate_response, TerminateResponse { running: true }); + + server.shutdown().await?; + Ok(()) +} diff --git a/codex-rs/exec-server/tests/relay.rs b/codex-rs/exec-server/tests/relay.rs new file mode 100644 index 0000000000000000000000000000000000000000..d37f76331a20043b5d8322e1e99b5855e74cd509 --- /dev/null +++ b/codex-rs/exec-server/tests/relay.rs @@ -0,0 +1,423 @@ +mod common; + +#[path = "common/relay.rs"] +mod relay_support; + +#[path = "relay/registration_retry_tests.rs"] +mod registration_retry; + +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; + +use anyhow::Context; +use anyhow::Result; +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD; +use codex_exec_server::EnvironmentConnectionState; +use codex_exec_server::EnvironmentManager; +use codex_exec_server::EnvironmentReadyInfo; +use codex_exec_server::ExecParams; +use codex_exec_server::ExecResponse; +use codex_exec_server::ExecServerError; +use codex_exec_server::ExecServerRuntimePaths; +use codex_exec_server::FsReadFileParams; +use codex_exec_server::NoiseChannelPublicKey; +use codex_exec_server::NoiseRendezvousConnectBundle; +use codex_exec_server::NoiseRendezvousConnectProvider; +use codex_exec_server::ProcessId; +use codex_exec_server::RemoteEnvironmentConfig; +use codex_exec_server_protocol::ProcessSandboxType; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_http_client::cache_system_proxy_route_for_test; +use codex_protocol::capabilities::CapabilityRootLocation; +use codex_protocol::capabilities::SelectedCapabilityRoot; +use codex_utils_path_uri::PathUri; +use futures::future::BoxFuture; +use pretty_assertions::assert_eq; +use relay_support::ENVIRONMENT_ID; +use relay_support::EXECUTOR_REGISTRATION_ID; +use relay_support::HARNESS_KEY_AUTHORIZATION; +use relay_support::RelayTest; +use relay_support::TEST_TIMEOUT; +use relay_support::accept_websocket; +use relay_support::proxy_relay_frames; +use relay_support::registered_executor_public_key; +use relay_support::static_registry_auth_provider; +use tempfile::TempDir; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::sync::mpsc; +use tokio::sync::watch; +use tokio::task::JoinSet; +use tokio::time::timeout; +use tokio_util::task::AbortOnDropHandle; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +struct FreshBundleNoiseConnectProvider { + websocket_url: String, + executor_public_key: NoiseChannelPublicKey, + calls: AtomicUsize, +} + +impl FreshBundleNoiseConnectProvider { + fn calls(&self) -> usize { + self.calls.load(Ordering::Relaxed) + } +} + +impl NoiseRendezvousConnectProvider for FreshBundleNoiseConnectProvider { + fn connect_bundle( + &self, + _: NoiseChannelPublicKey, + ) -> BoxFuture<'_, Result> { + let call = self.calls.fetch_add(1, Ordering::Relaxed) + 1; + let bundle = NoiseRendezvousConnectBundle { + websocket_url: self.websocket_url.clone(), + environment_id: ENVIRONMENT_ID.to_string(), + executor_registration_id: EXECUTOR_REGISTRATION_ID.to_string(), + executor_public_key: self.executor_public_key.clone(), + harness_key_authorization: format!("{HARNESS_KEY_AUTHORIZATION}-{call}"), + }; + Box::pin(async move { Ok(bundle) }) + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +#[serial_test::serial] +async fn failed_noise_environment_recovers_and_reconnects_after_ready_report() -> Result<()> { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let rendezvous_address = listener.local_addr()?; + let environment_rendezvous_url = + "ws://environment-noise-relay-system-proxy.invalid:8765/relay?role=environment"; + let harness_rendezvous_url = + "ws://harness-noise-relay-system-proxy.invalid:8765/relay?role=harness"; + let proxy_listener = TcpListener::bind("127.0.0.1:0").await?; + let proxy_url = format!("http://{}", proxy_listener.local_addr()?); + for rendezvous_url in [environment_rendezvous_url, harness_rendezvous_url] { + let proxy_resolution_url = rendezvous_url.replacen("ws://", "http://", /*count*/ 1); + cache_system_proxy_route_for_test(&proxy_resolution_url, proxy_url.clone()); + } + let (proxy_request_tx, mut proxy_request_rx) = mpsc::unbounded_channel(); + let _proxy_task = AbortOnDropHandle::new(tokio::spawn(async move { + let mut proxy_connections = JoinSet::new(); + while let Ok((mut client, _)) = proxy_listener.accept().await { + let proxy_request_tx = proxy_request_tx.clone(); + proxy_connections.spawn(async move { + let mut request = Vec::new(); + let mut byte = [0_u8; 1]; + while !request.ends_with(b"\r\n\r\n") { + client.read_exact(&mut byte).await?; + request.push(byte[0]); + } + let request_line = String::from_utf8(request)? + .lines() + .next() + .context("system proxy should receive a CONNECT request")? + .to_string(); + proxy_request_tx + .send(request_line) + .map_err(|_| anyhow::anyhow!("system proxy request receiver was dropped"))?; + let mut target = TcpStream::connect(rendezvous_address).await?; + client + .write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n") + .await?; + tokio::io::copy_bidirectional(&mut client, &mut target).await?; + Ok::<(), anyhow::Error>(()) + }); + } + })); + let registry = MockServer::start().await; + Mock::given(method("POST")) + .and(path(format!( + "/cloud/environment/{ENVIRONMENT_ID}/register" + ))) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "environment_id": ENVIRONMENT_ID, + "url": environment_rendezvous_url, + "security_profile": "noise_hybrid_ik_v1", + "executor_registration_id": EXECUTOR_REGISTRATION_ID, + }))) + .expect(1) + .mount(®istry) + .await; + Mock::given(method("POST")) + .and(path(format!( + "/cloud/environment/{ENVIRONMENT_ID}/validate" + ))) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "valid": true, + }))) + .expect(2) + .mount(®istry) + .await; + + let (codex_exe, codex_linux_sandbox_exe) = common::current_test_binary_helper_paths()?; + let runtime_paths = ExecServerRuntimePaths::new(codex_exe, codex_linux_sandbox_exe)?; + let http_client_factory = HttpClientFactory::new(OutboundProxyPolicy::RespectSystemProxy); + let config = RemoteEnvironmentConfig::new( + registry.uri(), + ENVIRONMENT_ID.to_string(), + static_registry_auth_provider(), + http_client_factory.clone(), + )?; + let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel(); + let remote_environment = AbortOnDropHandle::new(tokio::spawn( + codex_exec_server::run_remote_environment_until_shutdown( + config, + runtime_paths, + async move { + let _ = shutdown_rx.await; + }, + ), + )); + + let environment_websocket = accept_websocket(&listener, "environment").await?; + let environment_proxy_request = + "CONNECT environment-noise-relay-system-proxy.invalid:8765 HTTP/1.1"; + let harness_proxy_request = "CONNECT harness-noise-relay-system-proxy.invalid:8765 HTTP/1.1"; + assert_eq!( + timeout(TEST_TIMEOUT, proxy_request_rx.recv()).await?, + Some(environment_proxy_request.to_string()) + ); + let provider = Arc::new(FreshBundleNoiseConnectProvider { + websocket_url: harness_rendezvous_url.to_string(), + executor_public_key: registered_executor_public_key(®istry).await?, + calls: AtomicUsize::new(0), + }); + let manager = Arc::new(EnvironmentManager::without_environments( + http_client_factory, + )); + let environment = manager + .materialize_pending_noise_environment(ENVIRONMENT_ID.to_string(), provider.clone())?; + let mut connection_state = environment + .subscribe_connection_state() + .context("remote environment connection state")?; + + let capability_root = TempDir::new()?; + let skill_file = capability_root.path().join("SKILL.md"); + let skill_contents = b"# Recovered capability\n"; + std::fs::write(&skill_file, skill_contents)?; + let selected_capability_roots = vec![SelectedCapabilityRoot { + id: "executor-plugin".to_string(), + location: CapabilityRootLocation::Environment { + environment_id: ENVIRONMENT_ID.to_string(), + path: PathUri::from_host_native_path(capability_root.path())?, + }, + }]; + assert!( + manager + .resolve_selected_capability_roots(&selected_capability_roots, &HashMap::new()) + .await + .is_empty() + ); + assert_eq!(provider.calls(), 0); + manager.report_environment_provisioning_status( + ENVIRONMENT_ID.to_string(), + Err("first provisioning attempt failed".to_string()), + provider.clone(), + )?; + timeout(TEST_TIMEOUT, async { + while !environment.startup_finished() { + tokio::task::yield_now().await; + } + }) + .await + .expect("failed capability startup should record its completion"); + assert_eq!(provider.calls(), 0); + let reported = manager + .report_environment_provisioning_status( + ENVIRONMENT_ID.to_string(), + Ok(EnvironmentReadyInfo { + selected_capability_roots: selected_capability_roots.clone(), + }), + provider.clone(), + )? + .context("ready report should apply to the pending environment")?; + assert!(Arc::ptr_eq(&environment, &reported)); + assert_eq!(provider.calls(), 0); + // Capability resolution must retry on its own, without an explicit connection call. + let resolved_roots = tokio::spawn({ + let manager = Arc::clone(&manager); + let selected_capability_roots = selected_capability_roots.clone(); + async move { + manager + .resolve_selected_capability_roots(&selected_capability_roots, &HashMap::new()) + .await + } + }); + let harness_websocket = accept_websocket(&listener, "harness").await?; + assert_eq!( + timeout(TEST_TIMEOUT, proxy_request_rx.recv()).await?, + Some(harness_proxy_request.to_string()) + ); + let first_relay = tokio::spawn(proxy_relay_frames( + environment_websocket, + harness_websocket, + Arc::new(Mutex::new(Vec::new())), + )); + let resolved_roots = timeout(TEST_TIMEOUT, resolved_roots) + .await + .context("capability resolution should recover after Ready")??; + assert_eq!( + resolved_roots + .iter() + .map(|root| root.selected_root().clone()) + .collect::>(), + selected_capability_roots + ); + let [resolved_root] = resolved_roots.as_slice() else { + anyhow::bail!("the recovered capability root should resolve"); + }; + assert!(Arc::ptr_eq(resolved_root.environment(), &environment)); + let recovered_skill = resolved_root + .environment() + .get_filesystem() + .read_file( + &PathUri::from_host_native_path(skill_file)?, + Default::default(), + /*sandbox*/ None, + ) + .await?; + assert_eq!(recovered_skill, skill_contents.to_vec()); + let initial_info = environment.info().await?; + assert_eq!( + environment.selected_capability_roots(), + selected_capability_roots + ); + assert_eq!(provider.calls(), 1); + assert_eq!( + next_connection_state(&mut connection_state).await?, + EnvironmentConnectionState::Connected + ); + + first_relay.abort(); + let _ = first_relay.await; + assert_eq!( + next_connection_state(&mut connection_state).await?, + EnvironmentConnectionState::Disconnected + ); + let first_reconnected_websocket = accept_websocket(&listener, "reconnected peer").await?; + let second_reconnected_websocket = accept_websocket(&listener, "reconnected peer").await?; + let mut reconnect_proxy_requests = vec![ + timeout(TEST_TIMEOUT, proxy_request_rx.recv()) + .await? + .context("first reconnected peer should use the system proxy")?, + timeout(TEST_TIMEOUT, proxy_request_rx.recv()) + .await? + .context("second reconnected peer should use the system proxy")?, + ]; + reconnect_proxy_requests.sort(); + assert_eq!( + reconnect_proxy_requests, + vec![ + environment_proxy_request.to_string(), + harness_proxy_request.to_string(), + ] + ); + let second_relay = tokio::spawn(proxy_relay_frames( + first_reconnected_websocket, + second_reconnected_websocket, + Arc::new(Mutex::new(Vec::new())), + )); + assert_eq!( + next_connection_state(&mut connection_state).await?, + EnvironmentConnectionState::Connected + ); + let recovered_info = environment.info().await?; + + assert_eq!(recovered_info, initial_info); + assert_eq!( + environment.selected_capability_roots(), + selected_capability_roots + ); + assert_eq!(provider.calls(), 2); + registry.verify().await; + + second_relay.abort(); + let _ = second_relay.await; + let _ = shutdown_tx.send(()); + timeout(TEST_TIMEOUT, remote_environment).await???; + Ok(()) +} + +async fn next_connection_state( + state: &mut watch::Receiver, +) -> Result { + timeout(TEST_TIMEOUT, state.changed()).await??; + Ok(*state.borrow_and_update()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn remote_environment_routes_encrypted_exec_server_rpc() -> Result<()> { + let relay = RelayTest::new().await?; + let (codex_exe, codex_linux_sandbox_exe) = common::current_test_binary_helper_paths()?; + let runtime_paths = ExecServerRuntimePaths::new(codex_exe, codex_linux_sandbox_exe)?; + let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel(); + let remote_environment = AbortOnDropHandle::new(tokio::spawn( + codex_exec_server::run_remote_environment_until_shutdown( + relay.config()?, + runtime_paths, + async move { + let _ = shutdown_rx.await; + }, + ), + )); + let connection = relay.connect().await?; + let client = &connection.client; + + let exec_params = ExecParams { + metadata: Default::default(), + process_id: ProcessId::from("proc-1"), + argv: vec!["true".to_string()], + cwd: PathUri::from_host_native_path(std::env::current_dir()?)?, + shell_snapshot: None, + env_policy: None, + env: HashMap::new(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }; + let response = client.exec(exec_params).await?; + assert_eq!( + response, + ExecResponse { + process_id: ProcessId::from("proc-1"), + sandbox_type: Some(ProcessSandboxType::None), + } + ); + + let temp_dir = TempDir::new()?; + let large_file_path = temp_dir.path().join("large-response.bin"); + let large_file_contents = vec![0x5a; 128 * 1024]; + std::fs::write(&large_file_path, &large_file_contents)?; + let read_response = client + .fs_read_file(FsReadFileParams { + path: PathUri::from_host_native_path(large_file_path)?, + follow_symlinks: None, + sandbox: None, + }) + .await?; + assert_eq!( + STANDARD.decode(read_response.data_base64)?, + large_file_contents + ); + connection.assert_encrypted()?; + connection.close().await; + let _ = shutdown_tx.send(()); + timeout(TEST_TIMEOUT, remote_environment).await???; + Ok(()) +} diff --git a/codex-rs/exec-server/tests/relay/registration_retry_tests.rs b/codex-rs/exec-server/tests/relay/registration_retry_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..99455a6cc7415399c506f2caf0128225846646b4 --- /dev/null +++ b/codex-rs/exec-server/tests/relay/registration_retry_tests.rs @@ -0,0 +1,210 @@ +//! Confirmed registration conflicts retry without replacing the Noise identity or session handler. + +use std::time::Duration; + +use codex_exec_server::ExecServerClient; +use codex_exec_server::NoiseChannelIdentity; +use codex_exec_server::NoiseRendezvousConnectArgs; +use pretty_assertions::assert_eq; +use tokio::sync::oneshot; +use tokio_tungstenite::accept_hdr_async; +use tokio_tungstenite::tungstenite::handshake::server::Request; +use tokio_tungstenite::tungstenite::handshake::server::Response; + +use super::*; + +type RemoteTask = AbortOnDropHandle>; +type RelayTask = AbortOnDropHandle>; + +struct RegistryFixture { + registry: MockServer, + listener: TcpListener, + requests: mpsc::UnboundedReceiver, +} + +impl RegistryFixture { + async fn new(statuses: Vec, response_delay: Duration) -> Result { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let rendezvous_url = format!("ws://{}", listener.local_addr()?); + let registry = MockServer::start().await; + let (request_tx, requests) = mpsc::unbounded_channel(); + let attempts = AtomicUsize::new(0); + Mock::given(method("POST")) + .and(path(format!("/cloud/environment/{ENVIRONMENT_ID}/register"))) + .respond_with(move |_: &wiremock::Request| { + let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1; + let _ = request_tx.send(attempt); + let status = statuses[(attempt - 1).min(statuses.len() - 1)]; + let mut response = ResponseTemplate::new(status).set_delay(response_delay); + if status == 200 { + response = response.set_body_json(serde_json::json!({ + "environment_id": ENVIRONMENT_ID, + "url": format!("{rendezvous_url}/relay?role=environment®istration={attempt}"), + "security_profile": "noise_hybrid_ik_v1", + "executor_registration_id": format!("registration-{attempt}"), + })); + } else { + response = response.set_body_json(serde_json::json!({ + "error": {"code": "registration_conflict", "message": "registration conflicted"}, + })); + } + response + }) + .mount(®istry) + .await; + Mock::given(method("POST")) + .and(path(format!( + "/cloud/environment/{ENVIRONMENT_ID}/validate" + ))) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "valid": true, + }))) + .mount(®istry) + .await; + Ok(Self { + registry, + listener, + requests, + }) + } + + fn start(&self) -> Result<(oneshot::Sender<()>, RemoteTask)> { + let config = RemoteEnvironmentConfig::new( + self.registry.uri(), + ENVIRONMENT_ID.to_string(), + static_registry_auth_provider(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )?; + let (codex_exe, sandbox_exe) = common::current_test_binary_helper_paths()?; + let runtime_paths = ExecServerRuntimePaths::new(codex_exe, sandbox_exe)?; + let (shutdown_tx, shutdown_rx) = oneshot::channel(); + let task = AbortOnDropHandle::new(tokio::spawn( + codex_exec_server::run_remote_environment_until_shutdown( + config, + runtime_paths, + async move { + let _ = shutdown_rx.await; + }, + ), + )); + Ok((shutdown_tx, task)) + } + + async fn connect( + &self, + registration_id: &str, + harness_identity: &NoiseChannelIdentity, + resume_session_id: Option, + ) -> Result<(ExecServerClient, RelayTask)> { + let environment = accept_websocket(&self.listener, "environment").await?; + let args = NoiseRendezvousConnectArgs { + bundle: NoiseRendezvousConnectBundle { + websocket_url: format!("ws://{}/relay?role=harness", self.listener.local_addr()?), + environment_id: ENVIRONMENT_ID.to_string(), + executor_registration_id: registration_id.to_string(), + executor_public_key: registered_executor_public_key(&self.registry).await?, + harness_key_authorization: HARNESS_KEY_AUTHORIZATION.to_string(), + }, + harness_identity: harness_identity.clone(), + client_name: "registration-retry-test".to_string(), + connect_timeout: TEST_TIMEOUT, + initialize_timeout: TEST_TIMEOUT, + resume_session_id, + http_client_factory: HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + }; + let client = AbortOnDropHandle::new(tokio::spawn(async move { + ExecServerClient::connect_noise_rendezvous(args).await + })); + let harness = accept_websocket(&self.listener, "harness").await?; + let relay = AbortOnDropHandle::new(tokio::spawn(proxy_relay_frames( + environment, + harness, + Arc::new(Mutex::new(Vec::new())), + ))); + Ok((timeout(TEST_TIMEOUT, client).await???, relay)) + } + + async fn registered_keys(&self) -> Result> { + self.registry + .received_requests() + .await + .context("registry should retain requests")? + .iter() + .filter(|request| request.url.path().ends_with("/register")) + .map(|request| { + let body: serde_json::Value = serde_json::from_slice(&request.body)?; + Ok(serde_json::from_value(body["executor_public_key"].clone())?) + }) + .collect() + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn registration_retries_preserve_noise_identity_and_initialized_session() -> Result<()> { + let fixture = RegistryFixture::new(vec![503, 200, 503, 503, 200], Duration::ZERO).await?; + let (shutdown, remote) = fixture.start()?; + let harness_identity = NoiseChannelIdentity::generate()?; + let (client, first_relay) = fixture + .connect( + "registration-2", + &harness_identity, + /*resume_session_id*/ None, + ) + .await?; + let session_id = client.session_id().context("initialized session ID")?; + let environment_info = client.force_environment_info().await?; + let key = registered_executor_public_key(&fixture.registry).await?; + assert_eq!(fixture.registered_keys().await?, vec![key.clone(); 2]); + + first_relay.abort(); + let _ = first_relay.await; + drop(client); + let (socket, _) = timeout(TEST_TIMEOUT, fixture.listener.accept()).await??; + let rejected = http::Response::builder() + .status(http::StatusCode::UNAUTHORIZED) + .body(Some("expired registration".to_string()))?; + let stale_url = accept_hdr_async(socket, |request: &Request, _: Response| { + assert_eq!( + request + .uri() + .path_and_query() + .map(http::uri::PathAndQuery::as_str), + Some("/relay?role=environment®istration=2") + ); + Err(rejected) + }); + assert!(timeout(TEST_TIMEOUT, stale_url).await?.is_err()); + + let (resumed, second_relay) = fixture + .connect( + "registration-5", + &harness_identity, + Some(session_id.clone()), + ) + .await?; + assert_eq!(resumed.session_id(), Some(session_id)); + assert_eq!(resumed.force_environment_info().await?, environment_info); + assert_eq!(fixture.registered_keys().await?, vec![key; 5]); + + assert!(shutdown.send(()).is_ok(), "remote task exited"); + timeout(TEST_TIMEOUT, remote).await???; + drop(resumed); + second_relay.abort(); + let _ = second_relay.await; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn shutdown_interrupts_in_flight_registration() -> Result<()> { + let mut fixture = RegistryFixture::new(vec![503], Duration::from_secs(60)).await?; + let (shutdown, remote) = fixture.start()?; + assert_eq!( + timeout(TEST_TIMEOUT, fixture.requests.recv()).await?, + Some(1) + ); + assert!(shutdown.send(()).is_ok(), "remote task exited"); + timeout(Duration::from_secs(1), remote) + .await + .context("shutdown must not wait for the registration request timeout")???; + Ok(()) +} diff --git a/codex-rs/exec-server/tests/selected_capability_roots.rs b/codex-rs/exec-server/tests/selected_capability_roots.rs new file mode 100644 index 0000000000000000000000000000000000000000..3c621e1c82b6219e4e65453c303bc03561424774 --- /dev/null +++ b/codex-rs/exec-server/tests/selected_capability_roots.rs @@ -0,0 +1,69 @@ +#![cfg(unix)] + +mod common; + +use std::collections::HashMap; +use std::sync::Arc; + +use codex_exec_server_test_support::environment_manager_without_environments; +use codex_protocol::capabilities::CapabilityRootLocation; +use codex_protocol::capabilities::SelectedCapabilityRoot; +use codex_utils_path_uri::PathUri; +use common::exec_server::exec_server; +use pretty_assertions::assert_eq; + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn selected_capability_roots_use_captured_handle_after_replacement() -> anyhow::Result<()> { + let mut executor = exec_server().await?; + let manager = environment_manager_without_environments(); + let selected_root = SelectedCapabilityRoot { + id: "demo@1".to_string(), + location: CapabilityRootLocation::Environment { + environment_id: "tools".to_string(), + path: PathUri::parse("file:///plugins/demo")?, + }, + }; + + manager.upsert_environment( + "tools".to_string(), + executor.websocket_url().to_string(), + /*connect_timeout*/ None, + )?; + let environment_a = manager + .get_environment("tools") + .expect("executor A should be registered"); + environment_a.wait_until_ready().await?; + + let unavailable = manager + .resolve_selected_capability_roots( + std::slice::from_ref(&selected_root), + &HashMap::from([("tools".to_string(), None)]), + ) + .await; + assert!(unavailable.is_empty()); + + let captured_environments = + HashMap::from([("tools".to_string(), Some(Arc::clone(&environment_a)))]); + // Replace only the process-local handle; the stable environment ID and executor stay the same. + manager.upsert_environment( + "tools".to_string(), + executor.websocket_url().to_string(), + /*connect_timeout*/ None, + )?; + + let available = manager + .resolve_selected_capability_roots( + std::slice::from_ref(&selected_root), + &captured_environments, + ) + .await; + let [resolved] = available.as_slice() else { + anyhow::bail!("selected root should resolve through its stable environment"); + }; + + assert_eq!(resolved.selected_root(), &selected_root); + assert!(Arc::ptr_eq(resolved.environment(), &environment_a)); + + executor.shutdown().await?; + Ok(()) +} diff --git a/codex-rs/exec-server/tests/support/BUILD.bazel b/codex-rs/exec-server/tests/support/BUILD.bazel new file mode 100644 index 0000000000000000000000000000000000000000..6bb55d52be6bdbd81402cde94fef58902e5f19ef --- /dev/null +++ b/codex-rs/exec-server/tests/support/BUILD.bazel @@ -0,0 +1,8 @@ +load("//:defs.bzl", "codex_rust_crate") + +codex_rust_crate( + name = "support", + compile_data = ["//codex-rs/exec-server:src/proto/codex.exec_server.relay.v1.rs"], + crate_name = "codex_exec_server_test_support", + crate_srcs = glob(["*.rs"]), +) diff --git a/codex-rs/exec-server/tests/support/Cargo.toml b/codex-rs/exec-server/tests/support/Cargo.toml new file mode 100644 index 0000000000000000000000000000000000000000..073d4e341141ea59286b5563b4b6af1930d7ca51 --- /dev/null +++ b/codex-rs/exec-server/tests/support/Cargo.toml @@ -0,0 +1,24 @@ +[package] +name = "codex-exec-server-test-support" +version.workspace = true +edition.workspace = true +license.workspace = true + +[lib] +path = "lib.rs" +test = false +doctest = false + +[lints] +workspace = true + +[dependencies] +anyhow = { workspace = true } +codex-exec-server = { workspace = true } +codex-http-client = { workspace = true } +futures = { workspace = true } +prost = "0.14.3" +serde_json = { workspace = true } +tokio = { workspace = true, features = ["macros", "net", "time"] } +tokio-tungstenite = { workspace = true } +wiremock = { workspace = true } diff --git a/codex-rs/exec-server/tests/support/lib.rs b/codex-rs/exec-server/tests/support/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..bc29a2b17c8735c4b862ea118ba18f0a40313567 --- /dev/null +++ b/codex-rs/exec-server/tests/support/lib.rs @@ -0,0 +1,12 @@ +use codex_exec_server::EnvironmentManager; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; + +pub mod relay; + +/// Builds a manager without environments using the legacy outbound HTTP policy. +pub fn environment_manager_without_environments() -> EnvironmentManager { + EnvironmentManager::without_environments(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + )) +} diff --git a/codex-rs/exec-server/tests/support/relay.rs b/codex-rs/exec-server/tests/support/relay.rs new file mode 100644 index 0000000000000000000000000000000000000000..569f15dfb986675d652ab8eebb2cc357a91d6821 --- /dev/null +++ b/codex-rs/exec-server/tests/support/relay.rs @@ -0,0 +1,113 @@ +#[path = "../../src/proto/codex.exec_server.relay.v1.rs"] +mod relay_proto; + +use std::sync::Arc; +use std::sync::Mutex; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use codex_exec_server::NoiseChannelPublicKey; +use futures::SinkExt; +use futures::StreamExt; +use prost::Message as ProstMessage; +use relay_proto::RelayMessageFrame; +use relay_proto::relay_message_frame; +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 wiremock::MockServer; + +pub const TEST_TIMEOUT: Duration = Duration::from_secs(30); + +pub async fn accept_websocket( + listener: &TcpListener, + role: &str, +) -> Result> { + let (socket, _peer_addr) = timeout(TEST_TIMEOUT, listener.accept()) + .await + .with_context(|| format!("remote {role} should connect to fake rendezvous"))??; + timeout(TEST_TIMEOUT, accept_async(socket)) + .await + .with_context(|| format!("fake rendezvous should accept {role} websocket"))? + .map_err(Into::into) +} + +pub async fn registered_executor_public_key( + registry: &MockServer, +) -> Result { + let requests = registry + .received_requests() + .await + .context("wiremock should retain requests")?; + let request = requests + .iter() + .find(|request| request.url.path().ends_with("/register")) + .context("exec-server should register before connecting")?; + let body: serde_json::Value = serde_json::from_slice(&request.body)?; + let key = serde_json::from_value(body["executor_public_key"].clone())?; + Ok(key) +} + +pub async fn proxy_relay_frames( + mut environment: WebSocketStream, + mut harness: WebSocketStream, + captured_frames: Arc>>>, +) -> Result<()> { + loop { + tokio::select! { + message = environment.next() => { + let Some(message) = message else { + break; + }; + let message = message?; + capture_binary_frame(&captured_frames, &message); + harness.send(message).await?; + } + message = harness.next() => { + let Some(message) = message else { + break; + }; + let message = message?; + capture_binary_frame(&captured_frames, &message); + environment.send(message).await?; + } + } + } + Ok(()) +} + +fn capture_binary_frame(captured_frames: &Mutex>>, message: &Message) { + if let Message::Binary(bytes) = message { + captured_frames + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .push(bytes.to_vec()); + } +} + +pub fn assert_relay_data_is_encrypted(captured_frames: &Mutex>>) -> Result<()> { + let captured_frames = captured_frames + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let mut data_frames = 0; + for encoded in captured_frames.iter() { + let frame = RelayMessageFrame::decode(encoded.as_slice())?; + let Some(relay_message_frame::Body::Data(data)) = frame.body else { + continue; + }; + data_frames += 1; + let payload = String::from_utf8_lossy(&data.payload); + assert!(!payload.contains("initialize")); + assert!(!payload.contains("process/start")); + assert!(!payload.contains("noise-relay-test")); + } + assert!( + data_frames >= 4, + "expected encrypted request and response frames" + ); + Ok(()) +} diff --git a/codex-rs/exec-server/tests/unit/client_provisioning_tests.rs b/codex-rs/exec-server/tests/unit/client_provisioning_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..b4281c625ed207170870820836f1593622fc3d95 --- /dev/null +++ b/codex-rs/exec-server/tests/unit/client_provisioning_tests.rs @@ -0,0 +1,76 @@ +//! Exercises provisioning recovery while startup is publishing its previous failure. + +use super::*; +use crate::client_api::Deferred; +use crate::noise_channel::NoiseChannelIdentity; +use futures::poll; +use tokio::sync::oneshot; + +#[test_case::test_case(false; "initial startup")] +#[test_case::test_case(true; "subsequent attempt")] +#[tokio::test] +async fn caller_after_ready_retries_an_unpublished_provisioning_failure(reconnecting: bool) { + let (readiness, readiness_rx) = watch::channel(Some(Err("first failure".to_string()))); + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::Deferred(Box::new(Deferred { + readiness: readiness_rx, + transport: ExecServerTransportParams::NoiseRendezvous { + provider: Arc::new(FailingProvider), + identity: NoiseChannelIdentity::generate().expect("Noise identity"), + }, + })), + HttpClientFactory::new(codex_http_client::OutboundProxyPolicy::ReqwestDefault), + ); + let (failed_tx, failed_rx) = oneshot::channel(); + let (publish_tx, publish_rx) = oneshot::channel(); + let attempt = if reconnecting { + assert!(client.wait_until_ready().await.is_err()); + let attempt = Arc::new(ConnectionAttempt::default()); + *client.reconnect.lock().expect("reconnect lock") = Some(Arc::clone(&attempt)); + attempt + } else { + Arc::clone(&client.startup) + }; + let startup_client = client.clone(); + let startup = tokio::spawn(async move { + attempt + .result + .get_or_init(|| async { + let result = startup_client.connect_once(&attempt).await; + assert!(result.is_err()); + failed_tx.send(()).expect("failure observed"); + // Suspend at the publication boundary, as another executor thread could. + publish_rx.await.expect("publish startup result"); + result + }) + .await + .clone() + }); + failed_rx.await.expect("startup consumed the Failed report"); + readiness.send_replace(Some(Ok(()))); + + let mut after_ready = Box::pin(client.wait_until_ready()); + assert!(poll!(&mut after_ready).is_pending()); + publish_tx.send(()).expect("release startup"); + let error = after_ready.await.unwrap_err(); + assert!( + error.to_string().contains("provider reached after Ready"), + "post-Ready caller should retry the stale failure: {error}" + ); + assert!(startup.await.expect("startup task").is_err()); +} + +struct FailingProvider; + +impl crate::NoiseRendezvousConnectProvider for FailingProvider { + fn connect_bundle( + &self, + _: crate::NoiseChannelPublicKey, + ) -> BoxFuture<'_, Result> { + Box::pin(async { + Err(ExecServerError::Protocol( + "provider reached after Ready".to_string(), + )) + }) + } +} diff --git a/codex-rs/exec-server/tests/websocket.rs b/codex-rs/exec-server/tests/websocket.rs new file mode 100644 index 0000000000000000000000000000000000000000..d4b0fd8adc3d1f466aba874c0546d7afc72bf3ff --- /dev/null +++ b/codex-rs/exec-server/tests/websocket.rs @@ -0,0 +1,125 @@ +#![cfg(unix)] + +mod common; + +use codex_exec_server::InitializeParams; +use codex_exec_server::InitializeResponse; +use codex_exec_server_protocol::JSONRPCError; +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCResponse; +use common::exec_server::exec_server; +use pretty_assertions::assert_eq; +use tokio_tungstenite::connect_async; +use tokio_tungstenite::tungstenite::Error as WebSocketError; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::http::HeaderValue; +use tokio_tungstenite::tungstenite::http::StatusCode; +use tokio_tungstenite::tungstenite::http::header::ORIGIN; +use uuid::Uuid; + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_reports_malformed_websocket_json_and_keeps_running() -> anyhow::Result<()> { + let mut server = exec_server().await?; + server.send_raw_text("not-json").await?; + + let response = server + .wait_for_event(|event| matches!(event, JSONRPCMessage::Error(_))) + .await?; + let JSONRPCMessage::Error(JSONRPCError { id, error }) = response else { + panic!("expected malformed-message error response"); + }; + assert_eq!(id, codex_exec_server_protocol::RequestId::Integer(-1)); + assert_eq!(error.code, -32600); + assert!( + error + .message + .starts_with("failed to parse websocket JSON-RPC message from exec-server websocket"), + "unexpected malformed-message error: {}", + error.message + ); + + let initialize_id = server + .send_request( + "initialize", + serde_json::to_value(InitializeParams { + client_name: "exec-server-test".to_string(), + resume_session_id: None, + })?, + ) + .await?; + + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &initialize_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = response else { + panic!("expected initialize response after malformed input"); + }; + assert_eq!(id, initialize_id); + let initialize_response: InitializeResponse = serde_json::from_value(result)?; + Uuid::parse_str(&initialize_response.session_id)?; + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_accepts_binary_websocket_json() -> anyhow::Result<()> { + let mut server = exec_server().await?; + let initialize_id = codex_exec_server_protocol::RequestId::Integer(1); + let initialize = JSONRPCMessage::Request(codex_exec_server_protocol::JSONRPCRequest { + id: initialize_id.clone(), + method: "initialize".to_string(), + params: Some(serde_json::to_value(InitializeParams { + client_name: "exec-server-binary-test".to_string(), + resume_session_id: None, + })?), + trace: None, + }); + server + .send_raw_binary(serde_json::to_vec(&initialize)?) + .await?; + + let response = server + .wait_for_event(|event| { + matches!( + event, + JSONRPCMessage::Response(JSONRPCResponse { id, .. }) if id == &initialize_id + ) + }) + .await?; + let JSONRPCMessage::Response(JSONRPCResponse { id, result }) = response else { + panic!("expected initialize response for binary input"); + }; + assert_eq!(id, initialize_id); + let initialize_response: InitializeResponse = serde_json::from_value(result)?; + Uuid::parse_str(&initialize_response.session_id)?; + + server.shutdown().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn exec_server_rejects_browser_origin_websocket_handshake() -> anyhow::Result<()> { + let mut server = exec_server().await?; + let mut request = server.websocket_url().into_client_request()?; + request + .headers_mut() + .insert(ORIGIN, HeaderValue::from_static("https://evil.example")); + + let error = match connect_async(request).await { + Ok(_) => anyhow::bail!("browser-origin websocket handshake should be rejected"), + Err(error) => error, + }; + let WebSocketError::Http(response) = error else { + anyhow::bail!("browser-origin websocket handshake failed unexpectedly: {error}"); + }; + assert_eq!(response.status(), StatusCode::FORBIDDEN); + + server.shutdown().await?; + Ok(()) +} diff --git a/codex-rs/ext/agent/BUILD.bazel b/codex-rs/ext/agent/BUILD.bazel new file mode 100644 index 0000000000000000000000000000000000000000..21793fc95778e085127231f5bdb945cf61fb46d5 --- /dev/null +++ b/codex-rs/ext/agent/BUILD.bazel @@ -0,0 +1,6 @@ +load("//:defs.bzl", "codex_rust_crate") + +codex_rust_crate( + name = "agent", + crate_name = "codex_agent_extension", +) diff --git a/codex-rs/ext/agent/Cargo.toml b/codex-rs/ext/agent/Cargo.toml new file mode 100644 index 0000000000000000000000000000000000000000..6c3f10a43319a122e2da129e31ba34aa4f1df49a --- /dev/null +++ b/codex-rs/ext/agent/Cargo.toml @@ -0,0 +1,24 @@ +[package] +edition.workspace = true +license.workspace = true +name = "codex-agent-extension" +version.workspace = true + +[lib] +name = "codex_agent_extension" +path = "src/lib.rs" +doctest = false +test = false + +[lints] +workspace = true + +[dependencies] +codex-core = { workspace = true } +codex-protocol = { workspace = true } + +[dev-dependencies] +anyhow = { workspace = true } +core_test_support = { workspace = true } +pretty_assertions = { workspace = true } +tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } diff --git a/codex-rs/ext/agent/src/lib.rs b/codex-rs/ext/agent/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..bcd8fbb69ec5126ead8e663e95b9f2ce66e97da8 --- /dev/null +++ b/codex-rs/ext/agent/src/lib.rs @@ -0,0 +1,100 @@ +use codex_core::CodexThread; +use codex_core::NewThread; +use codex_core::StartIfIdleSubmission; +use codex_core::StartThreadOptions; +use codex_core::ThreadManager; +use codex_core::TurnInputRequest; +use codex_core::config::Config; +use codex_protocol::ThreadId; +use codex_protocol::error::CodexErr; +use codex_protocol::error::Result as CodexResult; +use codex_protocol::protocol::W3cTraceContext; +use codex_protocol::user_input::UserInput; +use std::sync::Arc; +use std::sync::Weak; + +/// A fully resolved agent invocation. +/// +/// Agent discovery owns rendering `prompt`, including any selected skill +/// references. The runtime only starts that prompt in isolated forked context. +pub struct AgentInvocation { + pub config: Config, + pub prompt: String, + pub parent_trace: Option, +} + +/// A spawned agent whose initial turn has been submitted. +pub struct AgentRun { + pub thread_id: ThreadId, + pub turn_id: String, + pub thread: Arc, +} + +/// Runs resolved agents in threads forked by the owning [`ThreadManager`]. +#[derive(Clone)] +pub struct AgentRunner { + thread_manager: Weak, +} + +impl AgentRunner { + pub fn new(thread_manager: Weak) -> Self { + Self { thread_manager } + } + + /// Starts a resolved agent in a fork of `parent_thread_id`. + pub async fn start( + &self, + parent_thread_id: ThreadId, + invocation: AgentInvocation, + ) -> CodexResult { + let AgentInvocation { + config, + prompt, + parent_trace, + } = invocation; + if prompt.trim().is_empty() { + return Err(CodexErr::InvalidRequest( + "agent prompt must not be empty".to_string(), + )); + } + + let thread_manager = self + .thread_manager + .upgrade() + .ok_or_else(|| CodexErr::UnsupportedOperation("thread manager dropped".to_string()))?; + let NewThread { + thread_id, thread, .. + } = thread_manager + .spawn_subagent( + parent_thread_id, + StartThreadOptions { + parent_trace: parent_trace.clone(), + ..StartThreadOptions::new(config) + }, + ) + .await?; + let turn_id = match thread + .start_turn_if_idle( + TurnInputRequest::user_input(vec![UserInput::Text { + text: prompt, + text_elements: Vec::new(), + }]) + .with_trace(parent_trace), + ) + .await? + { + StartIfIdleSubmission::Started { turn_id } => turn_id, + StartIfIdleSubmission::NotSubmitted { reason } => { + return Err(CodexErr::InvalidRequest(format!( + "agent prompt was not submitted: {reason:?}" + ))); + } + }; + + Ok(AgentRun { + thread_id, + turn_id, + thread, + }) + } +} diff --git a/codex-rs/ext/agent/tests/agent_service.rs b/codex-rs/ext/agent/tests/agent_service.rs new file mode 100644 index 0000000000000000000000000000000000000000..9b36d4a69bd1542e3056670c84d5de5311a6d9f9 --- /dev/null +++ b/codex-rs/ext/agent/tests/agent_service.rs @@ -0,0 +1,70 @@ +use anyhow::Result; +use codex_agent_extension::AgentInvocation; +use codex_agent_extension::AgentRunner; +use codex_protocol::protocol::EventMsg; +use core_test_support::responses; +use core_test_support::skip_if_no_network; +use core_test_support::test_codex::test_codex; +use core_test_support::wait_for_event; +use pretty_assertions::assert_eq; + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn starts_resolved_agent_prompt_in_forked_thread() -> Result<()> { + skip_if_no_network!(Ok(())); + + let server = responses::start_mock_server().await; + let response_mock = responses::mount_sse_once( + &server, + responses::sse(vec![ + responses::ev_response_created("agent-response"), + responses::ev_completed("agent-response"), + ]), + ) + .await; + let test = test_codex().build_with_auto_env(&server).await?; + let parent_thread_id = test.session_configured.session_id.into(); + let agent_runner = AgentRunner::new(std::sync::Arc::downgrade(&test.thread_manager)); + + let agent_run = agent_runner + .start( + parent_thread_id, + AgentInvocation { + config: test.config.clone(), + prompt: "Use $example-agent to inspect the current changes.".to_string(), + parent_trace: None, + }, + ) + .await?; + + assert_ne!(agent_run.thread_id, parent_thread_id); + assert_eq!( + agent_run + .thread + .config_snapshot() + .await + .forked_from_thread_id, + Some(parent_thread_id) + ); + let started = wait_for_event(&agent_run.thread, |event| { + matches!(event, EventMsg::TurnStarted(_)) + }) + .await; + let EventMsg::TurnStarted(started) = started else { + unreachable!("event predicate only matches turn started events"); + }; + assert_eq!(started.turn_id, agent_run.turn_id); + wait_for_event(&agent_run.thread, |event| { + matches!(event, EventMsg::TurnComplete(_)) + }) + .await; + + let request = response_mock.single_request(); + assert!( + request + .message_input_texts("user") + .iter() + .any(|text| text == "Use $example-agent to inspect the current changes.") + ); + + Ok(()) +} diff --git a/codex-rs/ext/connectors/BUILD.bazel b/codex-rs/ext/connectors/BUILD.bazel new file mode 100644 index 0000000000000000000000000000000000000000..304349b8a17c3709673fa15143cebd4c3a011669 --- /dev/null +++ b/codex-rs/ext/connectors/BUILD.bazel @@ -0,0 +1,6 @@ +load("//:defs.bzl", "codex_rust_crate") + +codex_rust_crate( + name = "connectors", + crate_name = "codex_connectors_extension", +) diff --git a/codex-rs/ext/connectors/Cargo.toml b/codex-rs/ext/connectors/Cargo.toml new file mode 100644 index 0000000000000000000000000000000000000000..1aead727731475bc9590ac90efb37c2ae7adef64 --- /dev/null +++ b/codex-rs/ext/connectors/Cargo.toml @@ -0,0 +1,24 @@ +[package] +edition.workspace = true +license.workspace = true +name = "codex-connectors-extension" +version.workspace = true + +[lib] +name = "codex_connectors_extension" +path = "src/lib.rs" +doctest = false +test = false + +[lints] +workspace = true + +[dependencies] +codex-connectors = { workspace = true } +codex-core-plugins = { workspace = true } +codex-file-system = { workspace = true } +codex-plugin = { workspace = true } +codex-utils-path-uri = { workspace = true } +serde_json = { workspace = true } +thiserror = { workspace = true } +tracing = { workspace = true } diff --git a/codex-rs/ext/connectors/src/executor_plugin.rs b/codex-rs/ext/connectors/src/executor_plugin.rs new file mode 100644 index 0000000000000000000000000000000000000000..c8d51e80ecac7e7fdaba7ee4ab76c48fa477ca0f --- /dev/null +++ b/codex-rs/ext/connectors/src/executor_plugin.rs @@ -0,0 +1,70 @@ +use codex_connectors::parse_plugin_app_config; +use codex_core_plugins::ResolvedExecutorPlugin; +use codex_file_system::ReadFileOptions; +use codex_plugin::AppDeclaration; +use codex_plugin::PluginResourceLocator; +use codex_utils_path_uri::PathUri; +use std::io; +use thiserror::Error; + +/// Loads connector declarations from a resolved plugin through its owning executor. +#[derive(Clone, Copy, Debug, Default)] +pub struct ExecutorPluginConnectorProvider; + +/// Failure to load connector declarations from an executor plugin. +#[derive(Debug, Error)] +pub enum ExecutorPluginConnectorProviderError { + #[error("failed to read app config for selected plugin `{plugin_id}` at `{path}`: {source}")] + ReadConfig { + plugin_id: String, + path: PathUri, + #[source] + source: io::Error, + }, + #[error("failed to parse app config for selected plugin `{plugin_id}` at `{path}`: {source}")] + ParseConfig { + plugin_id: String, + path: PathUri, + #[source] + source: serde_json::Error, + }, +} + +impl ExecutorPluginConnectorProvider { + /// Returns the connector declarations contributed by `plugin`. + #[tracing::instrument(name = "connectors.executor_plugin.declarations.load", skip_all)] + pub async fn load( + &self, + plugin: &ResolvedExecutorPlugin, + ) -> Result, ExecutorPluginConnectorProviderError> { + let resolved_plugin = plugin.plugin(); + let plugin_id = resolved_plugin.selected_root_id(); + let Some(PluginResourceLocator::Environment { + path: config_path, .. + }) = resolved_plugin.manifest().paths.apps.as_ref() + else { + return Ok(Vec::new()); + }; + let contents = plugin + .file_system() + .read_file_text( + config_path, + ReadFileOptions::default(), + /*sandbox*/ None, + ) + .await + .map_err(|source| ExecutorPluginConnectorProviderError::ReadConfig { + plugin_id: plugin_id.to_string(), + path: config_path.clone(), + source, + })?; + + parse_plugin_app_config(&contents).map_err(|source| { + ExecutorPluginConnectorProviderError::ParseConfig { + plugin_id: plugin_id.to_string(), + path: config_path.clone(), + source, + } + }) + } +} diff --git a/codex-rs/ext/connectors/src/lib.rs b/codex-rs/ext/connectors/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..f60e5f916a39b05335331f13e227ef1227ea001d --- /dev/null +++ b/codex-rs/ext/connectors/src/lib.rs @@ -0,0 +1,6 @@ +//! Executor-backed connector declaration loading. + +mod executor_plugin; + +pub use executor_plugin::ExecutorPluginConnectorProvider; +pub use executor_plugin::ExecutorPluginConnectorProviderError; diff --git a/codex-rs/ext/extension-api/BUILD.bazel b/codex-rs/ext/extension-api/BUILD.bazel new file mode 100644 index 0000000000000000000000000000000000000000..c79ad601a4c73efed4d58effb63838f1912967b0 --- /dev/null +++ b/codex-rs/ext/extension-api/BUILD.bazel @@ -0,0 +1,6 @@ +load("//:defs.bzl", "codex_rust_crate") + +codex_rust_crate( + name = "extension-api", + crate_name = "codex_extension_api", +) diff --git a/codex-rs/ext/extension-api/Cargo.toml b/codex-rs/ext/extension-api/Cargo.toml new file mode 100644 index 0000000000000000000000000000000000000000..a0d4fe1f4c52c95ec34aed8000c8a6fe2df5feee --- /dev/null +++ b/codex-rs/ext/extension-api/Cargo.toml @@ -0,0 +1,30 @@ +[package] +edition.workspace = true +license.workspace = true +name = "codex-extension-api" +version.workspace = true + +[lib] +name = "codex_extension_api" +path = "src/lib.rs" +test = false +doctest = false + +[lints] +workspace = true + +[dependencies] +codex-history = { workspace = true } +codex-config = { workspace = true } +codex-context-fragments = { workspace = true } +codex-exec-server-protocol = { workspace = true } +codex-mcp = { workspace = true } +codex-protocol = { workspace = true } +codex-tools = { workspace = true } +codex-utils-absolute-path = { workspace = true } +codex-utils-path-uri = { workspace = true } +serde_json = { workspace = true } + +[dev-dependencies] +pretty_assertions = { workspace = true } +tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } diff --git a/codex-rs/ext/extension-api/examples/enabled_extensions.rs b/codex-rs/ext/extension-api/examples/enabled_extensions.rs new file mode 100644 index 0000000000000000000000000000000000000000..45b178bc813128ad1cad72f4875d7860b8ed2008 --- /dev/null +++ b/codex-rs/ext/extension-api/examples/enabled_extensions.rs @@ -0,0 +1,97 @@ +#[path = "enabled_extensions/shared_state_extension.rs"] +mod shared_state_extension; + +use std::future::Future; +use std::pin::pin; +use std::task::Context; +use std::task::Poll; +use std::task::Waker; + +use codex_extension_api::ExtensionData; +use codex_extension_api::ExtensionRegistryBuilder; +use shared_state_extension::recorded_style_contributions; +use shared_state_extension::recorded_usage_contributions; + +fn main() { + // 1. Install the contributors for the thread-start input type this host exposes. + let mut builder = ExtensionRegistryBuilder::<()>::new(); + shared_state_extension::install(&mut builder); + let registry = builder.build(); + + // 2. The host decides which stores are shared. + let session_store = ExtensionData::new("session"); + let first_thread_store = ExtensionData::new("thread-1"); + let second_thread_store = ExtensionData::new("thread-2"); + + // 3. Reusing the same session store shares session state across threads. + let first_thread_fragments = block_on_ready(contribute_prompt( + ®istry, + &session_store, + &first_thread_store, + )); + block_on_ready(contribute_prompt( + ®istry, + &session_store, + &first_thread_store, + )); + block_on_ready(contribute_prompt( + ®istry, + &session_store, + &second_thread_store, + )); + + println!("first prompt fragments: {}", first_thread_fragments.len()); + println!( + "session style contributions: {}", + recorded_style_contributions(&session_store) + ); + println!( + "session usage contributions: {}", + recorded_usage_contributions(&session_store) + ); + println!( + "first thread style contributions: {}", + recorded_style_contributions(&first_thread_store) + ); + println!( + "first thread usage contributions: {}", + recorded_usage_contributions(&first_thread_store) + ); + println!( + "second thread style contributions: {}", + recorded_style_contributions(&second_thread_store) + ); + println!( + "second thread usage contributions: {}", + recorded_usage_contributions(&second_thread_store) + ); +} + +async fn contribute_prompt( + registry: &codex_extension_api::ExtensionRegistry<()>, + session_store: &ExtensionData, + thread_store: &ExtensionData, +) -> Vec { + let mut fragments = Vec::new(); + for contributor in registry.context_contributors() { + fragments.extend( + contributor + .contribute_thread_context(session_store, thread_store) + .await, + ); + } + fragments +} + +fn block_on_ready(future: F) -> F::Output +where + F: Future, +{ + let waker = Waker::noop(); + let mut context = Context::from_waker(waker); + let mut future = pin!(future); + match future.as_mut().poll(&mut context) { + Poll::Ready(output) => output, + Poll::Pending => panic!("example context contributors should complete immediately"), + } +} diff --git a/codex-rs/ext/extension-api/examples/enabled_extensions/shared_state_extension.rs b/codex-rs/ext/extension-api/examples/enabled_extensions/shared_state_extension.rs new file mode 100644 index 0000000000000000000000000000000000000000..0657f79940794440302ed4f31fd10171cb63b803 --- /dev/null +++ b/codex-rs/ext/extension-api/examples/enabled_extensions/shared_state_extension.rs @@ -0,0 +1,101 @@ +use std::sync::Arc; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; + +use codex_extension_api::ContentItemKind; +use codex_extension_api::ContextContributor; +use codex_extension_api::ExtensionData; +use codex_extension_api::ExtensionRegistryBuilder; +use codex_extension_api::PromptFragment; + +/// Installs the tutorial contributors used by the example host. +pub fn install(registry: &mut ExtensionRegistryBuilder<()>) { + registry.prompt_contributor(Arc::new(StyleContributor)); + registry.prompt_contributor(Arc::new(UsageContributor)); +} + +#[derive(Debug)] +struct StyleContributor; + +impl ContextContributor for StyleContributor { + fn contribute_thread_context<'a>( + &'a self, + session_store: &'a ExtensionData, + thread_store: &'a ExtensionData, + ) -> std::pin::Pin> + Send + 'a>> { + Box::pin(async move { + contribution_counts(session_store).record_style(); + contribution_counts(thread_store).record_style(); + + vec![PromptFragment::developer_policy( + "Prefer short answers unless the user asks for detail.", + ContentItemKind("example.style_instructions".to_string()), + )] + }) + } +} + +#[derive(Debug)] +struct UsageContributor; + +impl ContextContributor for UsageContributor { + fn contribute_thread_context<'a>( + &'a self, + session_store: &'a ExtensionData, + thread_store: &'a ExtensionData, + ) -> std::pin::Pin> + Send + 'a>> { + Box::pin(async move { + contribution_counts(session_store).record_usage(); + contribution_counts(thread_store).record_usage(); + + vec![PromptFragment::developer_capability( + "This extension can contribute more than one prompt fragment.", + ContentItemKind("example.usage_instructions".to_string()), + )] + }) + } +} + +/// Returns how many style contributions were recorded in `store`. +pub fn recorded_style_contributions(store: &ExtensionData) -> u64 { + store + .get::() + .map(|counts| counts.style()) + .unwrap_or_default() +} + +/// Returns how many usage contributions were recorded in `store`. +pub fn recorded_usage_contributions(store: &ExtensionData) -> u64 { + store + .get::() + .map(|counts| counts.usage()) + .unwrap_or_default() +} + +#[derive(Debug, Default)] +struct ContributionCounts { + style: AtomicU64, + usage: AtomicU64, +} + +impl ContributionCounts { + fn record_style(&self) { + self.style.fetch_add(1, Ordering::Relaxed); + } + + fn record_usage(&self) { + self.usage.fetch_add(1, Ordering::Relaxed); + } + + fn style(&self) -> u64 { + self.style.load(Ordering::Relaxed) + } + + fn usage(&self) -> u64 { + self.usage.load(Ordering::Relaxed) + } +} + +fn contribution_counts(store: &ExtensionData) -> Arc { + store.get_or_init::(Default::default) +} diff --git a/codex-rs/ext/extension-api/notes.md b/codex-rs/ext/extension-api/notes.md new file mode 100644 index 0000000000000000000000000000000000000000..e73b106f6af0414a5121e7b03564602cf5904f2a --- /dev/null +++ b/codex-rs/ext/extension-api/notes.md @@ -0,0 +1,14 @@ +Everything becomes a good contributor design, which contributors do we need? + +git attribution Context +memories Context + Tool + Output +guardian Context + Request +goal Tool + Runtime +image generation Tool + Output +skills Context + Turn +personality Context +plugins / apps / connectors Context + Turn +shell snapshot Runtime +web search Tool +AGENTS.md Context (Runtime too only if you want eager refresh/cache behavior) +future sandboxing probably Request + Runtime diff --git a/codex-rs/ext/extension-api/src/allowed_tools.rs b/codex-rs/ext/extension-api/src/allowed_tools.rs new file mode 100644 index 0000000000000000000000000000000000000000..25abab2244e5a713737f863d3ea8fbef74de7578 --- /dev/null +++ b/codex-rs/ext/extension-api/src/allowed_tools.rs @@ -0,0 +1,24 @@ +//! A startup ceiling on the tools a thread may advertise or execute. + +use crate::ToolName; + +/// Supply through `ExtensionDataInit` before starting a thread. The host captures +/// this value once; changing extension state later cannot widen the tool set. +/// Callers must supply it again when resuming a thread. +/// +/// An absent value keeps ordinary tool setup. An empty list permits no tools. +/// Names include their namespace; a plain name uses the default namespace. +/// This only removes tools: configuration and permission checks still apply. +/// Generated tools such as Code Mode's `exec` and `wait` must also be listed. +#[derive(Clone, Debug, Default)] +pub struct AllowedTools(pub Vec); + +impl AllowedTools { + pub fn contains(&self, tool: &ToolName) -> bool { + self.0.iter().any(|allowed| { + allowed.name == tool.name + && (allowed.namespace == tool.namespace + || (allowed.is_default_namespace() && tool.is_default_namespace())) + }) + } +} diff --git a/codex-rs/ext/extension-api/src/capabilities/conversation_history.rs b/codex-rs/ext/extension-api/src/capabilities/conversation_history.rs new file mode 100644 index 0000000000000000000000000000000000000000..68d7783e2d91bce1f8851c2027b2e36e9322a748 --- /dev/null +++ b/codex-rs/ext/extension-api/src/capabilities/conversation_history.rs @@ -0,0 +1,48 @@ +use codex_history::RetainedContext; + +use codex_protocol::models::ResponseItem; + +/// Read-only conversation-history snapshot supplied by the extension host. +/// +/// Implementations should retain the host's existing snapshot storage rather than +/// copying response payloads into an extension-owned collection. +pub trait ConversationHistorySnapshot: Send + Sync { + /// Returns the generation of the history captured by this snapshot. + fn history_version(&self) -> u64; + + /// Host-owned revision captured with this snapshot. Advances on user messages and + /// history resets, but stays unchanged for compaction and internal context. + fn user_message_revision(&self) -> u64; + + /// Returns the snapshot's response items in conversation order. + fn items(&self) -> Box + Send + '_>; + + /// Host-owned retained facts captured atomically with the parent model window. + /// These facts may be available while review still uses a legacy transcript. + fn retained_context(&self) -> Option<&RetainedContext> { + None + } + + /// Whether review uses the parent checkpoint and model window instead of a legacy transcript. + /// Checkpoint compatibility is independent of access to retained user evidence. + fn uses_parent_context_for_review(&self) -> bool { + self.retained_context().is_some() + } + + /// Producer compatibility recorded on the latest opaque checkpoint. Missing provenance + /// must not be inferred from the currently selected model, including after resume. + fn latest_compaction_model_hash(&self) -> Option<&str> { + None + } + + /// Original review evidence retained across parent compaction, in conversation order. + /// Hosts without separate retention provide their current history. + fn review_items(&self) -> Box + Send + '_> { + self.items() + } + + /// Changes whenever offsets into the retained review evidence become invalid. + fn review_history_version(&self) -> u64 { + self.history_version() + } +} diff --git a/codex-rs/ext/extension-api/src/capabilities/events.rs b/codex-rs/ext/extension-api/src/capabilities/events.rs new file mode 100644 index 0000000000000000000000000000000000000000..c19624a1704a7a91d486714eb52b5762a53863c2 --- /dev/null +++ b/codex-rs/ext/extension-api/src/capabilities/events.rs @@ -0,0 +1,38 @@ +use codex_protocol::protocol::Event; + +/// Extension warning with an explicit thread target and optional turn correlation. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ExtensionWarning { + /// Stable host-owned thread identifier used for delivery. + pub thread_id: String, + /// Stable host-owned turn identifier when the warning arose in a turn callback. + pub turn_id: Option, + /// Concise warning message for the user. + pub message: String, +} + +/// Host-provided fire-and-forget sink for extension-generated events. +/// +/// Extensions construct protocol events with the correlation id appropriate for +/// the callback they are handling, then leave persistence, ordering, transport +/// fanout, and logging decisions to the host. +pub trait ExtensionEventSink: Send + Sync { + /// Queue one protocol event for host-owned delivery. + fn emit(&self, event: Event); + + /// Queue one warning for host-owned delivery. + /// + /// Implementations must use [`ExtensionWarning::thread_id`] for routing. The optional + /// [`ExtensionWarning::turn_id`] is correlation metadata and does not identify a thread. + fn emit_warning(&self, warning: ExtensionWarning); +} + +/// Event sink used when the host does not expose extension event emission. +#[derive(Debug, Default, Clone, Copy)] +pub struct NoopExtensionEventSink; + +impl ExtensionEventSink for NoopExtensionEventSink { + fn emit(&self, _event: Event) {} + + fn emit_warning(&self, _warning: ExtensionWarning) {} +} diff --git a/codex-rs/ext/extension-api/src/capabilities/metrics.rs b/codex-rs/ext/extension-api/src/capabilities/metrics.rs new file mode 100644 index 0000000000000000000000000000000000000000..67f395b48059c1b897f16106da7bb8922d8d0ac6 --- /dev/null +++ b/codex-rs/ext/extension-api/src/capabilities/metrics.rs @@ -0,0 +1,21 @@ +/// Host-provided metrics capability for extension-owned behavior. +/// +/// Implementations are expected to attach the host's session attribution before +/// forwarding samples to the configured metrics backend. +pub trait ExtensionMetrics: Send + Sync { + /// Increments a counter with optional extension-provided tags. + fn counter(&self, name: &str, inc: i64, tags: &[(&str, &str)]); + + /// Records one histogram sample with optional extension-provided tags. + fn histogram(&self, name: &str, value: i64, tags: &[(&str, &str)]); + + /// Records a histogram with explicit buckets, preserving host attribution. + /// All callers of the same metric name must use the same boundaries. + fn histogram_with_boundaries( + &self, + name: &str, + value: i64, + boundaries: &[f64], + tags: &[(&str, &str)], + ); +} diff --git a/codex-rs/ext/extension-api/src/capabilities/mod.rs b/codex-rs/ext/extension-api/src/capabilities/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..d933738616189c8e8f5a65cf9d2336f18bd7d1f8 --- /dev/null +++ b/codex-rs/ext/extension-api/src/capabilities/mod.rs @@ -0,0 +1,13 @@ +mod conversation_history; +mod events; +mod metrics; +mod response_items; + +pub use conversation_history::ConversationHistorySnapshot; +pub use events::ExtensionEventSink; +pub use events::ExtensionWarning; +pub use events::NoopExtensionEventSink; +pub use metrics::ExtensionMetrics; +pub use response_items::NoopResponseItemInjector; +pub use response_items::ResponseItemInjectionFuture; +pub use response_items::ResponseItemInjector; diff --git a/codex-rs/ext/extension-api/src/capabilities/response_items.rs b/codex-rs/ext/extension-api/src/capabilities/response_items.rs new file mode 100644 index 0000000000000000000000000000000000000000..6c300e2bf14b4cab34ee60ee82fac0d2df485408 --- /dev/null +++ b/codex-rs/ext/extension-api/src/capabilities/response_items.rs @@ -0,0 +1,33 @@ +use std::future::Future; +use std::pin::Pin; + +use codex_protocol::models::ResponseInputItem; + +/// Future returned when an extension asks the host to inject model-visible input. +pub type ResponseItemInjectionFuture<'a> = + Pin>> + Send + 'a>>; + +/// Host-provided helper for extensions that need to steer the active model turn. +/// +/// Implementations should inject the supplied response items into the active turn +/// when one can accept same-turn model input. If injection is unavailable, they +/// return the unchanged items to the caller. +pub trait ResponseItemInjector: Send + Sync { + fn inject_response_items<'a>( + &'a self, + items: Vec, + ) -> ResponseItemInjectionFuture<'a>; +} + +/// Injector used when a host does not expose same-turn model steering. +#[derive(Debug, Default, Clone, Copy)] +pub struct NoopResponseItemInjector; + +impl ResponseItemInjector for NoopResponseItemInjector { + fn inject_response_items<'a>( + &'a self, + items: Vec, + ) -> ResponseItemInjectionFuture<'a> { + Box::pin(std::future::ready(Err(items))) + } +} diff --git a/codex-rs/ext/extension-api/src/contributors.rs b/codex-rs/ext/extension-api/src/contributors.rs new file mode 100644 index 0000000000000000000000000000000000000000..ad102973baa9c0bbabf421f1406d720d19d2b75a --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors.rs @@ -0,0 +1,383 @@ +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; + +use codex_context_fragments::ContextualUserFragment; +use codex_protocol::items::TurnItem; +use codex_protocol::protocol::TokenUsageInfo; +use codex_tools::ToolCall; +use codex_tools::ToolExecutor; + +use crate::ExtensionData; +use crate::ExtensionMetrics; + +mod approval_review; +mod context; +mod mcp; +mod prompt; +mod skill_invocation; +mod thread_lifecycle; +mod tool_lifecycle; +mod turn_input; +mod turn_lifecycle; +mod world_state; + +pub use approval_review::ApprovalDecision; +pub use approval_review::ApprovalDecisionInput; +pub use approval_review::GuardianV2Enabled; +pub use approval_review::SynchronousApprovalReviewer; +pub use context::TurnContextContributionInput; +pub use mcp::McpServerContribution; +pub use mcp::McpServerContributionContext; +pub use mcp::SelectedPluginIdentity; +pub use mcp::SelectedPluginSnapshot; +pub use prompt::PromptFragment; +pub use prompt::PromptSlot; +pub use skill_invocation::SkillInvocationInput; +pub use skill_invocation::SkillInvocationKind; +pub use thread_lifecycle::ThreadIdleCause; +pub use thread_lifecycle::ThreadIdleInput; +pub use thread_lifecycle::ThreadOriginator; +pub use thread_lifecycle::ThreadReadyInput; +pub use thread_lifecycle::ThreadResumeInput; +pub use thread_lifecycle::ThreadStartInput; +pub use thread_lifecycle::ThreadStopInput; +pub use tool_lifecycle::McpToolContext; +pub use tool_lifecycle::McpToolResultInput; +pub use tool_lifecycle::McpToolSource; +pub use tool_lifecycle::ToolCallOutcome; +pub use tool_lifecycle::ToolFinishInput; +pub use tool_lifecycle::ToolLifecycleFuture; +pub use tool_lifecycle::ToolStartInput; +pub use turn_input::TurnInputContext; +pub use turn_input::TurnInputEnvironment; +pub use turn_lifecycle::TurnAbortInput; +pub use turn_lifecycle::TurnErrorInput; +pub use turn_lifecycle::TurnStartInput; +pub use turn_lifecycle::TurnStopInput; +pub use world_state::PreviousWorldStateSection; +pub use world_state::RenderedWorldStateFragment; +pub use world_state::WorldStateContributionInput; +pub use world_state::WorldStateSectionContribution; + +/// Boxed, sendable future returned by asynchronous extension contributors. +pub type ExtensionFuture<'a, T> = Pin + Send + 'a>>; + +/// Extension contribution that resolves runtime MCP servers from host config. +/// +/// Contributors run in registration order. Later contributions for the same +/// name replace earlier ones. Implementations must contribute only names they +/// own and must apply any source-specific policy before returning a server. +/// Thread-scoped resolution exposes the host-seeded thread inputs; global +/// resolution exposes none and must not imply a local fallback. Thread inputs +/// are frozen for the runtime and do not include lifecycle-contributor state. +/// Auto-discovered plugin servers are resolved by the plugin manager. A +/// thread-selected plugin contribution must carry its own package provenance. +pub trait McpServerContributor: Send + Sync { + /// Stable identity used for registration provenance and conflict diagnostics. + fn id(&self) -> &'static str; + + fn contribute<'a>( + &'a self, + context: McpServerContributionContext<'a, C>, + ) -> ExtensionFuture<'a, Vec>; +} + +/// Extension contribution that adds prompt fragments during prompt assembly. +/// +/// Implementations should use the method matching the scope needed by the +/// fragment: thread/session context for stable inputs, and turn context for +/// fragments that depend on turn-local host state. +pub trait ContextContributor: Send + Sync { + /// Returns thread-scoped context using the supplied extension state. + fn contribute_thread_context<'a>( + &'a self, + session_store: &'a ExtensionData, + thread_store: &'a ExtensionData, + ) -> ExtensionFuture<'a, Vec> { + Box::pin(async move { + let _self = self; + let _session_store = session_store; + let _thread_store = thread_store; + Vec::new() + }) + } + + fn contribute_turn_context<'a>( + &'a self, + input: TurnContextContributionInput<'a>, + ) -> ExtensionFuture<'a, Vec> { + Box::pin(async move { + let _self = self; + let _input = input; + Vec::new() + }) + } + + fn contribute_world_state<'a>( + &'a self, + input: WorldStateContributionInput<'a>, + ) -> ExtensionFuture<'a, Vec> { + Box::pin(async move { + let _self = self; + let _input = input; + Vec::new() + }) + } +} + +/// Contributor for host-owned thread lifecycle gates. +/// +/// Implementations should use these callbacks to seed, rehydrate, or flush +/// extension-private thread state and retain any session capabilities supplied +/// by the host. Other heavy dependencies belong on the extension value. +pub trait ThreadLifecycleContributor: Send + Sync { + /// Called after host startup has initialized the thread-scoped store. + fn on_thread_start<'a>(&'a self, input: ThreadStartInput<'a, C>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _input = input; + }) + } + + /// Called after the initialized thread is registered with its host. + fn on_thread_ready<'a>(&'a self, input: ThreadReadyInput<'a, C>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _input = input; + }) + } + + /// Called after the host constructs a runtime from persisted history. + fn on_thread_resume<'a>(&'a self, input: ThreadResumeInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _input = input; + }) + } + + /// Called after the host has drained immediately pending thread work. + /// + /// Implementations may use host capabilities captured by the extension to + /// submit follow-up input. The host remains responsible for deciding + /// whether that input starts a turn, is queued, or is ignored. + fn on_thread_idle<'a>(&'a self, input: ThreadIdleInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _input = input; + }) + } + + /// Called during runtime teardown, before the host closes persistent history. + /// Contributors must cancel and join their background work before returning. + fn on_thread_stop<'a>(&'a self, input: ThreadStopInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _input = input; + }) + } +} + +/// Contributor for host-owned turn lifecycle gates. +/// +/// Implementations should use these callbacks to seed, observe, or clear +/// extension-private turn state. The host exposes stable identifiers and +/// extension stores instead of core runtime objects. +pub trait TurnLifecycleContributor: Send + Sync { + /// Called after turn-scoped extension stores are created, before the task + /// for the turn starts running. + fn on_turn_start<'a>(&'a self, input: TurnStartInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _input = input; + }) + } + + /// Observes a completed item without changing it or delaying streamed deltas. + fn on_item_completed<'a>( + &'a self, + _thread_store: &'a ExtensionData, + _turn_store: &'a ExtensionData, + _item: &'a TurnItem, + ) -> ExtensionFuture<'a, ()> { + Box::pin(std::future::ready(())) + } + + /// Called before the host drops the completed turn runtime and turn store. + fn on_turn_stop<'a>(&'a self, input: TurnStopInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _input = input; + }) + } + + /// Called after the host aborts a running turn. + fn on_turn_abort<'a>(&'a self, input: TurnAbortInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _input = input; + }) + } + + /// Called when the host observes an error for a running turn. + fn on_turn_error<'a>(&'a self, input: TurnErrorInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _input = input; + }) + } +} + +/// Extension contribution that can add turn-local model input. +/// +/// Implementations should resolve only the model-visible input they own and +/// must preserve authority boundaries for external resources. Expensive or +/// host-specific dependencies belong on the extension value installed by the +/// host, not in this input. +pub trait TurnInputContributor: Send + Sync { + /// Returns additional contextual fragments for one submitted turn. The optional metrics + /// capability is bound to the effective model for that turn. + fn contribute<'a>( + &'a self, + input: TurnInputContext<'a>, + extension_metrics: Option>, + session_store: &'a ExtensionData, + thread_store: &'a ExtensionData, + turn_store: &'a ExtensionData, + ) -> ExtensionFuture<'a, Vec>>; +} + +/// Contributor for host-owned configuration changes. +/// +/// Implementations should treat the supplied values as immutable before/after +/// snapshots of the effective thread configuration. +pub trait ConfigContributor: Send + Sync { + /// Called after the host commits a changed thread configuration. + fn on_config_changed( + &self, + _session_store: &ExtensionData, + _thread_store: &ExtensionData, + _previous_config: &C, + _new_config: &C, + ) { + } +} + +/// Contributor for token usage checkpoints reported by the model provider. +/// +/// Implementations should keep this callback cheap. The host calls it after +/// updating cached token usage and before emitting the corresponding client +/// token-count notification. +pub trait TokenUsageContributor: Send + Sync { + /// Called each time the host records token usage from a model response. + fn on_token_usage<'a>( + &'a self, + _session_store: &'a ExtensionData, + _thread_store: &'a ExtensionData, + _turn_store: &'a ExtensionData, + _token_usage: &'a TokenUsageInfo, + ) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _inputs = (_session_store, _thread_store, _turn_store, _token_usage); + }) + } +} + +/// Contributor for skill invocations observed by the host or an owning extension. +/// +/// Implementations should treat the skill resource as an opaque identity and keep this callback +/// cheap because it runs inline with skill loading or command dispatch. +pub trait SkillInvocationContributor: Send + Sync { + /// Whether this contributor needs a snapshot of host-owned skills. + /// + /// The default preserves legacy discovery for contributors that do not explicitly opt out. + fn requires_host_skill_discovery(&self) -> bool { + true + } + + /// Called after one explicit skill load or deduplicated implicit skill invocation is observed. + fn on_skill_invocation<'a>( + &'a self, + _input: SkillInvocationInput<'a>, + ) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let _self = self; + let _input = _input; + }) + } +} + +/// Extension contribution that exposes native tools owned by a feature. +pub trait ToolContributor: Send + Sync { + /// Returns native tools bound to the supplied extension state. + fn tools( + &self, + session_store: &ExtensionData, + thread_store: &ExtensionData, + ) -> Vec ToolExecutor>>>; + + /// Returns native tools bound to one sampling step. + fn tools_for_step( + &self, + session_store: &ExtensionData, + thread_store: &ExtensionData, + _step_store: &ExtensionData, + ) -> Vec ToolExecutor>>> { + self.tools(session_store, thread_store) + } +} + +/// Contributor for host-owned tool lifecycle gates. +/// +/// Implementations can observe tool execution and process MCP responses without +/// rewriting the invocation. Use `ToolContributor` for owning a tool implementation +/// and hooks for policy that changes tool payloads. +pub trait ToolLifecycleContributor: Send + Sync { + /// Called after pre-tool hooks finalize an invocation and before execution. + /// + /// Calls blocked by hooks, or whose hook-provided input cannot be applied, + /// do not reach this callback. + fn on_tool_start<'a>(&'a self, _input: ToolStartInput<'a>) -> ToolLifecycleFuture<'a> { + Box::pin(std::future::ready(())) + } + + /// Runs before the MCP result is sent to the client and model. + fn on_mcp_tool_result<'a>(&'a self, _input: McpToolResultInput<'a>) -> ToolLifecycleFuture<'a> { + Box::pin(std::future::ready(())) + } + + /// Called after the tool call returns, is blocked, fails, or is cancelled. + /// + /// A matching start callback does not exist when execution is blocked, + /// hook-provided input cannot be applied, or cancellation wins first. + fn on_tool_finish<'a>(&'a self, _input: ToolFinishInput<'a>) -> ToolLifecycleFuture<'a> { + Box::pin(std::future::ready(())) + } +} + +/// Owns the complete approval decision, including whether to consult a reviewer. +/// Returning `None` leaves the request to the next contributor. +pub trait ApprovalReviewContributor: Send + Sync { + /// Claims one request, including a handoff to the user. + fn decide<'a>( + &'a self, + _input: &'a ApprovalDecisionInput<'_>, + ) -> ExtensionFuture<'a, Option> { + Box::pin(std::future::ready(None)) + } +} + +/// Ordered post-processing contribution for one parsed turn item. +/// +/// Implementations may mutate the item before it is emitted and may use the +/// explicitly exposed thread- and turn-lifetime stores when they need durable +/// extension-private state. +pub trait TurnItemContributor: Send + Sync { + fn contribute<'a>( + &'a self, + thread_store: &'a ExtensionData, + turn_store: &'a ExtensionData, + item: &'a mut TurnItem, + ) -> ExtensionFuture<'a, Result<(), String>>; +} diff --git a/codex-rs/ext/extension-api/src/contributors/approval_review.rs b/codex-rs/ext/extension-api/src/contributors/approval_review.rs new file mode 100644 index 0000000000000000000000000000000000000000..ebff9c7a56ad8ee550798d6840acc7d536250069 --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors/approval_review.rs @@ -0,0 +1,48 @@ +//! Request-scoped approval decisions. A review only satisfies the review gate; the host enforces permissions. + +use std::sync::Arc; + +use codex_protocol::ThreadId; + +use crate::ExtensionData; + +/// Thread-local state installed only after Guardian V2's async classifier initializes. +pub struct GuardianV2Enabled; + +/// Guardian's choice for one approval. Synchronous results pass through unchanged. +#[derive(Clone, Debug, PartialEq)] +pub enum ApprovalDecision { + /// Existing async evidence allows this action without synchronous review. + Allow, + Reviewed(codex_protocol::protocol::ReviewDecision), + AskUser, +} + +/// Runs an extension-owned synchronous review already bound to an action by the host. +/// Implementations must not reuse an async score. `None` requests the host user +/// flow when automatic review exhausts its input budget and host policy permits it. +pub trait SynchronousApprovalReviewer: Send + Sync { + fn review( + &self, + reason: codex_protocol::approvals::GuardianReviewReason, + ) -> crate::ExtensionFuture<'_, Option>; +} + +/// Inputs to Guardian's policy choice. Conversation and scores stay thread-owned. +pub struct ApprovalDecisionInput<'a> { + pub approval_id: &'a str, + /// Host tool invocation being approved, absent for approvals without a tool call. + pub tool_call_id: Option<&'a str>, + pub action: &'a serde_json::Value, + pub thread_id: ThreadId, + pub thread_store: &'a ExtensionData, + pub category: codex_protocol::openai_models::GuardianScope, + pub approval_policy: codex_protocol::protocol::AskForApproval, + pub approvals_reviewer: codex_protocol::config_types::ApprovalsReviewer, + pub require_guardian: bool, + /// Existing retry and sensitive-action rules require a synchronous review. + pub require_fresh_review: bool, + pub full_access: bool, + pub metrics: Option>, + pub synchronous_reviewer: &'a dyn SynchronousApprovalReviewer, +} diff --git a/codex-rs/ext/extension-api/src/contributors/context.rs b/codex-rs/ext/extension-api/src/contributors/context.rs new file mode 100644 index 0000000000000000000000000000000000000000..78dd8ae7b5bfe9dee1f292a806d0d485fd6ce15f --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors/context.rs @@ -0,0 +1,20 @@ +use codex_protocol::ThreadId; + +use crate::ExtensionData; + +/// Host context available while extensions contribute turn-scoped context fragments. +#[derive(Clone, Copy)] +pub struct TurnContextContributionInput<'a> { + /// Stable host-owned thread identifier. + pub thread_id: ThreadId, + /// Stable host-owned turn identifier. + pub turn_id: &'a str, + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, + /// Store scoped to this turn. + pub turn_store: &'a ExtensionData, + /// Usable context window of the captured model for this context build, when known. + pub model_context_window: Option, +} diff --git a/codex-rs/ext/extension-api/src/contributors/mcp.rs b/codex-rs/ext/extension-api/src/contributors/mcp.rs new file mode 100644 index 0000000000000000000000000000000000000000..aa3afdc41af07bea54abf9a9980fa7e3c61154f5 --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors/mcp.rs @@ -0,0 +1,165 @@ +use codex_config::McpServerConfig; +use codex_exec_server_protocol::ExecutorCapabilityDiscoverySnapshot; +use codex_protocol::capabilities::SelectedCapabilityRoot; +use codex_protocol::protocol::SessionSource; + +use crate::ExtensionData; +use crate::ExtensionDataInit; + +/// Input supplied while resolving MCP server contributions. +/// +/// Thread-scoped implementations can read stable host inputs through [`Self::thread_init`] and +/// keep their cache in [`Self::thread_store`]. Implementations should not retain borrowed context +/// after contribution completes. +pub struct McpServerContributionContext<'a, C> { + /// Host configuration visible during MCP resolution. + config: &'a C, + /// Extension-owned data for the active thread, when resolution is thread-scoped. + thread_store: Option<&'a ExtensionData>, + /// Stable host inputs for the active thread, when resolution is thread-scoped. + thread_init: Option<&'a ExtensionDataInit>, + /// Source of the active thread, when supplied by the host runtime. + session_source: Option<&'a SessionSource>, + /// Effective request originator for the active thread, when resolution is thread-scoped. + originator: Option<&'a str>, + /// Selected roots resolved against ready environments for this exact step. + ready_selected_capability_roots: Option<&'a [SelectedCapabilityRoot]>, + /// Executor-materialized capability files shared by all consumers in this exact step. + executor_capability_discovery: Option<&'a ExecutorCapabilityDiscoverySnapshot>, +} + +impl Clone for McpServerContributionContext<'_, C> { + fn clone(&self) -> Self { + *self + } +} + +impl Copy for McpServerContributionContext<'_, C> {} + +impl<'a, C> McpServerContributionContext<'a, C> { + /// Creates context for resolution that is not associated with a running thread. + pub fn global(config: &'a C) -> Self { + Self { + config, + thread_store: None, + thread_init: None, + session_source: None, + originator: None, + ready_selected_capability_roots: None, + executor_capability_discovery: None, + } + } + + /// Creates context for one model step using only currently available environments. + pub fn for_step( + config: &'a C, + thread_init: &'a ExtensionDataInit, + thread_store: &'a ExtensionData, + originator: &'a str, + ready_selected_capability_roots: &'a [SelectedCapabilityRoot], + executor_capability_discovery: Option<&'a ExecutorCapabilityDiscoverySnapshot>, + ) -> Self { + Self { + config, + thread_store: Some(thread_store), + thread_init: Some(thread_init), + session_source: None, + originator: Some(originator), + ready_selected_capability_roots: Some(ready_selected_capability_roots), + executor_capability_discovery, + } + } + + /// Attaches the stable source of the active thread to this contribution. + pub fn with_session_source(mut self, session_source: &'a SessionSource) -> Self { + self.session_source = Some(session_source); + self + } + + /// Returns the host configuration visible during resolution. + pub fn config(&self) -> &'a C { + self.config + } + + /// Returns extension-owned state when resolving for a running thread. + pub fn thread_store(&self) -> Option<&'a ExtensionData> { + self.thread_store + } + + /// Returns stable host inputs when resolving for a running thread. + pub fn thread_init(&self) -> Option<&'a ExtensionDataInit> { + self.thread_init + } + + /// Returns the active thread's source when supplied by the host runtime. + pub fn session_source(&self) -> Option<&'a SessionSource> { + self.session_source + } + + /// Returns the effective request originator when resolving for a running thread. + pub fn originator(&self) -> Option<&'a str> { + self.originator + } + + /// Returns selected roots resolved against the ready environments for this model step. + pub fn ready_selected_capability_roots(&self) -> Option<&'a [SelectedCapabilityRoot]> { + self.ready_selected_capability_roots + } + + /// Returns the executor-materialized capability files for this model step, when enabled. + pub fn executor_capability_discovery(&self) -> Option<&'a ExecutorCapabilityDiscoverySnapshot> { + self.executor_capability_discovery + } +} + +/// Validated plugin identities projected for the current set of selected roots. +#[derive(Clone, Debug, Default)] +pub struct SelectedPluginSnapshot { + pub plugins: Vec, + /// Selected plugin roots suppressed by the effective plugin feature policy. + pub disabled_plugin_roots: Vec, +} + +/// The configured identity of a plugin resolved from one selected root. +#[derive(Clone, Debug)] +pub struct SelectedPluginIdentity { + pub selected_root_id: String, + pub plugin_id: String, +} + +/// One extension-owned overlay for the runtime MCP server configuration. +#[derive(Clone, Debug)] +pub enum McpServerContribution { + /// Adds or replaces a named MCP server. + Set { + name: String, + config: Box, + }, + /// Adds an ordinary extension-owned server with its own HTTP protocol mode. + /// The mode applies only if this registration wins server resolution; it + /// does not grant controller-owned Apps cache or environment authority. + SetWithProtocolMode { + name: String, + config: Box, + protocol_mode: crate::McpProtocolMode, + }, + /// Registers the controller-owned Apps server under its reserved name. + HostedApps { config: Box }, + /// Registers a server declared by a plugin selected for this thread. + SelectedPlugin { + name: String, + plugin_id: String, + plugin_display_name: String, + selection_order: usize, + config: Box, + }, + /// Records a plugin selected for this thread and any connector IDs it declares. + SelectedPluginPackage { + selected_root_id: String, + plugin_id: String, + plugin_display_name: String, + connector_ids: Vec, + }, + /// Removes a named MCP server. + Remove { name: String }, +} diff --git a/codex-rs/ext/extension-api/src/contributors/prompt.rs b/codex-rs/ext/extension-api/src/contributors/prompt.rs new file mode 100644 index 0000000000000000000000000000000000000000..73a12e3a6a0a0e92de248fca8a214c2ecacef39e --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors/prompt.rs @@ -0,0 +1,65 @@ +// All this file should be replaced by the existing fragment implementation ofc + +use codex_context_fragments::AnnotatedContent; +use codex_context_fragments::RenderedFragment; +use codex_protocol::models::ContentItemKind; + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +pub enum PromptSlot { + DeveloperPolicy, + DeveloperCapabilities, + /// Text inside the context-window message, supplied by `contribute_thread_context`. + ContextWindow, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct PromptFragment { + slot: PromptSlot, + text: String, + content_kind: ContentItemKind, +} + +impl PromptFragment { + /// Creates a prompt fragment for the given slot. + pub fn new(slot: PromptSlot, text: impl Into, content_kind: ContentItemKind) -> Self { + Self { + slot, + text: text.into(), + content_kind, + } + } + + /// Creates a developer-policy prompt fragment. + pub fn developer_policy(text: impl Into, content_kind: ContentItemKind) -> Self { + Self::new(PromptSlot::DeveloperPolicy, text, content_kind) + } + + /// Creates a developer-capabilities prompt fragment. + pub fn developer_capability(text: impl Into, content_kind: ContentItemKind) -> Self { + Self::new(PromptSlot::DeveloperCapabilities, text, content_kind) + } + + /// Returns the target prompt slot. + pub fn slot(&self) -> PromptSlot { + self.slot + } + + /// Returns the model-visible text. + pub fn text(&self) -> &str { + &self.text + } + + /// Returns the producer-owned classification of the model-visible text. + pub fn content_kind(&self) -> &ContentItemKind { + &self.content_kind + } +} + +impl From for RenderedFragment { + fn from(fragment: PromptFragment) -> Self { + Self::new( + "developer", + AnnotatedContent::input_text(fragment.text, fragment.content_kind), + ) + } +} diff --git a/codex-rs/ext/extension-api/src/contributors/skill_invocation.rs b/codex-rs/ext/extension-api/src/contributors/skill_invocation.rs new file mode 100644 index 0000000000000000000000000000000000000000..5479c54177909ca6607dd9f534f2b7b524286abf --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors/skill_invocation.rs @@ -0,0 +1,26 @@ +use crate::ExtensionData; + +/// Input supplied when the host or an extension observes one skill invocation. +pub struct SkillInvocationInput<'a> { + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, + /// Store scoped to this turn runtime. + pub turn_store: &'a ExtensionData, + /// Current turn submission id. + pub turn_id: &'a str, + /// Main prompt path or opaque resource id for the invoked skill. + pub skill_resource: &'a str, + /// How the skill invocation was initiated. + pub kind: SkillInvocationKind, +} + +/// How an observed skill invocation was initiated. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum SkillInvocationKind { + /// The user explicitly mentioned the skill. + Explicit, + /// The model read the skill instructions or ran one of its scripts. + Implicit, +} diff --git a/codex-rs/ext/extension-api/src/contributors/thread_lifecycle.rs b/codex-rs/ext/extension-api/src/contributors/thread_lifecycle.rs new file mode 100644 index 0000000000000000000000000000000000000000..b614d536f6318d8c50520d0995c96906b826deaf --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors/thread_lifecycle.rs @@ -0,0 +1,84 @@ +use std::sync::Arc; + +use crate::ExtensionData; +use crate::ExtensionMetrics; +use codex_mcp::McpResourceClient; +use codex_protocol::protocol::SessionSource; +use codex_protocol::protocol::TurnEnvironmentSelection; + +/// Trusted, host-resolved billing attribution for a thread. +/// +/// Extensions may forward this value to first-party APIs. It is seeded by Core +/// after resolving persisted and host-provided originator state, rather than +/// from model- or tool-controlled input. +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct ThreadOriginator(pub String); + +/// Input supplied when the host starts a runtime for a thread. +pub struct ThreadStartInput<'a, C> { + /// Host configuration visible at thread start. + pub config: &'a C, + /// Source that created the session for this thread. + pub session_source: &'a SessionSource, + /// Whether persistent thread-scoped state is available for this thread. + pub persistent_thread_state_available: bool, + /// Execution environments selected for this thread. + pub environments: &'a [TurnEnvironmentSelection], + /// MCP resource access supplied by the host for this session. + pub mcp_resource_client: Option>, + /// Session-attributed metrics supplied by the host. + pub extension_metrics: Option>, + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, +} + +/// Input supplied after the host has registered a fully initialized thread. +pub struct ThreadReadyInput<'a, C> { + /// Host configuration visible after thread registration. + pub config: &'a C, + /// Source that created the session for this thread. + pub session_source: &'a SessionSource, + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, +} + +/// Input supplied when the host resumes an existing thread. +pub struct ThreadResumeInput<'a> { + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, +} + +/// Why a thread has no immediately pending work. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum ThreadIdleCause { + /// The previous turn completed and automatic follow-up work can run. + Completed, + /// The user interrupted the previous turn. + Interrupted, + /// The previous turn ended with a terminal error. + Failed, +} + +/// Input supplied when the host has no immediately pending thread work. +pub struct ThreadIdleInput<'a> { + /// Why the thread became idle. + pub cause: ThreadIdleCause, + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, +} + +/// Input supplied during runtime teardown, before persistent history closes. +pub struct ThreadStopInput<'a> { + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, +} diff --git a/codex-rs/ext/extension-api/src/contributors/tool_lifecycle.rs b/codex-rs/ext/extension-api/src/contributors/tool_lifecycle.rs new file mode 100644 index 0000000000000000000000000000000000000000..eb8c8154d43544f6e5c82c1d531a5a7b29fca1cb --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors/tool_lifecycle.rs @@ -0,0 +1,187 @@ +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; + +use codex_config::McpServerConfig; +use codex_mcp::McpServerSource; +use codex_mcp::PreparedMcpCall; +use codex_mcp::ResolvedMcpServer; +use codex_protocol::mcp::CallToolResult; +use codex_tools::ToolCallSource; +use codex_tools::ToolName; +use codex_tools::ToolPayload; +use codex_utils_path_uri::PathUri; + +use crate::ConversationHistorySnapshot; +use crate::ExtensionData; + +/// Future returned by one tool-lifecycle callback. +pub type ToolLifecycleFuture<'a> = Pin + Send + 'a>>; + +/// Extension-facing outcome for a finished tool call. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum ToolCallOutcome { + /// The tool returned a normal output. + Completed { + /// The tool output's own success marker for telemetry/logging. + success: bool, + }, + /// The tool was blocked by host policy before the handler ran. + Blocked, + /// The tool did not produce a normal output. + Failed { + /// Whether the host reached the tool handler before the failure. + handler_executed: bool, + }, + /// The host cancelled the tool before normal completion. Cancellation can + /// win before the dispatch path accepts the call, so contributors should not + /// assume a matching start callback exists. + Aborted, +} + +/// Provenance captured from the immutable MCP call selected for one tool invocation. +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum McpToolSource { + /// A connector routed through the host-owned Codex Apps MCP server. + Connector, + /// An MCP server whose frozen registration matches the active Codex configuration. + Config, + /// An MCP server registered by a locally loaded plugin. + Plugin { + /// Identifier of the plugin that owns this MCP server. + id: String, + /// Host-local plugin root captured with the exact server registration. + root: PathUri, + }, + /// An executor-selected plugin whose root has not been attested by the host. + SelectedPlugin, + /// A compatibility or extension registration without user-owned provenance. + Other, +} + +/// Read-only metadata and provenance captured from the MCP call that will execute. +#[derive(Clone, Debug)] +pub struct McpToolContext { + tool: crate::McpToolInfo, + source: McpToolSource, +} + +impl McpToolContext { + /// Snapshots a prepared call without exposing its executable client to extensions. + /// + /// Configured servers retain their provenance only when their captured connection + /// still matches the host configuration for the current tool invocation. + pub fn from_prepared_call( + call: &PreparedMcpCall, + configured_server: Option<&McpServerConfig>, + ) -> Self { + let tool = call.tool_info().clone(); + let registration = call.config().mcp_server_catalog.server(call.server_name()); + let source = if tool.connector_id.is_some() && call.is_host_owned_apps() { + McpToolSource::Connector + } else if call.is_selected_plugin_server() { + McpToolSource::SelectedPlugin + } else if let Some(McpServerSource::Plugin(plugin)) = + registration.map(ResolvedMcpServer::source) + && Some(plugin.plugin_id()) == call.plugin_id() + && let Some(root) = plugin.host_root() + { + McpToolSource::Plugin { + id: plugin.plugin_id().to_owned(), + root: root.clone(), + } + } else if registration.is_some_and(|server| { + matches!(server.source(), McpServerSource::Config) + && configured_server.is_some_and(|configured| server.config() == configured) + }) { + McpToolSource::Config + } else { + McpToolSource::Other + }; + + Self { tool, source } + } + + /// Returns frozen metadata for the exact model-visible MCP tool being executed. + pub fn tool_info(&self) -> &crate::McpToolInfo { + &self.tool + } + + /// Returns the registration source captured with the executable call. + pub fn source(&self) -> &McpToolSource { + &self.source + } +} + +/// Input supplied when the host starts executing one tool call. +pub struct ToolStartInput<'a> { + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, + /// Store scoped to this turn runtime. + pub turn_store: &'a ExtensionData, + /// Current turn submission id. + pub turn_id: &'a str, + /// Trusted causal root of the owning turn, absent when unknown or ambiguous. + pub root_turn_id: Option<&'a str>, + /// Model-visible tool call id. + pub call_id: &'a str, + /// Responses item that issued this call or started its code-mode cell. + /// Hosts preserve the original wrapper identity across yields and waits. + pub originating_item_id: Option<&'a codex_protocol::ResponseItemId>, + /// Tool name as routed by the host. + pub tool_name: &'a ToolName, + /// Read-only metadata and provenance from the exact MCP call that will execute. + pub mcp_tool: Option<&'a McpToolContext>, + /// Finalized tool arguments, including any pre-tool-use hook rewrites. + /// + /// Payloads can contain sensitive plaintext and must not be logged. + pub payload: &'a ToolPayload, + /// Shared read-only snapshot taken after pre-tool hooks have completed. + pub conversation_history: Arc, + /// Source that issued the tool call. + pub source: ToolCallSource, +} + +/// Input supplied after an MCP server responds, before the host reports completion. +pub struct McpToolResultInput<'a> { + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, + /// Store scoped to this turn runtime. + pub turn_store: &'a ExtensionData, + /// Current turn submission id. + pub turn_id: &'a str, + /// Host tool call id, also used in the MCP completion notification. + pub call_id: &'a str, + /// Read-only metadata and provenance from the exact MCP call that executed. + pub mcp_tool: &'a McpToolContext, + /// Tool arguments after host-side rewriting, including file uploads. + pub arguments: &'a serde_json::Value, + /// Server response, including `_meta`. Changes feed the normal client and model output paths. + /// + /// Arguments and results can contain sensitive plaintext and must not be logged. + pub result: &'a mut CallToolResult, +} + +/// Input supplied when the host finishes executing one tool call. +pub struct ToolFinishInput<'a> { + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, + /// Store scoped to this turn runtime. + pub turn_store: &'a ExtensionData, + /// Current turn submission id. + pub turn_id: &'a str, + /// Model-visible tool call id. + pub call_id: &'a str, + /// Tool name as routed by the host. + pub tool_name: &'a ToolName, + /// Source that issued the tool call. + pub source: ToolCallSource, + /// Host-observed result of the tool call. + pub outcome: ToolCallOutcome, +} diff --git a/codex-rs/ext/extension-api/src/contributors/turn_input.rs b/codex-rs/ext/extension-api/src/contributors/turn_input.rs new file mode 100644 index 0000000000000000000000000000000000000000..61437279cdf103f5f172d73bd5f095ee8173f16c --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors/turn_input.rs @@ -0,0 +1,27 @@ +use codex_protocol::user_input::UserInput; +use codex_utils_path_uri::PathUri; +use std::marker::PhantomData; + +/// Host-owned turn environment summary visible to turn-input contributors. +#[derive(Debug, Clone)] +pub struct TurnInputEnvironment<'a> { + /// Stable host environment id used to route executor-scoped capabilities. + pub environment_id: String, + /// Effective working directory for this turn in the environment. + pub cwd: PathUri, + /// Whether this is the primary environment for the turn. + pub is_primary: bool, + // TODO(anp): Replace the marker with callback-scoped environment access. + pub _lifetime: PhantomData<&'a ()>, +} + +/// Turn facts supplied before the host records turn-local model input items. +#[derive(Debug, Clone)] +pub struct TurnInputContext<'a> { + /// Stable host-owned turn identifier. + pub turn_id: String, + /// User input submitted for this turn. + pub user_input: Vec, + /// Resolved turn environments, in host priority order. + pub environments: Vec>, +} diff --git a/codex-rs/ext/extension-api/src/contributors/turn_lifecycle.rs b/codex-rs/ext/extension-api/src/contributors/turn_lifecycle.rs new file mode 100644 index 0000000000000000000000000000000000000000..e2c9e4bf7bcfac17d1a1a9684629f71ba8539fcf --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors/turn_lifecycle.rs @@ -0,0 +1,58 @@ +use codex_protocol::config_types::CollaborationMode; +use codex_protocol::protocol::CodexErrorInfo; +use codex_protocol::protocol::TokenUsage; +use codex_protocol::protocol::TurnAbortReason; + +use crate::ExtensionData; + +/// Input supplied when the host starts a turn. +pub struct TurnStartInput<'a> { + /// Stable host-owned turn identifier. + pub turn_id: &'a str, + /// Effective collaboration mode for this turn. + pub collaboration_mode: &'a CollaborationMode, + /// Total token usage snapshot captured when the turn started. + pub token_usage_at_turn_start: &'a TokenUsage, + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, + /// Store scoped to this turn runtime. + pub turn_store: &'a ExtensionData, +} + +/// Input supplied when the host completes a turn. +pub struct TurnStopInput<'a> { + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, + /// Store scoped to this turn runtime. + pub turn_store: &'a ExtensionData, +} + +/// Input supplied when the host aborts a turn. +pub struct TurnAbortInput<'a> { + /// Reason the host aborted the turn. + pub reason: TurnAbortReason, + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, + /// Store scoped to this turn runtime. + pub turn_store: &'a ExtensionData, +} + +/// Input supplied when the host observes an error for a turn. +pub struct TurnErrorInput<'a> { + /// Stable host-owned turn identifier. + pub turn_id: &'a str, + /// Error surfaced by the host for this turn. + pub error: CodexErrorInfo, + /// Store scoped to the host session runtime. + pub session_store: &'a ExtensionData, + /// Store scoped to this thread runtime. + pub thread_store: &'a ExtensionData, + /// Store scoped to this turn runtime. + pub turn_store: &'a ExtensionData, +} diff --git a/codex-rs/ext/extension-api/src/contributors/world_state.rs b/codex-rs/ext/extension-api/src/contributors/world_state.rs new file mode 100644 index 0000000000000000000000000000000000000000..c916fac6d059f49db69f486d120e3934e5ec7824 --- /dev/null +++ b/codex-rs/ext/extension-api/src/contributors/world_state.rs @@ -0,0 +1,155 @@ +use std::sync::Arc; + +use codex_exec_server_protocol::ExecutorCapabilityDiscoverySnapshot; +use codex_protocol::ThreadId; +use codex_protocol::capabilities::SelectedCapabilityRoot; +use codex_protocol::openai_models::ModelInfo; +use codex_protocol::protocol::TurnEnvironmentSelection; +use serde_json::Value; + +use crate::ExtensionData; +use crate::ExtensionMetrics; + +/// Host state available while an extension contributes one sampling step's World State. +pub struct WorldStateContributionInput<'a> { + pub thread_id: ThreadId, + pub turn_id: &'a str, + /// Resolved model metadata captured for this sampling step, retained across discovery. + pub model_info: &'a ModelInfo, + pub environments: &'a [TurnEnvironmentSelection], + /// Selected roots whose stable environments are ready in this sampling step. + pub ready_selected_capability_roots: &'a [SelectedCapabilityRoot], + /// Executor-materialized capability files shared by all consumers in this exact step. + pub executor_capability_discovery: Option<&'a ExecutorCapabilityDiscoverySnapshot>, + /// Metrics bound to the captured model for this sampling step. + pub extension_metrics: Option>, + pub session_store: &'a ExtensionData, + pub thread_store: &'a ExtensionData, + pub turn_store: &'a ExtensionData, +} + +/// What the harness knows about the previous value of one extension-owned section. +pub enum PreviousWorldStateSection<'a> { + Absent, + Unknown, + Known(&'a Value), +} + +/// Plain model-visible data rendered by an extension-owned World State section. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct RenderedWorldStateFragment { + role: &'static str, + markers: (&'static str, &'static str), + body: String, +} + +impl RenderedWorldStateFragment { + pub fn new( + role: &'static str, + markers: (&'static str, &'static str), + body: impl Into, + ) -> Self { + Self { + role, + markers, + body: body.into(), + } + } + + pub fn role(&self) -> &'static str { + self.role + } + + pub fn markers(&self) -> (&'static str, &'static str) { + self.markers + } + + pub fn body(&self) -> &str { + &self.body + } +} + +type RenderDiff = dyn for<'a> Fn(PreviousWorldStateSection<'a>) -> Option + + Send + + Sync; +type LegacyFragmentMatcher = dyn Fn(&str, &str) -> bool + Send + Sync; + +/// One extension-owned World State section captured for a sampling step. +/// +/// The extension owns the stable ID, comparison snapshot, and diff rendering. The harness owns +/// persistence and the concrete model-context fragment envelope. +#[derive(Clone)] +pub struct WorldStateSectionContribution { + id: &'static str, + snapshot: Value, + render_diff: Arc, + matches_legacy_fragment: Arc, + matches_retained_fragment: Option>, +} + +impl WorldStateSectionContribution { + pub fn new( + id: &'static str, + snapshot: Value, + render_diff: impl for<'a> Fn( + PreviousWorldStateSection<'a>, + ) -> Option + + Send + + Sync + + 'static, + ) -> Self { + Self { + id, + snapshot, + render_diff: Arc::new(render_diff), + matches_legacy_fragment: Arc::new(|_, _| false), + matches_retained_fragment: None, + } + } + + pub fn with_legacy_matcher( + mut self, + matcher: impl Fn(&str, &str) -> bool + Send + Sync + 'static, + ) -> Self { + self.matches_legacy_fragment = Arc::new(matcher); + self + } + + /// Requires a matching model-visible fragment whenever a persisted snapshot is reused. + pub fn with_retained_fragment_matcher( + mut self, + matcher: impl Fn(&str, &str) -> bool + Send + Sync + 'static, + ) -> Self { + self.matches_retained_fragment = Some(Arc::new(matcher)); + self + } + + pub fn id(&self) -> &'static str { + self.id + } + + pub fn snapshot(&self) -> &Value { + &self.snapshot + } + + pub fn render_diff( + &self, + previous: PreviousWorldStateSection<'_>, + ) -> Option { + (self.render_diff)(previous) + } + + pub fn matches_legacy_fragment(&self, role: &str, text: &str) -> bool { + (self.matches_legacy_fragment)(role, text) + } + + pub fn has_retained_fragment_matcher(&self) -> bool { + self.matches_retained_fragment.is_some() + } + + pub fn matches_retained_fragment(&self, role: &str, text: &str) -> bool { + self.matches_retained_fragment + .as_ref() + .is_some_and(|matcher| matcher(role, text)) + } +} diff --git a/codex-rs/ext/extension-api/src/lib.rs b/codex-rs/ext/extension-api/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..72c56f26f530f83000573bb184b1595cfdcab146 --- /dev/null +++ b/codex-rs/ext/extension-api/src/lib.rs @@ -0,0 +1,106 @@ +mod allowed_tools; +mod capabilities; +mod contributors; +mod registry; +mod session_isolation; +mod state; +mod turn_admission; +mod user_instructions; + +pub use allowed_tools::AllowedTools; +pub use session_isolation::SessionIsolation; + +pub use capabilities::ConversationHistorySnapshot; +pub use capabilities::ExtensionEventSink; +pub use capabilities::ExtensionMetrics; +pub use capabilities::ExtensionWarning; +pub use capabilities::NoopExtensionEventSink; +pub use capabilities::NoopResponseItemInjector; +pub use capabilities::ResponseItemInjectionFuture; +pub use capabilities::ResponseItemInjector; +pub use codex_context_fragments::ContextualUserFragment; +pub use codex_mcp::McpProtocolMode; +pub use codex_mcp::ToolInfo as McpToolInfo; +pub use codex_protocol::models::ContentItemKind; +pub use codex_protocol::models::ResponseItem; +pub use codex_protocol::security_risk::SecurityRiskScore; +pub use codex_tools::ConversationHistory; +pub use codex_tools::ExtensionTurnItem; +pub use codex_tools::FunctionCallError; +pub use codex_tools::JsonToolOutput; +pub use codex_tools::NoopTurnItemEmitter; +pub use codex_tools::ResponsesApiTool; +pub use codex_tools::ToolCall; +pub use codex_tools::ToolCallSource; +pub use codex_tools::ToolEnvironment; +pub use codex_tools::ToolExecutor; +pub use codex_tools::ToolExecutorFuture; +pub use codex_tools::ToolName; +pub use codex_tools::ToolOutput; +pub use codex_tools::ToolPayload; +pub use codex_tools::ToolSpec; +pub use codex_tools::TurnItemEmissionFuture; +pub use codex_tools::TurnItemEmitter; +pub use codex_tools::parse_tool_input_schema; +pub use codex_tools::parse_tool_input_schema_without_compaction; +pub use contributors::ApprovalDecision; +pub use contributors::ApprovalDecisionInput; +pub use contributors::ApprovalReviewContributor; +pub use contributors::ConfigContributor; +pub use contributors::ContextContributor; +pub use contributors::ExtensionFuture; +pub use contributors::GuardianV2Enabled; +pub use contributors::McpServerContribution; +pub use contributors::McpServerContributionContext; +pub use contributors::McpServerContributor; +pub use contributors::McpToolContext; +pub use contributors::McpToolResultInput; +pub use contributors::McpToolSource; +pub use contributors::PreviousWorldStateSection; +pub use contributors::PromptFragment; +pub use contributors::PromptSlot; +pub use contributors::RenderedWorldStateFragment; +pub use contributors::SelectedPluginIdentity; +pub use contributors::SelectedPluginSnapshot; +pub use contributors::SkillInvocationContributor; +pub use contributors::SkillInvocationInput; +pub use contributors::SkillInvocationKind; +pub use contributors::SynchronousApprovalReviewer; +pub use contributors::ThreadIdleCause; +pub use contributors::ThreadIdleInput; +pub use contributors::ThreadLifecycleContributor; +pub use contributors::ThreadOriginator; +pub use contributors::ThreadReadyInput; +pub use contributors::ThreadResumeInput; +pub use contributors::ThreadStartInput; +pub use contributors::ThreadStopInput; +pub use contributors::TokenUsageContributor; +pub use contributors::ToolCallOutcome; +pub use contributors::ToolContributor; +pub use contributors::ToolFinishInput; +pub use contributors::ToolLifecycleContributor; +pub use contributors::ToolLifecycleFuture; +pub use contributors::ToolStartInput; +pub use contributors::TurnAbortInput; +pub use contributors::TurnContextContributionInput; +pub use contributors::TurnErrorInput; +pub use contributors::TurnInputContext; +pub use contributors::TurnInputContributor; +pub use contributors::TurnInputEnvironment; +pub use contributors::TurnItemContributor; +pub use contributors::TurnLifecycleContributor; +pub use contributors::TurnStartInput; +pub use contributors::TurnStopInput; +pub use contributors::WorldStateContributionInput; +pub use contributors::WorldStateSectionContribution; +pub use registry::ExtensionRegistry; +pub use registry::ExtensionRegistryBuilder; +pub use registry::empty_extension_registry; +pub use state::ExtensionData; +pub use state::ExtensionDataInit; +pub use turn_admission::TurnStartAdmission; +pub use user_instructions::Instructions; +pub use user_instructions::LoadInstructionsFuture; +pub use user_instructions::LoadedUserInstructions; +pub use user_instructions::ThreadInstructionsProvider; +pub use user_instructions::UserInstructionsProvider; diff --git a/codex-rs/ext/extension-api/src/registry.rs b/codex-rs/ext/extension-api/src/registry.rs new file mode 100644 index 0000000000000000000000000000000000000000..357cb681d64afbb2364acbc30713b91655ef56ab --- /dev/null +++ b/codex-rs/ext/extension-api/src/registry.rs @@ -0,0 +1,284 @@ +use std::sync::Arc; + +use crate::ApprovalReviewContributor; +use crate::ConfigContributor; +use crate::ContextContributor; +use crate::ExtensionEventSink; +use crate::McpServerContributor; +use crate::NoopExtensionEventSink; +use crate::SkillInvocationContributor; +use crate::ThreadLifecycleContributor; +use crate::TokenUsageContributor; +use crate::ToolContributor; +use crate::ToolLifecycleContributor; +use crate::TurnInputContributor; +use crate::TurnItemContributor; +use crate::TurnLifecycleContributor; +use crate::TurnStartAdmission; + +/// Mutable registry used while hosts register typed runtime contributions. +pub struct ExtensionRegistryBuilder { + registry: ExtensionRegistry, +} + +impl Default for ExtensionRegistryBuilder { + fn default() -> Self { + Self { + registry: ExtensionRegistry { + event_sink: Arc::new(NoopExtensionEventSink), + turn_start_admission: None, + thread_lifecycle_contributors: Vec::new(), + turn_lifecycle_contributors: Vec::new(), + config_contributors: Vec::new(), + token_usage_contributors: Vec::new(), + skill_invocation_contributors: Vec::new(), + approval_review_contributors: Vec::new(), + context_contributors: Vec::new(), + mcp_server_contributors: Vec::new(), + turn_input_contributors: Vec::new(), + tool_contributors: Vec::new(), + tool_lifecycle_contributors: Vec::new(), + turn_item_contributors: Vec::new(), + }, + } + } +} + +impl ExtensionRegistryBuilder { + /// Creates an empty registry builder. + pub fn new() -> Self { + Self::default() + } + + /// Creates an empty registry builder with a host-provided event sink. + pub fn with_event_sink(event_sink: Arc) -> Self { + let mut builder = Self::default(); + builder.registry.event_sink = event_sink; + builder + } + + /// Returns the host event sink to pass into extension constructors. + pub fn event_sink(&self) -> Arc { + Arc::clone(&self.registry.event_sink) + } + + /// Installs the host gate for turn-input submissions that start a new turn. + pub fn turn_start_admission(&mut self, admission: Arc) { + self.registry.turn_start_admission = Some(admission); + } + + /// Registers one approval-review contributor. + pub fn approval_review_contributor(&mut self, contributor: Arc) { + self.registry.approval_review_contributors.push(contributor); + } + + /// Registers one thread-lifecycle contributor. + pub fn thread_lifecycle_contributor( + &mut self, + contributor: Arc>, + ) { + self.registry + .thread_lifecycle_contributors + .push(contributor); + } + + /// Registers one turn-lifecycle contributor. + pub fn turn_lifecycle_contributor(&mut self, contributor: Arc) { + self.registry.turn_lifecycle_contributors.push(contributor); + } + + /// Registers one config contributor. + pub fn config_contributor(&mut self, contributor: Arc>) { + self.registry.config_contributors.push(contributor); + } + + /// Registers one token-usage contributor. + pub fn token_usage_contributor(&mut self, contributor: Arc) { + self.registry.token_usage_contributors.push(contributor); + } + + /// Registers one skill-invocation contributor. + pub fn skill_invocation_contributor( + &mut self, + contributor: Arc, + ) { + self.registry + .skill_invocation_contributors + .push(contributor); + } + + /// Registers one prompt contributor. + pub fn prompt_contributor(&mut self, contributor: Arc) { + self.registry.context_contributors.push(contributor); + } + + /// Registers one runtime MCP server contributor. + pub fn mcp_server_contributor(&mut self, contributor: Arc>) { + self.registry.mcp_server_contributors.push(contributor); + } + + /// Registers one turn-input contributor. + pub fn turn_input_contributor(&mut self, contributor: Arc) { + self.registry.turn_input_contributors.push(contributor); + } + + /// Registers one native tool contributor. + pub fn tool_contributor(&mut self, contributor: Arc) { + self.registry.tool_contributors.push(contributor); + } + + /// Registers one tool-lifecycle contributor. + pub fn tool_lifecycle_contributor(&mut self, contributor: Arc) { + self.registry.tool_lifecycle_contributors.push(contributor); + } + + /// Registers one ordered turn-item contributor. + pub fn turn_item_contributor(&mut self, contributor: Arc) { + self.registry.turn_item_contributors.push(contributor); + } + + /// Finishes construction and returns the immutable registry. + pub fn build(self) -> ExtensionRegistry { + self.registry + } +} + +/// Immutable typed registry produced after extensions are installed. +pub struct ExtensionRegistry { + event_sink: Arc, + turn_start_admission: Option>, + thread_lifecycle_contributors: Vec>>, + turn_lifecycle_contributors: Vec>, + config_contributors: Vec>>, + token_usage_contributors: Vec>, + skill_invocation_contributors: Vec>, + context_contributors: Vec>, + mcp_server_contributors: Vec>>, + turn_input_contributors: Vec>, + tool_contributors: Vec>, + tool_lifecycle_contributors: Vec>, + turn_item_contributors: Vec>, + approval_review_contributors: Vec>, +} + +impl ExtensionRegistry { + /// Copies the registered contributors into a builder for host-specific additions. + pub fn to_builder(&self) -> ExtensionRegistryBuilder { + ExtensionRegistryBuilder { + registry: Self { + event_sink: self.event_sink.clone(), + turn_start_admission: self.turn_start_admission.clone(), + thread_lifecycle_contributors: self.thread_lifecycle_contributors.clone(), + turn_lifecycle_contributors: self.turn_lifecycle_contributors.clone(), + config_contributors: self.config_contributors.clone(), + token_usage_contributors: self.token_usage_contributors.clone(), + skill_invocation_contributors: self.skill_invocation_contributors.clone(), + context_contributors: self.context_contributors.clone(), + mcp_server_contributors: self.mcp_server_contributors.clone(), + turn_input_contributors: self.turn_input_contributors.clone(), + tool_contributors: self.tool_contributors.clone(), + tool_lifecycle_contributors: self.tool_lifecycle_contributors.clone(), + turn_item_contributors: self.turn_item_contributors.clone(), + approval_review_contributors: self.approval_review_contributors.clone(), + }, + } + } + + /// Acquires the host's turn-start permit, or an empty permit for ungated hosts. + /// A missing permit rejects the start before Core consumes pending input. + pub fn admit_turn_start(&self) -> Option> { + match &self.turn_start_admission { + Some(admission) => admission.admit_turn_start(), + None => Some(Box::new(())), + } + } + + /// Returns the host event sink retained by this registry. + pub fn event_sink(&self) -> Arc { + Arc::clone(&self.event_sink) + } + + /// Returns the registered thread-lifecycle contributors. + pub fn thread_lifecycle_contributors(&self) -> &[Arc>] { + &self.thread_lifecycle_contributors + } + + /// Returns the registered turn-lifecycle contributors. + pub fn turn_lifecycle_contributors(&self) -> &[Arc] { + &self.turn_lifecycle_contributors + } + + /// Returns the registered config contributors. + pub fn config_contributors(&self) -> &[Arc>] { + &self.config_contributors + } + + /// Returns the registered token-usage contributors. + pub fn token_usage_contributors(&self) -> &[Arc] { + &self.token_usage_contributors + } + + /// Returns the registered skill-invocation contributors. + pub fn skill_invocation_contributors(&self) -> &[Arc] { + &self.skill_invocation_contributors + } + + /// Whether any installed skill contributor needs a snapshot of host-owned skills. + /// + /// Registries without skill contributors retain legacy host discovery behavior. + pub fn requires_host_skill_discovery(&self) -> bool { + self.skill_invocation_contributors.is_empty() + || self + .skill_invocation_contributors + .iter() + .any(|contributor| contributor.requires_host_skill_discovery()) + } + + /// Returns the first claimed decision in registration order. + pub async fn decide_approval( + &self, + input: &crate::ApprovalDecisionInput<'_>, + ) -> Option { + for contributor in &self.approval_review_contributors { + if let Some(decision) = contributor.decide(input).await { + return Some(decision); + } + } + None + } + + /// Returns the registered prompt contributors. + pub fn context_contributors(&self) -> &[Arc] { + &self.context_contributors + } + + /// Returns the registered runtime MCP server contributors. + pub fn mcp_server_contributors(&self) -> &[Arc>] { + &self.mcp_server_contributors + } + + /// Returns the registered turn-input contributors. + pub fn turn_input_contributors(&self) -> &[Arc] { + &self.turn_input_contributors + } + + /// Returns the registered native tool contributors. + pub fn tool_contributors(&self) -> &[Arc] { + &self.tool_contributors + } + + /// Returns the registered tool-lifecycle contributors. + pub fn tool_lifecycle_contributors(&self) -> &[Arc] { + &self.tool_lifecycle_contributors + } + + /// Returns the registered ordered turn-item contributors. + pub fn turn_item_contributors(&self) -> &[Arc] { + &self.turn_item_contributors + } +} + +/// Creates an empty shared registry for hosts that do not register contributions. +pub fn empty_extension_registry() -> Arc> { + Arc::new(ExtensionRegistryBuilder::new().build()) +} diff --git a/codex-rs/ext/extension-api/src/session_isolation.rs b/codex-rs/ext/extension-api/src/session_isolation.rs new file mode 100644 index 0000000000000000000000000000000000000000..8318517a54986cbc520ad85c023eb3670896529b --- /dev/null +++ b/codex-rs/ext/extension-api/src/session_isolation.rs @@ -0,0 +1,15 @@ +//! Host-supplied isolation for internal runtimes, independent of their attribution. +//! Isolation only removes inherited capabilities; it never grants review authority. + +/// Runtime policy supplied through `ExtensionDataInit` before session startup. +/// The host captures this value once so later extension-state changes cannot alter it. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum SessionIsolation { + /// Use the host's ordinary instruction providers, extensions and execution rules. + #[default] + Inherit, + /// Start without inherited instruction providers or extensions, retain only managed + /// execution rules, and omit executor-discovered MCP servers. Explicitly supplied + /// instructions and permissions remain the responsibility of the internal caller. + Isolated, +} diff --git a/codex-rs/ext/extension-api/src/state.rs b/codex-rs/ext/extension-api/src/state.rs new file mode 100644 index 0000000000000000000000000000000000000000..7d6e79dc5872968147c1aba00b4aa08d1efdae4c --- /dev/null +++ b/codex-rs/ext/extension-api/src/state.rs @@ -0,0 +1,147 @@ +use std::any::Any; +use std::any::TypeId; +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::PoisonError; + +type ErasedData = Arc; + +/// Typed values supplied before an [`ExtensionData`] scope is created. +/// +/// Hosts may retain a clone when later operations must use the same initial +/// inputs. Cloning freezes the attachment map and shares each value by `Arc`; +/// values with interior mutability remain shared. This type does not install +/// extensions or provide persistence. +#[derive(Clone, Debug, Default)] +pub struct ExtensionDataInit { + entries: HashMap, +} + +impl ExtensionDataInit { + /// Creates an empty extension data initializer. + pub fn new() -> Self { + Self::default() + } + + /// Stores `value` as the initial attachment of type `T`. + pub fn insert(&mut self, value: T) -> Option> + where + T: Any + Send + Sync, + { + self.entries + .insert(TypeId::of::(), Arc::new(value)) + .map(downcast_data) + } + + /// Returns a host-supplied initial attachment without creating a mutable scope. + pub fn get(&self) -> Option> + where + T: Any + Send + Sync, + { + let value = self.entries.get(&TypeId::of::())?.clone(); + Some(downcast_data(value)) + } +} + +/// Typed extension-owned data attached to one host object. +#[derive(Debug)] +pub struct ExtensionData { + level_id: String, + entries: Mutex>, +} + +impl ExtensionData { + /// Creates an empty attachment map for one host-owned scope. + pub fn new(level_id: impl Into) -> Self { + Self::new_with_init(level_id, ExtensionDataInit::default()) + } + + /// Creates an attachment map seeded with host-supplied initial data. + pub fn new_with_init(level_id: impl Into, init: ExtensionDataInit) -> Self { + Self { + level_id: level_id.into(), + entries: Mutex::new(init.entries), + } + } + + /// Returns the host identity for the scope this data is attached to. + pub fn level_id(&self) -> &str { + &self.level_id + } + + /// Returns the attached value of type `T`, if one exists. + pub fn get(&self) -> Option> + where + T: Any + Send + Sync, + { + let value = self.entries().get(&TypeId::of::())?.clone(); + Some(downcast_data(value)) + } + + /// Returns the attached value of type `T`, inserting one from `init` when absent. + /// + /// The initializer runs while this map is locked, so it should stay cheap; + /// heavyweight lazy work belongs inside the attached value itself. + pub fn get_or_init(&self, init: impl FnOnce() -> T) -> Arc + where + T: Any + Send + Sync, + { + let mut entries = self.entries(); + let value = entries + .entry(TypeId::of::()) + .or_insert_with(|| Arc::new(init())); + downcast_data(Arc::clone(value)) + } + + /// Stores `value` as the attachment of type `T`, returning any previous value. + pub fn insert(&self, value: T) -> Option> + where + T: Any + Send + Sync, + { + self.entries() + .insert(TypeId::of::(), Arc::new(value)) + .map(downcast_data) + } + + /// Stores `value` only when `should_insert` accepts the current attachment. + /// + /// The predicate and insertion happen while this map is locked, so concurrent + /// callers cannot replace a value after checking a stale attachment. + pub fn insert_if(&self, value: T, should_insert: impl FnOnce(Option<&T>) -> bool) -> bool + where + T: Any + Send + Sync, + { + let mut entries = self.entries(); + let existing = entries + .get(&TypeId::of::()) + .map(|value| downcast_data::(Arc::clone(value))); + if !should_insert(existing.as_deref()) { + return false; + } + entries.insert(TypeId::of::(), Arc::new(value)); + true + } + + /// Removes and returns the attached value of type `T`, if one exists. + pub fn remove(&self) -> Option> + where + T: Any + Send + Sync, + { + self.entries().remove(&TypeId::of::()).map(downcast_data) + } + + fn entries(&self) -> std::sync::MutexGuard<'_, HashMap> { + self.entries.lock().unwrap_or_else(PoisonError::into_inner) + } +} + +fn downcast_data(value: ErasedData) -> Arc +where + T: Any + Send + Sync, +{ + let Ok(value) = value.downcast::() else { + unreachable!("typed extension data stored an incompatible value"); + }; + value +} diff --git a/codex-rs/ext/extension-api/src/turn_admission.rs b/codex-rs/ext/extension-api/src/turn_admission.rs new file mode 100644 index 0000000000000000000000000000000000000000..9d5de89fc7fc77045dd39725a1bbc38a5847e7a9 --- /dev/null +++ b/codex-rs/ext/extension-api/src/turn_admission.rs @@ -0,0 +1,12 @@ +//! Lets Core turn-input submissions participate in a host's shutdown drain. + +/// A host-provided gate checked before Core starts a turn-input submission. +/// +/// Implementations return a permit for work admitted before shutdown and +/// Core retains it through submission. `None` skips the start without consuming +/// pending input. Steering an existing turn does not acquire a new permit. +/// Memory-only mailbox wakeups and parent-delegated subagent input bypass this +/// gate so delegated work can finish before exit. Automatic starts remain gated. +pub trait TurnStartAdmission: std::fmt::Debug + Send + Sync { + fn admit_turn_start(&self) -> Option>; +} diff --git a/codex-rs/ext/extension-api/src/user_instructions.rs b/codex-rs/ext/extension-api/src/user_instructions.rs new file mode 100644 index 0000000000000000000000000000000000000000..b6c9a4b36fa753726984ed9f05532dc0baa80800 --- /dev/null +++ b/codex-rs/ext/extension-api/src/user_instructions.rs @@ -0,0 +1,58 @@ +use std::future::Future; +use std::pin::Pin; + +use codex_utils_absolute_path::AbsolutePathBuf; + +/// Instructions supplied by the host. +/// +/// Filesystem-backed instructions retain their absolute source path for the +/// app-server `instructionSources` API. Other host-provided instructions do +/// not report a filesystem source. +// TODO(anp): Replace the absolute path with a more general instruction-source +// abstraction when non-filesystem providers need first-class attribution. +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct Instructions { + /// Model-visible instruction text. + pub text: String, + /// Absolute filesystem path reported through `instructionSources`, if any. + pub source: Option, +} + +/// Result of loading host-provided user instructions. +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct LoadedUserInstructions { + /// Loaded instructions, or `None` when the provider has no applicable text. + pub instructions: Option, + /// Recoverable loading problems that should be surfaced to the host. + /// Providers own suppression of recurring warnings; Core forwards each returned warning. + pub warnings: Vec, +} + +/// Future returned by an instruction provider. +pub type LoadInstructionsFuture<'a> = + Pin + Send + 'a>>; + +/// Loads host-provided instructions that apply to one root thread. +/// +/// These instructions follow the global [`UserInstructionsProvider`] snapshot +/// and precede repository instructions. A result with no instructions or +/// blank instructions clears only the thread-scoped contribution. Core retains +/// the provider and reads it at startup and when capturing model-request context. +/// Implementations own fetching and caching; repeated reads should be cheap and +/// return a coherent snapshot. On a recoverable fetch failure, return the last +/// usable snapshot with warnings rather than an empty result that clears it. +pub trait ThreadInstructionsProvider: Send + Sync { + /// Loads the current snapshot for the provider's root thread. + fn load_thread_instructions(&self) -> LoadInstructionsFuture<'_>; +} + +/// Loads global user instructions shared across root threads. +/// +/// Core reads this provider at startup and when capturing model-request context. +/// Implementations own fetching and caching, so repeated reads should be cheap. +/// Implementations should return any recoverable loading problems as warnings +/// while still returning usable fallback instructions when available. +pub trait UserInstructionsProvider: Send + Sync { + /// Loads the current global snapshot for a root runtime. + fn load_user_instructions(&self) -> LoadInstructionsFuture<'_>; +} diff --git a/codex-rs/ext/extension-api/tests/capabilities.rs b/codex-rs/ext/extension-api/tests/capabilities.rs new file mode 100644 index 0000000000000000000000000000000000000000..8e015cce231915f70d3186a3482b183a3f2b14ed --- /dev/null +++ b/codex-rs/ext/extension-api/tests/capabilities.rs @@ -0,0 +1,23 @@ +use codex_extension_api::NoopResponseItemInjector; +use codex_extension_api::ResponseItemInjector; +use codex_protocol::models::ContentItem; +use codex_protocol::models::ResponseInputItem; +use pretty_assertions::assert_eq; + +#[tokio::test] +async fn noop_response_item_injector_returns_original_items() { + let items = vec![ResponseInputItem::Message { + role: "user".to_string(), + content: vec![ContentItem::InputText { + text: "keep this input".to_string(), + }], + phase: None, + }]; + + let returned_items = NoopResponseItemInjector + .inject_response_items(items.clone()) + .await + .expect_err("noop injector should reject same-turn injection"); + + assert_eq!(returned_items, items); +} diff --git a/codex-rs/ext/extension-api/tests/registry.rs b/codex-rs/ext/extension-api/tests/registry.rs new file mode 100644 index 0000000000000000000000000000000000000000..c27718b9a560b44d5d3bb7bf19f6b08c9105222d --- /dev/null +++ b/codex-rs/ext/extension-api/tests/registry.rs @@ -0,0 +1,438 @@ +#![allow(clippy::expect_used)] + +use std::sync::Arc; +use std::sync::Mutex; + +use codex_extension_api::ApprovalReviewContributor; +use codex_extension_api::ConfigContributor; +use codex_extension_api::ContentItemKind; +use codex_extension_api::ContextContributor; +use codex_extension_api::ContextualUserFragment; +use codex_extension_api::ConversationHistorySnapshot; +use codex_extension_api::ExtensionData; +use codex_extension_api::ExtensionDataInit; +use codex_extension_api::ExtensionEventSink; +use codex_extension_api::ExtensionFuture; +use codex_extension_api::ExtensionMetrics; +use codex_extension_api::ExtensionRegistryBuilder; +use codex_extension_api::ExtensionWarning; +use codex_extension_api::McpServerContributionContext; +use codex_extension_api::PromptFragment; +use codex_extension_api::ResponseItem; +use codex_extension_api::SkillInvocationContributor; +use codex_extension_api::ThreadLifecycleContributor; +use codex_extension_api::TokenUsageContributor; +use codex_extension_api::ToolCall; +use codex_extension_api::ToolContributor; +use codex_extension_api::ToolExecutor; +use codex_extension_api::ToolLifecycleContributor; +use codex_extension_api::TurnContextContributionInput; +use codex_extension_api::TurnInputContext; +use codex_extension_api::TurnInputContributor; +use codex_extension_api::TurnItemContributor; +use codex_extension_api::TurnLifecycleContributor; +use codex_protocol::items::HookPromptItem; +use codex_protocol::items::TurnItem; +use codex_protocol::protocol::Event; +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::SessionSource; +use codex_protocol::protocol::SubAgentSource; +use codex_protocol::protocol::WarningEvent; +use pretty_assertions::assert_eq; + +struct AllContributors; + +#[test] +fn mcp_contribution_context_identifies_the_running_thread() { + let config = (); + let thread_init = ExtensionDataInit::new(); + let thread_store = ExtensionData::new("child-thread"); + let session_source = SessionSource::SubAgent(SubAgentSource::Review); + + let thread_context = McpServerContributionContext::for_step( + &config, + &thread_init, + &thread_store, + "codex_work_cca", + &[], + /*executor_capability_discovery*/ None, + ) + .with_session_source(&session_source); + + assert_eq!(thread_context.session_source(), Some(&session_source)); + assert_eq!( + McpServerContributionContext::global(&config).session_source(), + None + ); +} + +impl ContextContributor for AllContributors { + fn contribute_thread_context<'a>( + &'a self, + _session_store: &'a ExtensionData, + _thread_store: &'a ExtensionData, + ) -> ExtensionFuture<'a, Vec> { + Box::pin(std::future::ready(Vec::new())) + } +} + +impl ThreadLifecycleContributor<()> for AllContributors {} + +impl TurnLifecycleContributor for AllContributors {} + +impl ConfigContributor<()> for AllContributors {} + +impl TokenUsageContributor for AllContributors {} + +impl SkillInvocationContributor for AllContributors {} + +struct ExecutorOnlySkillContributor; + +impl SkillInvocationContributor for ExecutorOnlySkillContributor { + fn requires_host_skill_discovery(&self) -> bool { + false + } +} + +#[test] +fn host_skill_discovery_preserves_legacy_and_host_contributor_behavior() { + assert!( + ExtensionRegistryBuilder::<()>::new() + .build() + .requires_host_skill_discovery() + ); + + let mut executor_only = ExtensionRegistryBuilder::<()>::new(); + executor_only.skill_invocation_contributor(Arc::new(ExecutorOnlySkillContributor)); + assert!(!executor_only.build().requires_host_skill_discovery()); + + let mut mixed = ExtensionRegistryBuilder::<()>::new(); + mixed.skill_invocation_contributor(Arc::new(ExecutorOnlySkillContributor)); + mixed.skill_invocation_contributor(Arc::new(AllContributors)); + assert!(mixed.build().requires_host_skill_discovery()); +} + +impl TurnInputContributor for AllContributors { + fn contribute<'a>( + &'a self, + input: TurnInputContext<'a>, + _extension_metrics: Option>, + _session_store: &'a ExtensionData, + _thread_store: &'a ExtensionData, + _turn_store: &'a ExtensionData, + ) -> ExtensionFuture<'a, Vec>> { + Box::pin(async move { + let _self = self; + let _input = input; + Vec::new() + }) + } +} + +impl ToolContributor for AllContributors { + fn tools( + &self, + _session_store: &ExtensionData, + _thread_store: &ExtensionData, + ) -> Vec ToolExecutor>>> { + Vec::new() + } +} + +impl ToolLifecycleContributor for AllContributors {} + +impl TurnItemContributor for AllContributors { + fn contribute<'a>( + &'a self, + _thread_store: &'a ExtensionData, + _turn_store: &'a ExtensionData, + _item: &'a mut TurnItem, + ) -> ExtensionFuture<'a, Result<(), String>> { + Box::pin(async move { + let _self = self; + Ok(()) + }) + } +} + +impl ApprovalReviewContributor for AllContributors { + fn decide<'a>( + &'a self, + _input: &'a codex_extension_api::ApprovalDecisionInput<'_>, + ) -> ExtensionFuture<'a, Option> { + Box::pin(async { Some(codex_extension_api::ApprovalDecision::AskUser) }) + } +} + +impl codex_extension_api::SynchronousApprovalReviewer for AllContributors { + fn review( + &self, + _reason: codex_protocol::approvals::GuardianReviewReason, + ) -> ExtensionFuture<'_, Option> { + Box::pin(std::future::ready(Some( + codex_protocol::protocol::ReviewDecision::Approved, + ))) + } +} + +#[tokio::test] +async fn build_round_trips_every_contributor_category() { + let contributor = Arc::new(AllContributors); + let mut builder = ExtensionRegistryBuilder::<()>::new(); + builder.thread_lifecycle_contributor(contributor.clone()); + builder.turn_lifecycle_contributor(contributor.clone()); + builder.config_contributor(contributor.clone()); + builder.token_usage_contributor(contributor.clone()); + builder.skill_invocation_contributor(contributor.clone()); + builder.prompt_contributor(contributor.clone()); + builder.turn_input_contributor(contributor.clone()); + builder.tool_contributor(contributor.clone()); + builder.tool_lifecycle_contributor(contributor.clone()); + builder.turn_item_contributor(contributor.clone()); + builder.approval_review_contributor(contributor); + let registry = builder.build(); + + assert_eq!(registry.thread_lifecycle_contributors().len(), 1); + assert_eq!(registry.turn_lifecycle_contributors().len(), 1); + assert_eq!(registry.config_contributors().len(), 1); + assert_eq!(registry.token_usage_contributors().len(), 1); + assert_eq!(registry.skill_invocation_contributors().len(), 1); + assert_eq!(registry.context_contributors().len(), 1); + assert_eq!(registry.turn_input_contributors().len(), 1); + assert_eq!(registry.tool_contributors().len(), 1); + assert_eq!(registry.tool_lifecycle_contributors().len(), 1); + assert_eq!(registry.turn_item_contributors().len(), 1); + let thread_store = ExtensionData::new("thread"); + let input = codex_extension_api::ApprovalDecisionInput { + approval_id: "approval-1", + tool_call_id: None, + action: &serde_json::Value::Null, + thread_id: codex_protocol::ThreadId::new(), + thread_store: &thread_store, + category: codex_protocol::openai_models::GuardianScope::Shell, + approval_policy: codex_protocol::protocol::AskForApproval::OnRequest, + approvals_reviewer: codex_protocol::config_types::ApprovalsReviewer::AutoReview, + require_guardian: false, + require_fresh_review: false, + full_access: false, + metrics: None, + synchronous_reviewer: &AllContributors, + }; + assert_eq!( + registry.decide_approval(&input).await, + Some(codex_extension_api::ApprovalDecision::AskUser) + ); +} + +impl ConversationHistorySnapshot for AllContributors { + fn history_version(&self) -> u64 { + 0 + } + + fn user_message_revision(&self) -> u64 { + 0 + } + + fn items(&self) -> Box + Send + '_> { + Box::new(std::iter::empty()) + } +} + +struct NamedContextContributor(&'static str); + +impl ContextContributor for NamedContextContributor { + fn contribute_thread_context<'a>( + &'a self, + _session_store: &'a ExtensionData, + _thread_store: &'a ExtensionData, + ) -> ExtensionFuture<'a, Vec> { + Box::pin(std::future::ready(vec![PromptFragment::developer_policy( + self.0, + ContentItemKind("test.thread_context".to_string()), + )])) + } +} + +struct NamedTurnContextContributor(&'static str); + +impl ContextContributor for NamedTurnContextContributor { + fn contribute_turn_context<'a>( + &'a self, + _input: TurnContextContributionInput<'a>, + ) -> ExtensionFuture<'a, Vec> { + Box::pin(std::future::ready(vec![ + PromptFragment::developer_capability( + self.0, + ContentItemKind("test.turn_context".to_string()), + ), + ])) + } +} + +struct RecordingTurnItemContributor { + name: &'static str, + calls: Arc>>, +} + +impl TurnItemContributor for RecordingTurnItemContributor { + fn contribute<'a>( + &'a self, + _thread_store: &'a ExtensionData, + _turn_store: &'a ExtensionData, + _item: &'a mut TurnItem, + ) -> ExtensionFuture<'a, Result<(), String>> { + Box::pin(async move { + self.calls + .lock() + .expect("turn item calls lock should not be poisoned") + .push(self.name); + Ok(()) + }) + } +} + +#[tokio::test] +async fn contributors_preserve_registration_order() { + let turn_item_calls = Arc::new(Mutex::new(Vec::new())); + let mut builder = ExtensionRegistryBuilder::<()>::new(); + builder.prompt_contributor(Arc::new(NamedContextContributor("first"))); + builder.prompt_contributor(Arc::new(NamedContextContributor("second"))); + builder.prompt_contributor(Arc::new(NamedTurnContextContributor("turn-first"))); + builder.prompt_contributor(Arc::new(NamedTurnContextContributor("turn-second"))); + for name in ["first", "second"] { + builder.turn_item_contributor(Arc::new(RecordingTurnItemContributor { + name, + calls: Arc::clone(&turn_item_calls), + })); + } + let registry = builder.build(); + let session_store = ExtensionData::new("session"); + let thread_store = ExtensionData::new("thread"); + let turn_store = ExtensionData::new("turn"); + + let mut fragments = Vec::new(); + for contributor in registry.context_contributors() { + fragments.extend( + contributor + .contribute_thread_context(&session_store, &thread_store) + .await, + ); + } + for contributor in registry.context_contributors() { + fragments.extend( + contributor + .contribute_turn_context(TurnContextContributionInput { + thread_id: codex_protocol::ThreadId::default(), + turn_id: turn_store.level_id(), + session_store: &session_store, + thread_store: &thread_store, + turn_store: &turn_store, + model_context_window: Some(123), + }) + .await, + ); + } + let mut item = TurnItem::HookPrompt(HookPromptItem { + id: "item".to_string(), + fragments: Vec::new(), + }); + for contributor in registry.turn_item_contributors() { + contributor + .contribute(&thread_store, &turn_store, &mut item) + .await + .expect("turn item contribution should succeed"); + } + + assert_eq!( + fragments, + vec![ + PromptFragment::developer_policy( + "first", + ContentItemKind("test.thread_context".to_string()), + ), + PromptFragment::developer_policy( + "second", + ContentItemKind("test.thread_context".to_string()), + ), + PromptFragment::developer_capability( + "turn-first", + ContentItemKind("test.turn_context".to_string()), + ), + PromptFragment::developer_capability( + "turn-second", + ContentItemKind("test.turn_context".to_string()), + ), + ] + ); + assert_eq!( + turn_item_calls + .lock() + .expect("turn item calls lock") + .as_slice(), + ["first", "second"] + ); +} + +#[derive(Default)] +struct RecordingEventSink { + events: Mutex>, +} + +impl ExtensionEventSink for RecordingEventSink { + fn emit(&self, event: Event) { + let EventMsg::Warning(warning) = event.msg else { + panic!("test sink only accepts warning events"); + }; + self.events + .lock() + .expect("recording event sink lock should not be poisoned") + .push((event.id, warning.message)); + } + + fn emit_warning(&self, warning: ExtensionWarning) { + self.events + .lock() + .expect("recording event sink lock should not be poisoned") + .push((warning.thread_id, warning.message)); + } +} + +#[test] +fn custom_event_sink_survives_registry_build() { + let sink = Arc::new(RecordingEventSink::default()); + let builder = ExtensionRegistryBuilder::<()>::with_event_sink(sink.clone()); + + builder + .event_sink() + .emit(warning_event("builder", "before")); + let registry = builder.build(); + registry + .event_sink() + .emit(warning_event("registry", "after")); + registry.event_sink().emit_warning(ExtensionWarning { + thread_id: "thread".to_string(), + turn_id: Some("turn".to_string()), + message: "warning".to_string(), + }); + + assert_eq!( + sink.events + .lock() + .expect("recording event sink lock") + .as_slice(), + [ + ("builder".to_string(), "before".to_string()), + ("registry".to_string(), "after".to_string()), + ("thread".to_string(), "warning".to_string()), + ] + ); +} + +fn warning_event(id: &str, message: &str) -> Event { + Event { + id: id.to_string(), + msg: EventMsg::Warning(WarningEvent { + message: message.to_string(), + }), + } +} diff --git a/codex-rs/ext/extension-api/tests/state.rs b/codex-rs/ext/extension-api/tests/state.rs new file mode 100644 index 0000000000000000000000000000000000000000..6e7720c91374c7b080d7a1d8c9d9c3fafd8a0d7e --- /dev/null +++ b/codex-rs/ext/extension-api/tests/state.rs @@ -0,0 +1,140 @@ +use std::panic::AssertUnwindSafe; +use std::sync::Arc; +use std::sync::Barrier; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; + +use codex_extension_api::ExtensionData; +use pretty_assertions::assert_eq; + +#[test] +fn typed_values_can_be_inserted_replaced_and_removed() { + let data = ExtensionData::new("thread-1"); + + assert_eq!(data.insert(/*value*/ 41_u64), None); + assert_eq!(data.insert("alpha".to_string()), None); + assert_eq!(data.get::().as_deref(), Some(&41)); + assert_eq!( + data.get::().map(|value| value.as_str().to_string()), + Some("alpha".to_string()) + ); + + assert_eq!(data.insert(/*value*/ 42_u64).as_deref(), Some(&41)); + assert_eq!(data.get::().as_deref(), Some(&42)); + assert_eq!( + data.remove::() + .map(|value| value.as_str().to_string()), + Some("alpha".to_string()) + ); + assert_eq!(data.get::(), None); + assert_eq!(data.get::().as_deref(), Some(&42)); +} + +#[test] +fn conditional_insert_keeps_the_newest_concurrent_value() { + const CALLER_COUNT: u64 = 16; + + let data = Arc::new(ExtensionData::new("thread-1")); + let callers_ready = Arc::new(Barrier::new(CALLER_COUNT as usize)); + let handles = (0..CALLER_COUNT) + .map(|value| { + let data = Arc::clone(&data); + let callers_ready = Arc::clone(&callers_ready); + std::thread::spawn(move || { + callers_ready.wait(); + data.insert_if(value, |existing| { + existing.is_none_or(|existing| value > *existing) + }); + }) + }) + .collect::>(); + + for handle in handles { + handle.join().expect("insertion thread should succeed"); + } + + assert_eq!(data.get::().as_deref(), Some(&(CALLER_COUNT - 1))); + assert!(!data.insert_if(/*value*/ 0_u64, |existing| { + existing.is_none_or(|existing| *existing == 0) + })); + assert_eq!(data.get::().as_deref(), Some(&(CALLER_COUNT - 1))); +} + +#[test] +fn get_or_init_initializes_once_and_returns_shared_value() { + const CALLER_COUNT: usize = 8; + + #[derive(Debug, PartialEq, Eq)] + struct SharedValue(usize); + + let data = Arc::new(ExtensionData::new("session")); + let callers_started = Arc::new(AtomicUsize::new(0)); + let initialization_count = Arc::new(AtomicUsize::new(0)); + + let handles: [_; CALLER_COUNT] = std::array::from_fn(|_| { + let data = Arc::clone(&data); + let callers_started = Arc::clone(&callers_started); + let initialization_count = Arc::clone(&initialization_count); + std::thread::spawn(move || { + callers_started.fetch_add(1, Ordering::SeqCst); + data.get_or_init(|| { + initialization_count.fetch_add(1, Ordering::SeqCst); + // Keep the first initializer active until every worker has attempted + // get_or_init, forcing callers to overlap on the same missing entry. + while callers_started.load(Ordering::SeqCst) < CALLER_COUNT { + std::thread::yield_now(); + } + SharedValue(7) + }) + }) + }); + let values = handles + .into_iter() + .map(|handle| handle.join().expect("initializer thread should succeed")) + .collect::>(); + + assert_eq!(initialization_count.load(Ordering::SeqCst), 1); + assert_eq!( + values.iter().map(Arc::as_ref).collect::>(), + vec![&SharedValue(7); CALLER_COUNT] + ); + assert!( + values + .iter() + .skip(1) + .all(|value| Arc::ptr_eq(&values[0], value)) + ); +} + +#[test] +fn stores_are_isolated_and_preserve_level_id() { + let session_data = ExtensionData::new("root-1"); + let thread_data = ExtensionData::new("root-1"); + + session_data.insert(/*value*/ 17_u32); + thread_data.insert("thread value".to_string()); + + assert_eq!(session_data.level_id(), "root-1"); + assert_eq!(thread_data.level_id(), "root-1"); + assert_eq!(session_data.get::().as_deref(), Some(&17)); + assert_eq!(session_data.get::(), None); + assert_eq!(thread_data.get::(), None); + assert_eq!( + thread_data + .get::() + .map(|value| value.as_str().to_string()), + Some("thread value".to_string()) + ); +} + +#[test] +fn store_remains_usable_after_panicking_initializer() { + let data = ExtensionData::new("turn-1"); + + let result = std::panic::catch_unwind(AssertUnwindSafe(|| { + data.get_or_init::(|| panic!("initializer failed")); + })); + + assert!(result.is_err()); + assert_eq!(*data.get_or_init(|| 99_u64), 99); +} diff --git a/codex-rs/ext/git-attribution/BUILD.bazel b/codex-rs/ext/git-attribution/BUILD.bazel new file mode 100644 index 0000000000000000000000000000000000000000..0cb1ab5764c207e1c6a5a3f9700a3f91c8dcbc51 --- /dev/null +++ b/codex-rs/ext/git-attribution/BUILD.bazel @@ -0,0 +1,6 @@ +load("//:defs.bzl", "codex_rust_crate") + +codex_rust_crate( + name = "git-attribution", + crate_name = "codex_git_attribution", +) diff --git a/codex-rs/ext/git-attribution/Cargo.toml b/codex-rs/ext/git-attribution/Cargo.toml new file mode 100644 index 0000000000000000000000000000000000000000..e96d0dbcd577bcc0645f3f6175bbe31429f87159 --- /dev/null +++ b/codex-rs/ext/git-attribution/Cargo.toml @@ -0,0 +1,25 @@ +[package] +edition.workspace = true +license.workspace = true +name = "codex-git-attribution" +version.workspace = true + +[lib] +name = "codex_git_attribution" +path = "src/lib.rs" +doctest = false + +[lints] +workspace = true + +[dependencies] +codex-backend-client = { workspace = true } +codex-extension-api = { workspace = true } +codex-http-client = { workspace = true } +codex-login = { workspace = true } +serde_json = { workspace = true } +tokio = { workspace = true, features = ["time"] } + +[dev-dependencies] +tokio = { workspace = true, features = ["macros", "rt"] } +wiremock = { workspace = true } diff --git a/codex-rs/ext/git-attribution/src/git_attribution_tests.rs b/codex-rs/ext/git-attribution/src/git_attribution_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..0123f5b4e85d50d5a0e462d3ca3b4e4b7f6c9852 --- /dev/null +++ b/codex-rs/ext/git-attribution/src/git_attribution_tests.rs @@ -0,0 +1,144 @@ +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_login::AuthManager; +use codex_login::CodexAuth; +use codex_login::ExternalAuth; +use codex_login::ExternalAuthFuture; +use codex_login::ExternalAuthRefreshContext; +use tokio::sync::Notify; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use super::policy::resolve_attribution_policy; + +fn enterprise_auth_manager() -> Arc { + AuthManager::from_auth_for_testing(enterprise_auth("workspace-123")) +} + +fn http_client_factory() -> HttpClientFactory { + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault) +} + +fn enterprise_auth(account_id: &str) -> CodexAuth { + CodexAuth::from_external_chatgpt_tokens("e30.e30.c2ln", account_id, Some("enterprise")) + .expect("fake ChatGPT auth should parse") +} + +struct StaticExternalAuth(CodexAuth); + +impl ExternalAuth for StaticExternalAuth { + fn resolve(&self) -> ExternalAuthFuture<'_, CodexAuth> { + Box::pin(async { Ok(self.0.clone()) }) + } + + fn refresh(&self, _context: ExternalAuthRefreshContext) -> ExternalAuthFuture<'_, CodexAuth> { + self.resolve() + } +} + +async fn set_auth(auth_manager: &AuthManager, account_id: &str) { + auth_manager + .set_external_auth(Arc::new(StaticExternalAuth(enterprise_auth(account_id)))) + .await + .expect("auth refresh should succeed"); +} + +#[tokio::test] +async fn policy_resolution_recovers_after_unauthorized() { + let server = MockServer::start().await; + let request_count = Arc::new(AtomicUsize::new(0)); + Mock::given(method("GET")) + .and(path("/backend-api/wham/settings/user")) + .respond_with({ + let request_count = request_count.clone(); + move |_request: &wiremock::Request| { + if request_count.fetch_add(1, Ordering::SeqCst) == 0 { + ResponseTemplate::new(401) + } else { + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"commit_attribution_enabled": true})) + } + } + }) + .expect(2) + .mount(&server) + .await; + let auth_manager = enterprise_auth_manager(); + set_auth(auth_manager.as_ref(), "workspace-123").await; + + let policy = resolve_attribution_policy( + &auth_manager, + &format!("{}/backend-api", server.uri()), + &http_client_factory(), + ) + .await + .expect("policy resolution should not time out") + .expect("policy should resolve after auth recovery"); + + assert!(policy.enabled); + assert_eq!(request_count.load(Ordering::SeqCst), 2); + server.verify().await; +} + +#[tokio::test] +async fn policy_resolution_retries_after_auth_refresh() { + let server = MockServer::start().await; + let request_started = Arc::new(Notify::new()); + let request_count = Arc::new(AtomicUsize::new(0)); + Mock::given(method("GET")) + .and(path("/backend-api/wham/settings/user")) + .respond_with({ + let request_started = request_started.clone(); + let request_count = request_count.clone(); + move |_request: &wiremock::Request| match request_count.fetch_add(1, Ordering::SeqCst) { + 0 => { + request_started.notify_one(); + ResponseTemplate::new(200) + .set_delay(Duration::from_millis(100)) + .set_body_json(serde_json::json!({ + "commit_attribution_enabled": true, + })) + } + 1 => ResponseTemplate::new(401), + _ => ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "commit_attribution_enabled": true, + })), + } + }) + .expect(3) + .mount(&server) + .await; + let auth_manager = enterprise_auth_manager(); + let resolve = tokio::spawn({ + let auth_manager = auth_manager.clone(); + let base_url = format!("{}/backend-api", server.uri()); + async move { + resolve_attribution_policy(&auth_manager, &base_url, &http_client_factory()) + .await + .ok() + .flatten() + } + }); + + tokio::time::timeout(Duration::from_secs(5), request_started.notified()) + .await + .expect("first settings request should start"); + set_auth(auth_manager.as_ref(), "workspace-456").await; + + assert!( + resolve + .await + .expect("policy task should complete") + .expect("policy should resolve after refresh") + .enabled + ); + server.verify().await; +} diff --git a/codex-rs/ext/git-attribution/src/lib.rs b/codex-rs/ext/git-attribution/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..a0a3e4dcbf46eb7c72477b35b720d9314af75c77 --- /dev/null +++ b/codex-rs/ext/git-attribution/src/lib.rs @@ -0,0 +1,113 @@ +mod policy; +mod world_state; + +use std::sync::Arc; +use std::time::Instant; + +use codex_extension_api::ContextContributor; +use codex_extension_api::ExtensionFuture; +use codex_extension_api::ExtensionRegistryBuilder; +use codex_extension_api::WorldStateContributionInput; +use codex_extension_api::WorldStateSectionContribution; +use codex_http_client::HttpClientFactory; +use codex_login::AuthManager; + +use crate::policy::GitAttributionPolicy; +use crate::policy::GitAttributionRetry; +use crate::policy::POLICY_RETRY_DELAY; +use crate::policy::auth_generation; +use crate::policy::cached_attribution_policy; +use crate::policy::resolve_attribution_policy; +use crate::policy::retry_deferred; +use crate::world_state::git_attribution_world_state_section; + +/// Contributes model instructions for agent-created git commits and pull requests. +#[derive(Clone)] +struct GitAttributionExtension { + auth_manager: Arc, + base_url: String, + http_client_factory: HttpClientFactory, +} + +impl ContextContributor for GitAttributionExtension { + fn contribute_world_state<'a>( + &'a self, + input: WorldStateContributionInput<'a>, + ) -> ExtensionFuture<'a, Vec> { + Box::pin(async move { + let enabled = loop { + let current_auth_generation = auth_generation(self.auth_manager.as_ref()); + let policy = match cached_attribution_policy( + input.thread_store, + input.turn_store, + current_auth_generation, + ) { + Some(policy) => policy, + None if retry_deferred(input.thread_store, current_auth_generation) => { + GitAttributionPolicy { + auth_generation: current_auth_generation, + enabled: false, + } + } + None => { + match resolve_attribution_policy( + &self.auth_manager, + &self.base_url, + &self.http_client_factory, + ) + .await + { + Ok(Some(policy)) => { + input.thread_store.insert(policy.clone()); + policy + } + Ok(None) => { + let policy = GitAttributionPolicy { + auth_generation: current_auth_generation, + enabled: false, + }; + input.turn_store.insert(policy.clone()); + policy + } + Err(_) => { + let auth_generation = auth_generation(self.auth_manager.as_ref()); + if auth_generation == current_auth_generation { + input.thread_store.insert(GitAttributionRetry { + auth_generation, + retry_at: Instant::now() + POLICY_RETRY_DELAY, + }); + } + GitAttributionPolicy { + auth_generation: current_auth_generation, + enabled: false, + } + } + } + } + }; + if policy.auth_generation == auth_generation(self.auth_manager.as_ref()) { + break policy.enabled; + } + }; + vec![git_attribution_world_state_section(enabled)] + }) + } +} + +/// Installs the git-attribution contributor into the extension registry. +pub fn install( + registry: &mut ExtensionRegistryBuilder, + auth_manager: Arc, + base_url: String, + http_client_factory: HttpClientFactory, +) { + registry.prompt_contributor(Arc::new(GitAttributionExtension { + auth_manager, + base_url, + http_client_factory, + })); +} + +#[cfg(test)] +#[path = "git_attribution_tests.rs"] +mod tests; diff --git a/codex-rs/ext/git-attribution/src/policy.rs b/codex-rs/ext/git-attribution/src/policy.rs new file mode 100644 index 0000000000000000000000000000000000000000..cf607d64d6d6543281e1b88db7a80f63a1bd27ee --- /dev/null +++ b/codex-rs/ext/git-attribution/src/policy.rs @@ -0,0 +1,106 @@ +use std::sync::Arc; +use std::time::Duration; +use std::time::Instant; + +use codex_backend_client::Client as BackendClient; +use codex_extension_api::ExtensionData; +use codex_http_client::HttpClientFactory; +use codex_login::AuthManager; +use tokio::time::timeout; + +#[derive(Clone, Debug)] +pub(super) struct GitAttributionPolicy { + pub(super) auth_generation: u64, + pub(super) enabled: bool, +} + +pub(super) struct GitAttributionRetry { + pub(super) auth_generation: u64, + pub(super) retry_at: Instant, +} + +pub(super) fn retry_deferred(thread_store: &ExtensionData, auth_generation: u64) -> bool { + thread_store + .get::() + .is_some_and(|retry| { + retry.auth_generation == auth_generation && retry.retry_at > Instant::now() + }) +} + +pub(super) fn cached_attribution_policy( + thread_store: &ExtensionData, + turn_store: &ExtensionData, + auth_generation: u64, +) -> Option { + thread_store + .get::() + .filter(|policy| policy.auth_generation == auth_generation) + .or_else(|| { + turn_store + .get::() + .filter(|policy| policy.auth_generation == auth_generation) + }) + .map(|policy| policy.as_ref().clone()) +} + +#[cfg(not(test))] +const POLICY_RESOLUTION_TIMEOUT: Duration = Duration::from_secs(5); +#[cfg(test)] +const POLICY_RESOLUTION_TIMEOUT: Duration = Duration::from_millis(500); +pub(super) const POLICY_RETRY_DELAY: Duration = Duration::from_secs(30); + +pub(super) async fn resolve_attribution_policy( + auth_manager: &Arc, + base_url: &str, + http_client_factory: &HttpClientFactory, +) -> Result, tokio::time::error::Elapsed> { + timeout(POLICY_RESOLUTION_TIMEOUT, async { + let mut recovery_generation = auth_generation(auth_manager); + let mut auth_recovery = auth_manager.unauthorized_recovery(); + loop { + let auth_generation_at_start = auth_generation(auth_manager); + if auth_generation_at_start != recovery_generation { + auth_recovery = auth_manager.unauthorized_recovery(); + recovery_generation = auth_generation_at_start; + } + let auth = auth_manager.auth().await; + if auth_generation(auth_manager) != auth_generation_at_start { + continue; + } + let enabled = match auth { + Some(auth) if auth.uses_codex_backend() => { + let client = + BackendClient::from_auth(base_url, &auth, http_client_factory.clone()); + let settings = client.get_user_settings().await; + if auth_generation(auth_manager) != auth_generation_at_start { + continue; + } + match settings { + Ok(settings) => Some(settings.commit_attribution_enabled), + Err(err) if err.is_unauthorized() && auth_recovery.has_next() => { + if auth_recovery.next().await.is_ok() { + recovery_generation = auth_generation(auth_manager); + continue; + } + None + } + Err(_) => None, + } + } + Some(_) | None => Some(false), + }; + if auth_generation(auth_manager) != auth_generation_at_start { + continue; + } + return enabled.map(|enabled| GitAttributionPolicy { + auth_generation: auth_generation_at_start, + enabled, + }); + } + }) + .await +} + +pub(super) fn auth_generation(auth_manager: &AuthManager) -> u64 { + *auth_manager.auth_change_receiver().borrow() +} diff --git a/codex-rs/ext/git-attribution/src/world_state.rs b/codex-rs/ext/git-attribution/src/world_state.rs new file mode 100644 index 0000000000000000000000000000000000000000..4f96f80109c74c3371273de6533ef617485da040 --- /dev/null +++ b/codex-rs/ext/git-attribution/src/world_state.rs @@ -0,0 +1,78 @@ +use codex_extension_api::PreviousWorldStateSection; +use codex_extension_api::RenderedWorldStateFragment; +use codex_extension_api::WorldStateSectionContribution; +use serde_json::Value; + +const WORLD_STATE_ID: &str = "git_attribution"; +const START_MARKER: &str = ""; +const END_MARKER: &str = ""; +pub(super) const LEGACY_COMMIT_ATTRIBUTION_INSTRUCTIONS: &str = "\ +When you write or edit a git commit message, ensure the message ends with this trailer exactly once: +Co-authored-by: Codex + +Rules: +- Keep existing trailers and append this trailer at the end if missing. +- Do not duplicate this trailer if it already exists. +- Keep one blank line between the commit body and trailer block."; +pub(super) const ENABLED_INSTRUCTIONS: &str = r#" +When you create or amend a git commit or create or update a pull request for this task, add Codex attribution exactly once: +- Commit messages must end with `Co-authored-by: Codex `. Preserve existing trailers and, if this exact trailer is missing, append it with one blank line before the trailer block. +- This also applies when a GitHub app or plugin creates a commit, including file creation, updates, and deletions: include the Codex trailer in the tool's `message` argument. +- Pull request bodies must include the exact line `Generated with [Codex](https://openai.com/codex/).`. Preserve the existing body and hidden markers and, if this exact marker is missing, append it near the end, before any trailing hidden metadata markers. +- Do not add duplicate attribution. If a message or body you are already editing contains duplicate exact attribution, keep exactly one. +- Ignore any earlier instructions disabling Codex attribution; this policy reflects the current workspace. +- Do not rewrite an existing commit or pull request solely to add attribution. +"#; +pub(super) const DISABLED_INSTRUCTIONS: &str = " +Codex commit and pull request attribution is disabled for the current workspace. Ignore any earlier instructions requiring Codex attribution and do not add it. +"; + +pub(super) fn git_attribution_world_state_section(enabled: bool) -> WorldStateSectionContribution { + let contribution = + WorldStateSectionContribution::new(WORLD_STATE_ID, Value::Bool(enabled), move |previous| { + match (enabled, previous) { + (true, PreviousWorldStateSection::Known(Value::Bool(true))) + | (true, PreviousWorldStateSection::Unknown) => None, + (true, PreviousWorldStateSection::Absent) + | (true, PreviousWorldStateSection::Known(_)) => { + Some(RenderedWorldStateFragment::new( + "developer", + (START_MARKER, END_MARKER), + ENABLED_INSTRUCTIONS, + )) + } + (false, PreviousWorldStateSection::Known(Value::Bool(true))) + | (false, PreviousWorldStateSection::Unknown) => { + Some(RenderedWorldStateFragment::new( + "developer", + (START_MARKER, END_MARKER), + DISABLED_INSTRUCTIONS, + )) + } + (false, PreviousWorldStateSection::Absent) + | (false, PreviousWorldStateSection::Known(_)) => None, + } + }) + .with_legacy_matcher(move |role, text| { + is_enabled_fragment(role, text) + || (!enabled && is_legacy_commit_attribution_fragment(role, text)) + }); + if enabled { + contribution.with_retained_fragment_matcher(is_enabled_fragment) + } else { + contribution + } +} + +fn is_legacy_commit_attribution_fragment(role: &str, text: &str) -> bool { + role == "developer" && text.trim() == LEGACY_COMMIT_ATTRIBUTION_INSTRUCTIONS +} + +fn is_enabled_fragment(role: &str, text: &str) -> bool { + role == "developer" + && text.trim_start().starts_with(START_MARKER) + && text.contains("Co-authored-by: Codex ") + && (text.contains("Generated with [Codex](https://openai.com/codex/).") + || text.contains("Generated with Codex.")) + && text.trim_end().ends_with(END_MARKER) +} diff --git a/codex-rs/ext/goal/BUILD.bazel b/codex-rs/ext/goal/BUILD.bazel new file mode 100644 index 0000000000000000000000000000000000000000..c8f3c96c845d804c504ab03d00c9e2737c092f15 --- /dev/null +++ b/codex-rs/ext/goal/BUILD.bazel @@ -0,0 +1,13 @@ +load("//:defs.bzl", "codex_rust_crate") + +codex_rust_crate( + name = "goal", + compile_data = glob([ + "templates/**", + ]), + crate_name = "codex_goal_extension", + integration_compile_data_extra = [ + "src/accounting.rs", + "src/steering.rs", + ] + glob(["templates/**"]), +) diff --git a/codex-rs/ext/goal/Cargo.toml b/codex-rs/ext/goal/Cargo.toml new file mode 100644 index 0000000000000000000000000000000000000000..22ee1794fdf1341819bb8041882b84124ef2362e --- /dev/null +++ b/codex-rs/ext/goal/Cargo.toml @@ -0,0 +1,37 @@ +[package] +edition.workspace = true +license.workspace = true +name = "codex-goal-extension" +version.workspace = true + +[lib] +name = "codex_goal_extension" +path = "src/lib.rs" +test = false +doctest = false + +[lints] +workspace = true + +[dependencies] +codex-analytics = { workspace = true } +codex-core = { workspace = true } +codex-extension-api = { workspace = true } +codex-otel = { workspace = true } +codex-protocol = { workspace = true } +codex-rollout = { workspace = true } +codex-state = { workspace = true } +codex-tools = { workspace = true } +codex-utils-template = { workspace = true } +serde = { workspace = true, features = ["derive"] } +serde_json = { workspace = true } +tokio = { workspace = true, features = ["sync"] } +tracing = { workspace = true } + +[dev-dependencies] +anyhow = { workspace = true } +chrono = { workspace = true } +codex-utils-absolute-path = { workspace = true } +pretty_assertions = { workspace = true } +tempfile = { workspace = true } +tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } diff --git a/codex-rs/ext/goal/src/accounting.rs b/codex-rs/ext/goal/src/accounting.rs new file mode 100644 index 0000000000000000000000000000000000000000..b2a89ce2ed3029c807cfdbed90e11be709a307c9 --- /dev/null +++ b/codex-rs/ext/goal/src/accounting.rs @@ -0,0 +1,647 @@ +use codex_extension_api::ToolCallOutcome; +use codex_extension_api::ToolName; +use codex_protocol::config_types::ModeKind; +use codex_protocol::items::AgentMessageContent; +use codex_protocol::items::TurnItem; +use codex_protocol::models::MessagePhase; +use codex_protocol::protocol::TokenUsage; +use codex_state::ThreadGoalStatus; +use std::collections::HashMap; +use std::sync::Mutex; +use std::sync::PoisonError; +use std::sync::atomic::AtomicI64; +use std::sync::atomic::Ordering; +use std::time::Duration; +use std::time::Instant; +use tokio::sync::Semaphore; +use tokio::sync::SemaphorePermit; + +#[derive(Debug)] +pub(crate) struct GoalAccountingState { + inner: Mutex, + progress_accounting_lock: Semaphore, + descendant_token_usage: AtomicI64, +} + +#[derive(Debug)] +struct GoalAccountingInner { + current_turn_id: Option, + turns: HashMap, + wall_clock: GoalWallClockAccounting, + budget_limit_reported_goal_id: Option, + execution_failure_goal_id: Option, + consecutive_execution_failure_turns: u8, + automatic_goal_turn_id: Option, + consecutive_empty_turns: u8, + last_accounted_descendant_token_usage: i64, +} + +#[derive(Debug)] +struct GoalTurnAccounting { + current_token_usage: TokenUsage, + last_accounted_token_usage: TokenUsage, + active_goal_id: Option, + account_tokens: bool, + failed_execution: bool, + successful_tool: bool, + empty_final: bool, + has_activity: bool, +} + +#[derive(Debug)] +struct GoalWallClockAccounting { + last_accounted_at: Instant, + active_goal_id: Option, +} + +#[derive(Debug, Clone)] +pub(crate) struct GoalProgressSnapshot { + pub(crate) current_token_usage: TokenUsage, + current_descendant_token_usage: i64, + pub(crate) expected_goal_id: String, + pub(crate) time_delta_seconds: i64, + pub(crate) token_delta: i64, +} + +#[derive(Debug, Clone)] +pub(crate) struct IdleGoalProgressSnapshot { + current_descendant_token_usage: i64, + pub(crate) expected_goal_id: String, + pub(crate) time_delta_seconds: i64, + pub(crate) token_delta: i64, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum BudgetLimitedGoalDisposition { + KeepActive, + ClearActive, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct RecordedTokenDelta { + pub(crate) turn_delta: i64, + pub(crate) thread_unflushed_delta: i64, +} + +impl GoalAccountingState { + pub(crate) fn start_turn( + &self, + turn_id: impl Into, + collaboration_mode: ModeKind, + token_usage_at_turn_start: &TokenUsage, + ) { + let turn_id = turn_id.into(); + let mut inner = self.inner(); + inner.current_turn_id = Some(turn_id.clone()); + inner.turns.insert( + turn_id, + GoalTurnAccounting::new( + token_usage_at_turn_start.clone(), + !matches!(collaboration_mode, ModeKind::Plan), + ), + ); + } + + pub(crate) fn current_turn_id(&self) -> Option { + self.inner().current_turn_id.clone() + } + + pub(crate) fn record_tool_outcome( + &self, + turn_id: &str, + tool_name: &ToolName, + outcome: ToolCallOutcome, + ) { + let mut inner = self.inner(); + inner.consecutive_empty_turns = 0; + let Some(turn) = inner.turns.get_mut(turn_id) else { + return; + }; + turn.has_activity = true; + if turn.active_goal_id.is_none() { + return; + } + + match outcome { + ToolCallOutcome::Completed { success: true } => { + turn.successful_tool = true; + inner.execution_failure_goal_id = None; + inner.consecutive_execution_failure_turns = 0; + } + ToolCallOutcome::Failed { + handler_executed: true, + } if tool_name.is_default_namespace() && tool_name.name == "exec" => { + turn.failed_execution = true; + } + ToolCallOutcome::Completed { success: false } + | ToolCallOutcome::Failed { .. } + | ToolCallOutcome::Blocked + | ToolCallOutcome::Aborted => {} + } + } + + pub(crate) fn execution_failure_goal(&self, turn_id: &str) -> Option { + let mut inner = self.inner(); + let turn = inner.turns.get(turn_id)?; + let goal_id = turn.active_goal_id.clone()?; + if turn.successful_tool { + return None; + } + if !turn.failed_execution { + return None; + } + + if inner.execution_failure_goal_id.as_deref() != Some(goal_id.as_str()) { + inner.execution_failure_goal_id = Some(goal_id.clone()); + inner.consecutive_execution_failure_turns = 0; + } + inner.consecutive_execution_failure_turns = + inner.consecutive_execution_failure_turns.saturating_add(1); + (inner.consecutive_execution_failure_turns >= 3).then_some(goal_id) + } + + pub(crate) fn record_item(&self, turn_id: &str, item: &TurnItem) { + let mut inner = self.inner(); + let Some(turn) = inner.turns.get_mut(turn_id) else { + return; + }; + match item { + TurnItem::AgentMessage(message) => { + let has_text = message.content.iter().any(|content| match content { + AgentMessageContent::Text { text } => !text.trim().is_empty(), + }); + turn.has_activity |= has_text || message.questions.is_some(); + turn.empty_final |= + !has_text && !matches!(message.phase, Some(MessagePhase::Commentary)); + } + TurnItem::Reasoning(reasoning) => { + turn.has_activity |= reasoning + .summary_text + .iter() + .any(|text| !text.trim().is_empty()); + } + TurnItem::UserMessage(_) + | TurnItem::FunctionCallOutput(_) + | TurnItem::HookPrompt(_) + | TurnItem::Plan(_) + | TurnItem::CommandExecution(_) + | TurnItem::DynamicToolCall(_) + | TurnItem::CollabAgentToolCall(_) + | TurnItem::SubAgentActivity(_) + | TurnItem::WebSearch(_) + | TurnItem::ImageView(_) + | TurnItem::Extension(_) + | TurnItem::ImageGeneration(_) + | TurnItem::EnteredReviewMode(_) + | TurnItem::ExitedReviewMode(_) + | TurnItem::FileChange(_) + | TurnItem::McpToolCall(_) + | TurnItem::ContextCompaction(_) => turn.has_activity = true, + } + if turn.has_activity { + inner.consecutive_empty_turns = 0; + } + } + + pub(crate) fn mark_goal_continuation(&self, turn_id: String) { + self.inner().automatic_goal_turn_id = Some(turn_id); + } + + pub(crate) fn reset_empty_responses(&self) { + let mut inner = self.inner(); + inner.automatic_goal_turn_id = None; + inner.consecutive_empty_turns = 0; + } + + /// Evaluated under the goal-state permit after automatic admission records its turn ID. + pub(crate) fn empty_response_goal(&self, turn_id: &str) -> Option { + let mut inner = self.inner(); + let automatic = inner.automatic_goal_turn_id.as_deref() == Some(turn_id); + let turn = inner.turns.get_mut(turn_id)?; + let goal_id = turn.active_goal_id.clone()?; + let empty = automatic && turn.empty_final && !turn.has_activity; + turn.empty_final = false; + if !empty { + inner.consecutive_empty_turns = 0; + return None; + } + inner.consecutive_empty_turns = inner.consecutive_empty_turns.saturating_add(1); + (inner.consecutive_empty_turns >= 3).then_some(goal_id) + } + + /// Acquires the per-thread progress-accounting permit. + /// + /// Hold the returned permit from before taking a progress snapshot until after the persistent + /// usage write has succeeded and the snapshot has been marked accounted. This serializes + /// concurrent tool-completion hooks so only one hook can charge a given token or time delta. + pub(crate) async fn progress_accounting_permit( + &self, + ) -> Result, tokio::sync::AcquireError> { + self.progress_accounting_lock.acquire().await + } + + pub(crate) fn current_active_goal_id_for_turn(&self, turn_id: &str) -> Option { + let inner = self.inner(); + if inner.current_turn_id.as_deref() != Some(turn_id) { + return None; + } + let turn = inner.turns.get(turn_id)?; + if !turn.account_tokens { + return None; + } + turn.active_goal_id.clone() + } + + pub(crate) fn record_token_usage( + &self, + turn_id: impl Into, + total_usage: &TokenUsage, + ) -> Option { + let turn_id = turn_id.into(); + let mut inner = self.inner(); + let turn = inner.turns.get_mut(&turn_id)?; + turn.current_token_usage = total_usage.clone(); + if !turn.account_tokens { + return None; + } + + let delta = turn.token_delta_since_last_accounting(); + if delta <= 0 { + return None; + } + Some(RecordedTokenDelta { + turn_delta: delta, + thread_unflushed_delta: inner.thread_unflushed_token_delta(), + }) + } + + pub(crate) fn record_descendant_token_usage(&self, usage: &TokenUsage) { + let delta = goal_token_delta_for_usage(usage); + if delta > 0 { + self.descendant_token_usage + .fetch_add(delta, Ordering::Relaxed); + } + } + + pub(crate) fn mark_turn_goal_active(&self, turn_id: &str, goal_id: impl Into) { + let mut inner = self.inner(); + let goal_id = goal_id.into(); + if inner.budget_limit_reported_goal_id.as_deref() != Some(goal_id.as_str()) { + inner.budget_limit_reported_goal_id = None; + } + if let Some(turn) = inner.turns.get_mut(turn_id) { + turn.active_goal_id = Some(goal_id.clone()); + if inner.current_turn_id.as_deref() == Some(turn_id) { + if inner.wall_clock.active_goal_id.as_deref() != Some(goal_id.as_str()) { + inner.consecutive_empty_turns = 0; + inner.last_accounted_descendant_token_usage = + self.descendant_token_usage.load(Ordering::Relaxed); + } + inner.wall_clock.mark_active_goal(goal_id); + } + } + } + + pub(crate) fn mark_current_turn_goal_active( + &self, + goal_id: impl Into, + ) -> Option { + let mut inner = self.inner(); + let turn_id = inner.current_turn_id.clone()?; + let goal_id = goal_id.into(); + if inner.budget_limit_reported_goal_id.as_deref() != Some(goal_id.as_str()) { + inner.budget_limit_reported_goal_id = None; + } + let goal_changed = inner.wall_clock.active_goal_id.as_deref() != Some(goal_id.as_str()); + let turn = inner.turns.get_mut(turn_id.as_str())?; + if turn.active_goal_id.as_deref() != Some(goal_id.as_str()) { + turn.failed_execution = false; + turn.successful_tool = false; + } + turn.active_goal_id = Some(goal_id.clone()); + if goal_changed { + turn.reset_baseline_to_current(); + inner.automatic_goal_turn_id = None; + inner.consecutive_empty_turns = 0; + inner.last_accounted_descendant_token_usage = + self.descendant_token_usage.load(Ordering::Relaxed); + } + inner.wall_clock.mark_active_goal(goal_id); + Some(turn_id) + } + + pub(crate) fn mark_idle_goal_active(&self, goal_id: impl Into) { + let mut inner = self.inner(); + let goal_id = goal_id.into(); + if inner.budget_limit_reported_goal_id.as_deref() != Some(goal_id.as_str()) { + inner.budget_limit_reported_goal_id = None; + } + if inner.wall_clock.active_goal_id.as_deref() != Some(goal_id.as_str()) { + inner.consecutive_empty_turns = 0; + inner.last_accounted_descendant_token_usage = + self.descendant_token_usage.load(Ordering::Relaxed); + } + inner.wall_clock.mark_active_goal(goal_id); + } + + pub(crate) fn clear_current_turn_goal(&self) -> Option { + let mut inner = self.inner(); + let turn_id = inner.current_turn_id.clone()?; + if let Some(turn) = inner.turns.get_mut(turn_id.as_str()) { + turn.active_goal_id = None; + } + inner.wall_clock.clear_active_goal(); + inner.budget_limit_reported_goal_id = None; + inner.execution_failure_goal_id = None; + inner.consecutive_execution_failure_turns = 0; + inner.automatic_goal_turn_id = None; + inner.consecutive_empty_turns = 0; + Some(turn_id) + } + + pub(crate) fn clear_active_goal(&self) { + let mut inner = self.inner(); + if let Some(turn_id) = inner.current_turn_id.clone() + && let Some(turn) = inner.turns.get_mut(turn_id.as_str()) + { + turn.active_goal_id = None; + } + inner.wall_clock.clear_active_goal(); + inner.budget_limit_reported_goal_id = None; + inner.execution_failure_goal_id = None; + inner.consecutive_execution_failure_turns = 0; + inner.automatic_goal_turn_id = None; + inner.consecutive_empty_turns = 0; + } + + pub(crate) fn progress_snapshot(&self, turn_id: &str) -> Option { + let inner = self.inner(); + let turn = inner.turns.get(turn_id)?; + if !turn.account_tokens { + return None; + } + let expected_goal_id = turn.active_goal_id()?; + let current_descendant_token_usage = self.descendant_token_usage.load(Ordering::Relaxed); + let descendant_token_delta = current_descendant_token_usage + .saturating_sub(inner.last_accounted_descendant_token_usage); + let token_delta = turn + .token_delta_since_last_accounting() + .saturating_add(descendant_token_delta); + let time_delta_seconds = + if inner.wall_clock.active_goal_id.as_deref() == Some(expected_goal_id.as_str()) { + inner.wall_clock.time_delta_since_last_accounting() + } else { + 0 + }; + if time_delta_seconds == 0 && token_delta <= 0 { + return None; + } + Some(GoalProgressSnapshot { + current_token_usage: turn.current_token_usage.clone(), + current_descendant_token_usage, + expected_goal_id, + time_delta_seconds, + token_delta, + }) + } + + pub(crate) fn idle_progress_snapshot(&self) -> Option { + let inner = self.inner(); + let expected_goal_id = inner.wall_clock.active_goal_id.clone()?; + let time_delta_seconds = inner.wall_clock.time_delta_since_last_accounting(); + let current_descendant_token_usage = self.descendant_token_usage.load(Ordering::Relaxed); + let token_delta = current_descendant_token_usage + .saturating_sub(inner.last_accounted_descendant_token_usage); + if time_delta_seconds == 0 && token_delta <= 0 { + return None; + } + Some(IdleGoalProgressSnapshot { + current_descendant_token_usage, + expected_goal_id, + time_delta_seconds, + token_delta, + }) + } + + pub(crate) fn mark_progress_accounted_for_status( + &self, + turn_id: &str, + snapshot: &GoalProgressSnapshot, + status: ThreadGoalStatus, + budget_limited_goal_disposition: BudgetLimitedGoalDisposition, + ) { + let clear_active_goal = should_clear_active_goal(status, budget_limited_goal_disposition); + let mut inner = self.inner(); + if let Some(turn) = inner.turns.get_mut(turn_id) { + turn.last_accounted_token_usage = snapshot.current_token_usage.clone(); + if clear_active_goal { + turn.active_goal_id = None; + } + } + inner.last_accounted_descendant_token_usage = snapshot.current_descendant_token_usage; + inner.wall_clock.mark_accounted(snapshot.time_delta_seconds); + if clear_active_goal { + inner.wall_clock.clear_active_goal(); + } + if status != ThreadGoalStatus::BudgetLimited { + inner.budget_limit_reported_goal_id = None; + } + } + + pub(crate) fn finish_turn(&self, turn_id: &str) { + let mut inner = self.inner(); + inner.turns.remove(turn_id); + if inner.current_turn_id.as_deref() == Some(turn_id) { + inner.current_turn_id = None; + } + } + + pub(crate) fn mark_idle_progress_accounted_for_status( + &self, + snapshot: &IdleGoalProgressSnapshot, + status: ThreadGoalStatus, + budget_limited_goal_disposition: BudgetLimitedGoalDisposition, + ) { + let clear_active_goal = should_clear_active_goal(status, budget_limited_goal_disposition); + let mut inner = self.inner(); + inner.last_accounted_descendant_token_usage = snapshot.current_descendant_token_usage; + inner.wall_clock.mark_accounted(snapshot.time_delta_seconds); + if clear_active_goal { + inner.wall_clock.clear_active_goal(); + } + if status != ThreadGoalStatus::BudgetLimited { + inner.budget_limit_reported_goal_id = None; + } + } + + pub(crate) fn reset_idle_progress_baseline_and_clear_active_goal(&self) { + let mut inner = self.inner(); + inner.wall_clock.reset_baseline(); + inner.wall_clock.clear_active_goal(); + inner.budget_limit_reported_goal_id = None; + } + + pub(crate) fn mark_budget_limit_reported_if_new(&self, goal_id: &str) -> bool { + let mut inner = self.inner(); + if inner.budget_limit_reported_goal_id.as_deref() == Some(goal_id) { + return false; + } + inner.budget_limit_reported_goal_id = Some(goal_id.to_string()); + true + } + + fn inner(&self) -> std::sync::MutexGuard<'_, GoalAccountingInner> { + self.inner.lock().unwrap_or_else(PoisonError::into_inner) + } +} + +impl Default for GoalAccountingState { + fn default() -> Self { + Self { + inner: Mutex::new(GoalAccountingInner::default()), + progress_accounting_lock: Semaphore::new(/*permits*/ 1), + descendant_token_usage: AtomicI64::new(0), + } + } +} + +fn token_delta_since_last_accounting(last: &TokenUsage, current: &TokenUsage) -> i64 { + let delta = TokenUsage { + input_tokens: current.input_tokens.saturating_sub(last.input_tokens), + cached_input_tokens: current + .cached_input_tokens + .saturating_sub(last.cached_input_tokens), + cache_write_input_tokens: current + .cache_write_input_tokens + .saturating_sub(last.cache_write_input_tokens), + output_tokens: current.output_tokens.saturating_sub(last.output_tokens), + reasoning_output_tokens: current + .reasoning_output_tokens + .saturating_sub(last.reasoning_output_tokens), + total_tokens: current.total_tokens.saturating_sub(last.total_tokens), + codex_rollout_budget_units: None, + }; + goal_token_delta_for_usage(&delta) +} + +pub(crate) fn goal_token_delta_for_usage(usage: &TokenUsage) -> i64 { + usage + .input_tokens + .saturating_sub(usage.cached_input_tokens) + .saturating_add(usage.output_tokens.max(0)) +} + +impl Default for GoalAccountingInner { + fn default() -> Self { + Self { + current_turn_id: None, + turns: HashMap::new(), + wall_clock: GoalWallClockAccounting::new(), + budget_limit_reported_goal_id: None, + execution_failure_goal_id: None, + consecutive_execution_failure_turns: 0, + automatic_goal_turn_id: None, + consecutive_empty_turns: 0, + last_accounted_descendant_token_usage: 0, + } + } +} + +impl GoalAccountingInner { + fn thread_unflushed_token_delta(&self) -> i64 { + self.turns + .values() + .filter(|turn| turn.account_tokens) + .fold(0_i64, |total, turn| { + total.saturating_add(turn.token_delta_since_last_accounting().max(0)) + }) + } +} + +impl GoalTurnAccounting { + fn new(current_token_usage: TokenUsage, account_tokens: bool) -> Self { + Self { + last_accounted_token_usage: current_token_usage.clone(), + current_token_usage, + active_goal_id: None, + account_tokens, + failed_execution: false, + successful_tool: false, + empty_final: false, + has_activity: false, + } + } + + fn active_goal_id(&self) -> Option { + self.active_goal_id.clone() + } + + fn reset_baseline_to_current(&mut self) { + self.last_accounted_token_usage = self.current_token_usage.clone(); + } + + fn token_delta_since_last_accounting(&self) -> i64 { + token_delta_since_last_accounting( + &self.last_accounted_token_usage, + &self.current_token_usage, + ) + } +} + +impl GoalWallClockAccounting { + fn new() -> Self { + Self { + last_accounted_at: Instant::now(), + active_goal_id: None, + } + } + + fn time_delta_since_last_accounting(&self) -> i64 { + i64::try_from(self.last_accounted_at.elapsed().as_secs()).unwrap_or(i64::MAX) + } + + fn mark_accounted(&mut self, accounted_seconds: i64) { + if accounted_seconds <= 0 { + return; + } + let advance = Duration::from_secs(u64::try_from(accounted_seconds).unwrap_or(u64::MAX)); + self.last_accounted_at = self + .last_accounted_at + .checked_add(advance) + .unwrap_or_else(Instant::now); + } + + fn reset_baseline(&mut self) { + self.last_accounted_at = Instant::now(); + } + + fn mark_active_goal(&mut self, goal_id: impl Into) { + let goal_id = goal_id.into(); + if self.active_goal_id.as_deref() != Some(goal_id.as_str()) { + self.reset_baseline(); + self.active_goal_id = Some(goal_id); + } + } + + fn clear_active_goal(&mut self) { + self.active_goal_id = None; + self.reset_baseline(); + } +} + +fn should_clear_active_goal( + status: ThreadGoalStatus, + budget_limited_goal_disposition: BudgetLimitedGoalDisposition, +) -> bool { + match status { + ThreadGoalStatus::Active => false, + ThreadGoalStatus::BudgetLimited => matches!( + budget_limited_goal_disposition, + BudgetLimitedGoalDisposition::ClearActive + ), + ThreadGoalStatus::Paused + | ThreadGoalStatus::Blocked + | ThreadGoalStatus::UsageLimited + | ThreadGoalStatus::Complete => true, + } +} diff --git a/codex-rs/ext/goal/src/analytics.rs b/codex-rs/ext/goal/src/analytics.rs new file mode 100644 index 0000000000000000000000000000000000000000..82d34962d1999dbd9fd600a90ad759d6a4f26899 --- /dev/null +++ b/codex-rs/ext/goal/src/analytics.rs @@ -0,0 +1,77 @@ +use codex_analytics::AnalyticsEventsClient; +use codex_analytics::CodexGoalEvent; +use codex_analytics::GoalEventKind; + +#[derive(Clone)] +pub(crate) struct GoalAnalytics { + client: AnalyticsEventsClient, +} + +pub(crate) enum GoalEventAttribution<'a> { + Turn(&'a str), + NoTurn, +} + +impl GoalAnalytics { + pub(crate) fn new(client: AnalyticsEventsClient) -> Self { + Self { client } + } + + pub(crate) fn created( + &self, + goal: &codex_state::ThreadGoal, + attribution: GoalEventAttribution<'_>, + ) { + self.track(goal, attribution, GoalEventKind::Created); + } + + pub(crate) fn usage_accounted( + &self, + goal: &codex_state::ThreadGoal, + attribution: GoalEventAttribution<'_>, + ) { + self.track(goal, attribution, GoalEventKind::UsageAccounted); + } + + pub(crate) fn status_changed( + &self, + goal: &codex_state::ThreadGoal, + previous_status: Option, + attribution: GoalEventAttribution<'_>, + ) { + if previous_status.is_some_and(|status| status != goal.status) { + self.track(goal, attribution, GoalEventKind::StatusChanged); + } + } + + pub(crate) fn cleared(&self, goal: &codex_state::ThreadGoal) { + self.track(goal, GoalEventAttribution::NoTurn, GoalEventKind::Cleared); + } + + fn track( + &self, + goal: &codex_state::ThreadGoal, + attribution: GoalEventAttribution<'_>, + event_kind: GoalEventKind, + ) { + let (cumulative_tokens_accounted, cumulative_time_accounted_seconds) = match event_kind { + GoalEventKind::UsageAccounted => (Some(goal.tokens_used), Some(goal.time_used_seconds)), + GoalEventKind::Created | GoalEventKind::StatusChanged | GoalEventKind::Cleared => { + (None, None) + } + }; + self.client.track_goal_event(CodexGoalEvent { + thread_id: goal.thread_id.to_string(), + turn_id: match attribution { + GoalEventAttribution::Turn(turn_id) => Some(turn_id.to_string()), + GoalEventAttribution::NoTurn => None, + }, + goal_id: goal.goal_id.clone(), + event_kind, + goal_status: goal.status, + has_token_budget: goal.token_budget.is_some(), + cumulative_tokens_accounted, + cumulative_time_accounted_seconds, + }); + } +} diff --git a/codex-rs/ext/goal/src/api.rs b/codex-rs/ext/goal/src/api.rs new file mode 100644 index 0000000000000000000000000000000000000000..b414e9ccfb04eaefa074059e6e1e2c3a71ed46b7 --- /dev/null +++ b/codex-rs/ext/goal/src/api.rs @@ -0,0 +1,368 @@ +use std::collections::HashMap; +use std::fmt; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::PoisonError; +use std::sync::Weak; + +use codex_protocol::ThreadId; +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::ThreadGoal; +use codex_protocol::protocol::ThreadGoalStatus; +use codex_protocol::protocol::ThreadGoalUpdatedEvent; +use codex_protocol::protocol::validate_thread_goal_objective; +use codex_rollout::RolloutItem; + +use crate::runtime::GoalRuntimeHandle; +use crate::runtime::PreviousGoalSnapshot; +use crate::tool::fill_empty_thread_preview_if_possible; +use crate::tool::protocol_goal_from_state; +use crate::tool::state_status_from_protocol; +use crate::tool::validate_goal_budget; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum GoalServiceError { + InvalidRequest(String), + Internal(String), +} + +impl fmt::Display for GoalServiceError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::InvalidRequest(message) | Self::Internal(message) => f.write_str(message), + } + } +} + +impl std::error::Error for GoalServiceError {} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum GoalObjectiveUpdate<'a> { + Keep, + Set(&'a str), +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum GoalTokenBudgetUpdate { + Keep, + Set(Option), +} + +#[derive(Clone, Copy, Debug)] +pub struct GoalSetRequest<'a> { + pub thread_id: ThreadId, + pub objective: GoalObjectiveUpdate<'a>, + pub status: Option, + pub token_budget: GoalTokenBudgetUpdate, + pub max_goal_token_budget: Option, +} + +#[derive(Clone, Debug)] +pub struct GoalSetOutcome { + pub goal: ThreadGoal, + state_goal: codex_state::ThreadGoal, + previous_goal: Option, +} + +impl GoalSetOutcome { + pub fn thread_goal_updated_item(&self) -> RolloutItem { + RolloutItem::EventMsg(EventMsg::ThreadGoalUpdated(ThreadGoalUpdatedEvent { + thread_id: self.goal.thread_id, + turn_id: None, + goal: self.goal.clone(), + })) + } + + pub async fn apply_runtime_effects(&self, goal_service: &GoalService) { + if let Some(runtime) = goal_service.runtime_for_thread(self.goal.thread_id) + && let Err(err) = runtime + .apply_external_goal_set(self.state_goal.clone(), self.previous_goal.clone()) + .await + { + tracing::warn!("failed to apply external goal status runtime effects: {err}"); + } + } +} + +#[derive(Debug, Default)] +pub struct GoalService { + runtimes: Mutex>>, +} + +impl GoalService { + pub fn new() -> Self { + Self::default() + } + + /// Restores persisted goal state into the registered runtime for `thread_id`. + pub async fn restore_thread_runtime_after_resume( + &self, + thread_id: ThreadId, + ) -> Result<(), GoalServiceError> { + let runtime = self.runtime_for_thread(thread_id).ok_or_else(|| { + GoalServiceError::Internal(format!( + "goal runtime is unavailable for thread {thread_id}" + )) + })?; + runtime + .restore_after_resume() + .await + .map_err(GoalServiceError::Internal) + } + + /// Flushes any in-flight goal accounting before a fork copies the source goal snapshot. + pub async fn flush_thread_goal_progress_for_fork( + &self, + thread_id: ThreadId, + ) -> Result<(), GoalServiceError> { + let Some(runtime) = self.runtime_for_thread(thread_id) else { + return Ok(()); + }; + let _goal_state_permit = runtime + .goal_state_permit() + .await + .map_err(GoalServiceError::Internal)?; + runtime + .prepare_external_goal_mutation() + .await + .map_err(GoalServiceError::Internal) + } + + pub async fn get_thread_goal( + &self, + state_db: &codex_state::StateRuntime, + thread_id: ThreadId, + ) -> Result, GoalServiceError> { + state_db + .thread_goals() + .get_thread_goal(thread_id) + .await + .map(|goal| goal.map(protocol_goal_from_state)) + .map_err(|err| GoalServiceError::Internal(format!("failed to read thread goal: {err}"))) + } + + pub async fn set_thread_goal( + &self, + state_db: &codex_state::StateRuntime, + request: GoalSetRequest<'_>, + ) -> Result { + let GoalSetRequest { + thread_id, + objective, + status, + token_budget, + max_goal_token_budget, + } = request; + let status = status.map(state_status_from_protocol); + let objective = match objective { + GoalObjectiveUpdate::Keep => None, + GoalObjectiveUpdate::Set(objective) => Some(objective.trim()), + }; + let token_budget = match token_budget { + GoalTokenBudgetUpdate::Keep => None, + GoalTokenBudgetUpdate::Set(token_budget) => { + Some(token_budget.or(max_goal_token_budget)) + } + }; + + if let Some(objective) = objective { + validate_thread_goal_objective(objective).map_err(GoalServiceError::InvalidRequest)?; + } + if objective.is_some() || token_budget.is_some() { + validate_goal_budget(token_budget.flatten(), max_goal_token_budget) + .map_err(GoalServiceError::InvalidRequest)?; + } + + let runtime = self.runtime_for_thread(thread_id); + // Hold this through the prepare/write window so idle continuation cannot + // launch from goal state that this external mutation is about to change. + let _goal_state_permit = match runtime.as_ref() { + Some(runtime) => Some( + runtime + .goal_state_permit() + .await + .map_err(GoalServiceError::Internal)?, + ), + None => None, + }; + if let Some(runtime) = runtime.as_ref() + && let Err(err) = runtime.prepare_external_goal_mutation().await + { + tracing::warn!("failed to prepare external goal mutation: {err}"); + } + + let (goal, previous_goal) = if let Some(objective) = objective { + let existing_goal = state_db + .thread_goals() + .get_thread_goal(thread_id) + .await + .map_err(|err| { + GoalServiceError::Internal(format!("failed to read thread goal: {err}")) + })?; + if let Some(existing_goal) = existing_goal.as_ref() { + let previous_goal = PreviousGoalSnapshot::from(existing_goal); + state_db + .thread_goals() + .update_thread_goal( + thread_id, + codex_state::GoalUpdate { + objective: Some(objective.to_string()), + status, + token_budget, + expected_goal_id: Some(existing_goal.goal_id.clone()), + }, + ) + .await + .map_err(|err| { + GoalServiceError::Internal(format!("failed to update thread goal: {err}")) + })? + .ok_or_else(|| { + GoalServiceError::InvalidRequest(format!( + "cannot update goal for thread {thread_id}: no goal exists" + )) + }) + .map(|goal| (goal, Some(previous_goal)))? + } else { + state_db + .thread_goals() + .replace_thread_goal( + thread_id, + objective, + status.unwrap_or(codex_state::ThreadGoalStatus::Active), + token_budget.flatten().or(max_goal_token_budget), + ) + .await + .map_err(|err| { + GoalServiceError::Internal(format!("failed to replace thread goal: {err}")) + }) + .map(|goal| (goal, None))? + } + } else { + let existing_goal = state_db + .thread_goals() + .get_thread_goal(thread_id) + .await + .map_err(|err| { + GoalServiceError::Internal(format!("failed to read thread goal: {err}")) + })? + .ok_or_else(|| { + GoalServiceError::InvalidRequest(format!( + "cannot update goal for thread {thread_id}: no goal exists" + )) + })?; + let previous_goal = PreviousGoalSnapshot::from(&existing_goal); + let expected_goal_id = existing_goal.goal_id.clone(); + state_db + .thread_goals() + .update_thread_goal( + thread_id, + codex_state::GoalUpdate { + objective: None, + status, + token_budget, + expected_goal_id: Some(expected_goal_id), + }, + ) + .await + .map_err(|err| { + GoalServiceError::Internal(format!("failed to update thread goal: {err}")) + })? + .ok_or_else(|| { + GoalServiceError::InvalidRequest(format!( + "cannot update goal for thread {thread_id}: no goal exists" + )) + }) + .map(|goal| (goal, Some(previous_goal)))? + }; + + if let Some(runtime) = runtime.as_ref() { + runtime.clear_pending_turn_start_options().await; + } + + if objective.is_some() { + fill_empty_thread_preview_if_possible(state_db, thread_id, &goal).await; + } + Ok(GoalSetOutcome { + goal: protocol_goal_from_state(goal.clone()), + state_goal: goal, + previous_goal, + }) + } + + pub async fn clear_thread_goal( + &self, + state_db: &codex_state::StateRuntime, + thread_id: ThreadId, + ) -> Result { + let runtime = self.runtime_for_thread(thread_id); + // Hold this through the prepare/write window so idle continuation cannot + // launch from goal state that this external mutation is about to change. + let goal_state_permit = match runtime.as_ref() { + Some(runtime) => Some( + runtime + .goal_state_permit() + .await + .map_err(GoalServiceError::Internal)?, + ), + None => None, + }; + if let Some(runtime) = runtime.as_ref() + && let Err(err) = runtime.prepare_external_goal_mutation().await + { + tracing::warn!("failed to prepare external goal mutation: {err}"); + } + + let cleared_goal = state_db + .thread_goals() + .delete_thread_goal(thread_id) + .await + .map_err(|err| { + GoalServiceError::Internal(format!("failed to clear thread goal: {err}")) + })?; + let cleared = cleared_goal.is_some(); + if cleared && let Some(runtime) = runtime.as_ref() { + runtime.clear_pending_turn_start_options().await; + } + drop(goal_state_permit); + drop(runtime); + + if let (Some(runtime), Some(goal)) = (self.runtime_for_thread(thread_id), cleared_goal) + && let Err(err) = runtime.apply_external_goal_clear(goal).await + { + tracing::warn!("failed to apply external goal clear runtime effects: {err}"); + } + + Ok(cleared) + } + + pub(crate) fn register_runtime(&self, runtime: &Arc) { + self.runtimes() + .insert(runtime.thread_id().to_string(), Arc::downgrade(runtime)); + } + + pub(crate) fn unregister_runtime(&self, runtime: &Arc) { + let key = runtime.thread_id().to_string(); + let runtime = Arc::downgrade(runtime); + let mut runtimes = self.runtimes(); + if runtimes + .get(&key) + .is_some_and(|registered| registered.ptr_eq(&runtime)) + { + runtimes.remove(&key); + } + } + + pub(crate) fn runtime_for_thread(&self, thread_id: ThreadId) -> Option> { + let key = thread_id.to_string(); + let mut runtimes = self.runtimes(); + let runtime = runtimes.get(&key).and_then(Weak::upgrade); + if runtime.is_none() { + runtimes.remove(&key); + } + runtime + } + + fn runtimes(&self) -> std::sync::MutexGuard<'_, HashMap>> { + self.runtimes.lock().unwrap_or_else(PoisonError::into_inner) + } +} diff --git a/codex-rs/ext/goal/src/events.rs b/codex-rs/ext/goal/src/events.rs new file mode 100644 index 0000000000000000000000000000000000000000..ab9eda40563523195c79fb5cecf126751e9c53b9 --- /dev/null +++ b/codex-rs/ext/goal/src/events.rs @@ -0,0 +1,34 @@ +use std::sync::Arc; + +use codex_extension_api::ExtensionEventSink; +use codex_protocol::protocol::Event; +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::ThreadGoal; +use codex_protocol::protocol::ThreadGoalUpdatedEvent; + +#[derive(Clone)] +pub(crate) struct GoalEventEmitter { + sink: Arc, +} + +impl GoalEventEmitter { + pub(crate) fn new(sink: Arc) -> Self { + Self { sink } + } + + pub(crate) fn thread_goal_updated( + &self, + event_id: impl Into, + turn_id: Option, + goal: ThreadGoal, + ) { + self.sink.emit(Event { + id: event_id.into(), + msg: EventMsg::ThreadGoalUpdated(ThreadGoalUpdatedEvent { + thread_id: goal.thread_id, + turn_id, + goal, + }), + }); + } +} diff --git a/codex-rs/ext/goal/src/extension.rs b/codex-rs/ext/goal/src/extension.rs new file mode 100644 index 0000000000000000000000000000000000000000..f5dc85138a46f031be8be67f9fdcde5e945c762f --- /dev/null +++ b/codex-rs/ext/goal/src/extension.rs @@ -0,0 +1,626 @@ +use std::sync::Arc; +use std::sync::Weak; + +use codex_analytics::AnalyticsEventsClient; +use codex_core::ThreadManager; +use codex_core::TurnStartOptions; +use codex_extension_api::ConfigContributor; +use codex_extension_api::ExtensionData; +use codex_extension_api::ExtensionEventSink; +use codex_extension_api::ExtensionFuture; +use codex_extension_api::ExtensionRegistryBuilder; +use codex_extension_api::ThreadIdleInput; +use codex_extension_api::ThreadLifecycleContributor; +use codex_extension_api::ThreadResumeInput; +use codex_extension_api::ThreadStartInput; +use codex_extension_api::ThreadStopInput; +use codex_extension_api::TokenUsageContributor; +use codex_extension_api::ToolCall; +use codex_extension_api::ToolCallOutcome; +use codex_extension_api::ToolContributor; +use codex_extension_api::ToolExecutor; +use codex_extension_api::ToolFinishInput; +use codex_extension_api::ToolLifecycleContributor; +use codex_extension_api::ToolLifecycleFuture; +use codex_extension_api::TurnAbortInput; +use codex_extension_api::TurnErrorInput; +use codex_extension_api::TurnLifecycleContributor; +use codex_extension_api::TurnStartInput; +use codex_extension_api::TurnStopInput; +use codex_otel::MetricsClient; +use codex_protocol::ThreadId; +use codex_protocol::items::TurnItem; +use codex_protocol::protocol::CodexErrorInfo; +use codex_protocol::protocol::SessionSource; +use codex_protocol::protocol::SubAgentSource; +use codex_protocol::protocol::ThreadGoalStatus; +use codex_protocol::protocol::TokenUsageInfo; + +use crate::accounting::BudgetLimitedGoalDisposition; +use crate::accounting::GoalAccountingState; +use crate::analytics::GoalAnalytics; +use crate::api::GoalService; +use crate::events::GoalEventEmitter; +use crate::metrics::GoalMetrics; +use crate::runtime::ActiveGoalStopReason; +use crate::runtime::GoalRuntimeConfig; +use crate::runtime::GoalRuntimeHandle; +use crate::spec::CREATE_GOAL_TOOL_NAME; +use crate::spec::UPDATE_GOAL_TOOL_NAME; +use crate::steering::budget_limit_steering_item; +use crate::tool::GoalToolExecutor; + +#[derive(Clone, Debug)] +pub struct GoalExtensionConfig { + pub enabled: bool, + pub max_goal_token_budget: Option, +} + +#[derive(Clone)] +pub struct GoalExtension { + state_dbs: Arc, + analytics: GoalAnalytics, + event_emitter: GoalEventEmitter, + metrics: GoalMetrics, + thread_manager: Weak, + goal_service: Arc, + goal_config: Arc GoalExtensionConfig + Send + Sync>, +} + +impl std::fmt::Debug for GoalExtension { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("GoalExtension").finish_non_exhaustive() + } +} + +impl GoalExtension { + pub(crate) fn new_with_host_capabilities( + state_dbs: Arc, + analytics_events_client: AnalyticsEventsClient, + event_sink: Arc, + metrics_client: Option, + thread_manager: Weak, + goal_service: Arc, + goal_config: impl Fn(&C) -> GoalExtensionConfig + Send + Sync + 'static, + ) -> Self { + Self { + state_dbs, + analytics: GoalAnalytics::new(analytics_events_client), + event_emitter: GoalEventEmitter::new(event_sink), + metrics: GoalMetrics::new(metrics_client), + thread_manager, + goal_service, + goal_config: Arc::new(goal_config), + } + } +} + +impl ThreadLifecycleContributor for GoalExtension +where + C: Send + Sync + 'static, +{ + fn on_thread_start<'a>(&'a self, input: ThreadStartInput<'a, C>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let config = (self.goal_config)(input.config); + let enabled = config.enabled; + let tools_visible_for_thread = !matches!( + input.session_source, + SessionSource::SubAgent(SubAgentSource::Review) + ); + let tools_available_for_thread = + input.persistent_thread_state_available && tools_visible_for_thread; + input.thread_store.insert(config); + let accounting_state = input + .thread_store + .get_or_init::(GoalAccountingState::default); + let Ok(thread_id) = ThreadId::from_string(input.thread_store.level_id()) else { + return; + }; + let root_accounting_state = input + .session_source + .parent_thread_id() + .or_else(|| { + ThreadId::from_string(input.session_store.level_id()) + .ok() + .filter(|_| input.session_source.is_non_root_agent()) + }) + .and_then(|parent_thread_id| { + self.goal_service + .runtime_for_thread(parent_thread_id) + .or_else(|| { + ThreadId::from_string(input.session_store.level_id()) + .ok() + .and_then(|root_thread_id| { + self.goal_service.runtime_for_thread(root_thread_id) + }) + }) + }) + .map(|parent| { + parent + .root_accounting_state() + .unwrap_or_else(|| parent.accounting_state()) + }); + let runtime = input.thread_store.get_or_init::(|| { + GoalRuntimeHandle::new( + thread_id, + Arc::clone(&self.state_dbs), + self.event_emitter.clone(), + self.metrics.clone(), + self.thread_manager.clone(), + accounting_state, + GoalRuntimeConfig { + analytics: self.analytics.clone(), + enabled, + tools_available_for_thread, + tools_visible_for_thread, + root_accounting_state, + }, + ) + }); + runtime.set_enabled(enabled); + self.goal_service.register_runtime(&runtime); + }) + } + + fn on_thread_resume<'a>(&'a self, input: ThreadResumeInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let Some(runtime) = goal_runtime_handle(input.thread_store) else { + return; + }; + + if let Err(err) = runtime.restore_after_resume().await { + tracing::warn!( + "failed to restore goal runtime after thread resume for {}: {err}", + runtime.thread_id() + ); + } + }) + } + + fn on_thread_idle<'a>(&'a self, input: ThreadIdleInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let Some(runtime) = goal_runtime_handle(input.thread_store) else { + return; + }; + + if let Err(err) = runtime.continue_if_idle().await { + tracing::warn!( + "failed to continue active goal for idle thread {}: {err}", + runtime.thread_id() + ); + } + }) + } + + fn on_thread_stop<'a>(&'a self, input: ThreadStopInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + if let Some(runtime) = goal_runtime_handle(input.thread_store) { + self.goal_service.unregister_runtime(&runtime); + } + }) + } +} + +impl ConfigContributor for GoalExtension +where + C: Send + Sync + 'static, +{ + fn on_config_changed( + &self, + _session_store: &ExtensionData, + thread_store: &ExtensionData, + _previous_config: &C, + new_config: &C, + ) { + let config = (self.goal_config)(new_config); + let enabled = config.enabled; + thread_store.insert(config); + if let Some(runtime) = goal_runtime_handle(thread_store) { + runtime.set_enabled(enabled); + } + } +} + +impl TurnLifecycleContributor for GoalExtension +where + C: Send + Sync + 'static, +{ + fn on_turn_start<'a>(&'a self, input: TurnStartInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let Some(runtime) = goal_runtime_handle(input.thread_store) else { + return; + }; + if !runtime.is_enabled() { + return; + } + + if let Err(err) = self + .state_dbs + .thread_goals() + .clear_thread_goal_continuation_deferral(runtime.thread_id()) + .await + { + tracing::warn!("failed to clear deferred goal continuation: {err}"); + } + + let accounting = runtime.accounting_state(); + accounting.start_turn( + input.turn_id, + input.collaboration_mode.mode, + input.token_usage_at_turn_start, + ); + if matches!( + input.collaboration_mode.mode, + codex_protocol::config_types::ModeKind::Plan + ) { + accounting.clear_current_turn_goal(); + return; + } + let Ok(goal) = self + .state_dbs + .thread_goals() + .get_thread_goal(runtime.thread_id()) + .await + else { + return; + }; + if let Some(goal) = goal + && matches!( + goal.status, + codex_state::ThreadGoalStatus::Active + | codex_state::ThreadGoalStatus::BudgetLimited + ) + { + accounting.mark_turn_goal_active(input.turn_id, goal.goal_id); + } + }) + } + + fn on_item_completed<'a>( + &'a self, + thread_store: &'a ExtensionData, + turn_store: &'a ExtensionData, + item: &'a TurnItem, + ) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + if let Some(runtime) = goal_runtime_handle(thread_store) + && runtime.is_enabled() + { + runtime + .accounting_state() + .record_item(turn_store.level_id(), item); + } + }) + } + + fn on_turn_stop<'a>(&'a self, input: TurnStopInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let Some(runtime) = goal_runtime_handle(input.thread_store) else { + return; + }; + if !runtime.is_enabled() { + return; + } + + let turn_id = input.turn_store.level_id(); + if let Some(expected_goal_id) = + runtime.accounting_state().execution_failure_goal(turn_id) + && let Err(err) = runtime + .stop_active_goal_for_turn( + turn_id, + ActiveGoalStopReason::ExecutionUnavailable { expected_goal_id }, + ) + .await + { + input.thread_store.remove::(); + tracing::warn!( + "failed to stop active goal after repeated execution failures for {turn_id}: {err}" + ); + return; + } + if let Err(err) = runtime + .stop_active_goal_for_turn(turn_id, ActiveGoalStopReason::EmptyResponse) + .await + { + input.thread_store.remove::(); + tracing::warn!("failed to stop goal after empty responses for {turn_id}: {err}"); + return; + } + if let Err(err) = runtime + .account_active_goal_progress( + turn_id, + &format!("{turn_id}:turn-stop"), + codex_state::GoalAccountingMode::ActiveOnly, + BudgetLimitedGoalDisposition::ClearActive, + ) + .await + { + input.thread_store.remove::(); + tracing::warn!( + "failed to account active goal progress at turn stop for {turn_id}: {err}" + ); + return; + } + let accounting = runtime.accounting_state(); + if accounting + .current_active_goal_id_for_turn(turn_id) + .is_some() + && let Some(options) = input.thread_store.get::() + { + input.thread_store.insert_if( + TurnStartOptions { + parent_turn_id: Some(turn_id.to_string()), + ..options.as_ref().clone() + }, + |current| current.is_some(), + ); + } + accounting.finish_turn(turn_id); + }) + } + + fn on_turn_abort<'a>(&'a self, input: TurnAbortInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let Some(runtime) = goal_runtime_handle(input.thread_store) else { + return; + }; + runtime.accounting_state().reset_empty_responses(); + if !runtime.is_enabled() { + return; + } + + let turn_id = input.turn_store.level_id(); + input.thread_store.remove::(); + if let Err(err) = runtime + .account_active_goal_progress( + turn_id, + &format!("{turn_id}:turn-abort"), + codex_state::GoalAccountingMode::ActiveOnly, + BudgetLimitedGoalDisposition::ClearActive, + ) + .await + { + tracing::warn!( + "failed to account active goal progress after turn abort for {turn_id}: {err}" + ); + return; + } + runtime.accounting_state().finish_turn(turn_id); + }) + } + + fn on_turn_error<'a>(&'a self, input: TurnErrorInput<'a>) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let Some(runtime) = goal_runtime_handle(input.thread_store) else { + return; + }; + + let reason = match input.error { + CodexErrorInfo::UsageLimitExceeded => ActiveGoalStopReason::UsageLimit, + // The turn has ended because the error was non-retryable or its + // retries were exhausted. Block the goal to prevent automatic + // continuation from looping and consuming tokens, as can happen + // with compaction errors. + _ => ActiveGoalStopReason::TurnError, + }; + if let Err(err) = runtime + .stop_active_goal_for_turn(input.turn_id, reason) + .await + { + tracing::warn!( + error = ?input.error, + "failed to stop active goal after turn error: {err}" + ); + } + }) + } +} + +impl TokenUsageContributor for GoalExtension +where + C: Send + Sync + 'static, +{ + fn on_token_usage<'a>( + &'a self, + _session_store: &'a ExtensionData, + thread_store: &'a ExtensionData, + turn_store: &'a ExtensionData, + token_usage: &'a TokenUsageInfo, + ) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + let Some(runtime) = goal_runtime_handle(thread_store) else { + return; + }; + if !runtime.is_enabled() { + return; + } + + if let Some(root_accounting_state) = runtime.root_accounting_state() { + root_accounting_state.record_descendant_token_usage(&token_usage.last_token_usage); + } + let _ = runtime + .accounting_state() + .record_token_usage(turn_store.level_id(), &token_usage.total_token_usage); + }) + } +} + +impl ToolLifecycleContributor for GoalExtension +where + C: Send + Sync + 'static, +{ + fn on_tool_finish<'a>(&'a self, input: ToolFinishInput<'a>) -> ToolLifecycleFuture<'a> { + Box::pin(async move { + let Some(runtime) = goal_runtime_handle(input.thread_store) else { + return; + }; + if !runtime.is_enabled() { + return; + } + runtime.accounting_state().record_tool_outcome( + input.turn_id, + input.tool_name, + input.outcome, + ); + if input.tool_name.is_default_namespace() + && input.tool_name.name == CREATE_GOAL_TOOL_NAME + && matches!(input.outcome, ToolCallOutcome::Completed { success: true }) + { + input.thread_store.remove::(); + if let Ok(_goal_state_permit) = runtime.goal_state_permit().await + && let Some(thread_manager) = self.thread_manager.upgrade() + && let Ok(thread) = thread_manager.get_thread(runtime.thread_id()).await + && let Some(root_turn_id) = thread.active_turn_root(input.turn_id).await + { + input.thread_store.insert(TurnStartOptions { + root_turn_id: Some(root_turn_id), + parent_turn_id: Some(input.turn_id.to_string()), + ..Default::default() + }); + } + } + let should_count_for_goal_progress = + tool_attempt_counts_for_goal_progress(input.outcome) + && !(input.tool_name.is_default_namespace() + && input.tool_name.name == UPDATE_GOAL_TOOL_NAME); + if !should_count_for_goal_progress { + return; + } + let turn_id = input.turn_id; + let progress = match runtime + .account_active_goal_progress( + turn_id, + input.call_id, + codex_state::GoalAccountingMode::ActiveOnly, + BudgetLimitedGoalDisposition::KeepActive, + ) + .await + { + Ok(Some(progress)) => progress, + Ok(None) => return, + Err(err) => { + tracing::warn!( + "failed to account active goal progress after tool finish for {turn_id}: {err}" + ); + return; + } + }; + let goal = progress.goal; + if goal.status != ThreadGoalStatus::BudgetLimited { + return; + } + if !runtime + .accounting_state() + .mark_budget_limit_reported_if_new(progress.goal_id.as_str()) + { + return; + } + let item = budget_limit_steering_item(&goal); + runtime.inject_active_turn_steering(item).await; + }) + } +} + +impl ToolContributor for GoalExtension +where + C: Send + Sync + 'static, +{ + fn tools( + &self, + _session_store: &ExtensionData, + thread_store: &ExtensionData, + ) -> Vec< + Arc codex_extension_api::ToolExecutor>>, + > { + let Some(runtime) = goal_runtime_handle(thread_store) else { + return Vec::new(); + }; + if !runtime.tools_visible() { + return Vec::new(); + } + let max_goal_token_budget = thread_store + .get::() + .and_then(|config| config.max_goal_token_budget); + + let tools = [ + GoalToolExecutor::get( + runtime.thread_id(), + Arc::clone(&self.state_dbs), + runtime.accounting_state(), + self.analytics.clone(), + self.event_emitter.clone(), + self.metrics.clone(), + ), + GoalToolExecutor::create( + runtime.thread_id(), + Arc::clone(&self.state_dbs), + runtime.accounting_state(), + self.analytics.clone(), + self.event_emitter.clone(), + self.metrics.clone(), + max_goal_token_budget, + ), + GoalToolExecutor::update( + runtime.thread_id(), + Arc::clone(&self.state_dbs), + runtime.accounting_state(), + self.analytics.clone(), + self.event_emitter.clone(), + self.metrics.clone(), + ), + ]; + tools + .into_iter() + .map(|mut tool| { + tool.execution_allowed = runtime.tools_available(); + Arc::new(tool) as Arc ToolExecutor>> + }) + .collect() + } +} + +pub fn install_with_backend( + registry: &mut ExtensionRegistryBuilder, + state_dbs: Arc, + analytics_events_client: AnalyticsEventsClient, + metrics_client: Option, + thread_manager: Weak, + goal_service: Arc, + goal_config: impl Fn(&C) -> GoalExtensionConfig + Send + Sync + 'static, +) where + C: Send + Sync + 'static, +{ + let extension = Arc::new(GoalExtension::new_with_host_capabilities( + state_dbs, + analytics_events_client, + registry.event_sink(), + metrics_client, + thread_manager, + Arc::clone(&goal_service), + goal_config, + )); + registry.thread_lifecycle_contributor(extension.clone()); + registry.config_contributor(extension.clone()); + registry.turn_lifecycle_contributor(extension.clone()); + registry.token_usage_contributor(extension.clone()); + registry.tool_lifecycle_contributor(extension.clone()); + registry.tool_contributor(extension); +} + +fn goal_runtime_handle(thread_store: &ExtensionData) -> Option> { + thread_store.get::() +} + +fn tool_attempt_counts_for_goal_progress(outcome: ToolCallOutcome) -> bool { + match outcome { + ToolCallOutcome::Completed { .. } => true, + ToolCallOutcome::Failed { + handler_executed: true, + } => true, + ToolCallOutcome::Blocked + | ToolCallOutcome::Failed { + handler_executed: false, + } + | ToolCallOutcome::Aborted => false, + } +} diff --git a/codex-rs/ext/goal/src/lib.rs b/codex-rs/ext/goal/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..ccd091affce4646394f7c6e41f84b2153feb6654 --- /dev/null +++ b/codex-rs/ext/goal/src/lib.rs @@ -0,0 +1,28 @@ +//! Extension crate for the `/goal` feature. + +mod accounting; +mod analytics; +mod api; +mod events; +mod extension; +mod metrics; +mod runtime; +mod spec; +mod steering; +mod tool; + +pub use api::GoalObjectiveUpdate; +pub use api::GoalService; +pub use api::GoalServiceError; +pub use api::GoalSetOutcome; +pub use api::GoalSetRequest; +pub use api::GoalTokenBudgetUpdate; +pub use extension::GoalExtension; +pub use extension::GoalExtensionConfig; +pub use extension::install_with_backend; +pub use runtime::GoalRuntimeHandle; +pub use runtime::PreviousGoalSnapshot; +pub use spec::CREATE_GOAL_TOOL_NAME; +pub use spec::GET_GOAL_TOOL_NAME; +pub use spec::UPDATE_GOAL_TOOL_NAME; +pub use tool::CreateGoalRequest; diff --git a/codex-rs/ext/goal/src/metrics.rs b/codex-rs/ext/goal/src/metrics.rs new file mode 100644 index 0000000000000000000000000000000000000000..0bb227ef6185ad88d9328c8f0660531a18ccfe3c --- /dev/null +++ b/codex-rs/ext/goal/src/metrics.rs @@ -0,0 +1,84 @@ +use codex_otel::GOAL_BLOCKED_METRIC; +use codex_otel::GOAL_BUDGET_LIMITED_METRIC; +use codex_otel::GOAL_COMPLETED_METRIC; +use codex_otel::GOAL_CREATED_METRIC; +use codex_otel::GOAL_DURATION_SECONDS_METRIC; +use codex_otel::GOAL_RESUMED_METRIC; +use codex_otel::GOAL_TOKEN_COUNT_METRIC; +use codex_otel::GOAL_USAGE_LIMITED_METRIC; +use codex_otel::MetricsClient; + +#[derive(Clone, Default)] +pub(crate) struct GoalMetrics { + metrics_client: Option, +} + +impl GoalMetrics { + pub(crate) fn new(metrics_client: Option) -> Self { + Self { metrics_client } + } + + pub(crate) fn record_created(&self) { + let Some(metrics_client) = self.metrics_client.as_ref() else { + return; + }; + let _ = metrics_client.counter(GOAL_CREATED_METRIC, /*inc*/ 1, &[]); + } + + pub(crate) fn record_resumed(&self) { + let Some(metrics_client) = self.metrics_client.as_ref() else { + return; + }; + let _ = metrics_client.counter(GOAL_RESUMED_METRIC, /*inc*/ 1, &[]); + } + + pub(crate) fn record_resumed_if_status_changed( + &self, + previous_status: Option, + goal_status: codex_state::ThreadGoalStatus, + ) { + if goal_status == codex_state::ThreadGoalStatus::Active + && matches!( + previous_status, + Some( + codex_state::ThreadGoalStatus::Paused + | codex_state::ThreadGoalStatus::Blocked + | codex_state::ThreadGoalStatus::UsageLimited + ) + ) + { + self.record_resumed(); + } + } + + pub(crate) fn record_terminal_if_status_changed( + &self, + previous_status: Option, + goal: &codex_state::ThreadGoal, + ) { + if previous_status == Some(goal.status) { + return; + } + + let counter = match goal.status { + codex_state::ThreadGoalStatus::Blocked => GOAL_BLOCKED_METRIC, + codex_state::ThreadGoalStatus::UsageLimited => GOAL_USAGE_LIMITED_METRIC, + codex_state::ThreadGoalStatus::BudgetLimited => GOAL_BUDGET_LIMITED_METRIC, + codex_state::ThreadGoalStatus::Complete => GOAL_COMPLETED_METRIC, + codex_state::ThreadGoalStatus::Active | codex_state::ThreadGoalStatus::Paused => { + return; + } + }; + let Some(metrics_client) = self.metrics_client.as_ref() else { + return; + }; + let status_tag = [("status", goal.status.as_str())]; + let _ = metrics_client.counter(counter, /*inc*/ 1, &[]); + let _ = metrics_client.histogram(GOAL_TOKEN_COUNT_METRIC, goal.tokens_used, &status_tag); + let _ = metrics_client.histogram( + GOAL_DURATION_SECONDS_METRIC, + goal.time_used_seconds, + &status_tag, + ); + } +} diff --git a/codex-rs/ext/goal/src/runtime.rs b/codex-rs/ext/goal/src/runtime.rs new file mode 100644 index 0000000000000000000000000000000000000000..7d826ac2381df8b0dafa39ef66959b22527a4604 --- /dev/null +++ b/codex-rs/ext/goal/src/runtime.rs @@ -0,0 +1,683 @@ +use std::sync::Arc; +use std::sync::Weak; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +use codex_core::StartIfIdleSubmission; +use codex_core::ThreadManager; +use codex_core::TurnInput; +use codex_core::TurnInputRequest; +use codex_core::TurnStartOptions; +use codex_protocol::ThreadId; +use codex_protocol::models::ResponseItem; +use codex_protocol::protocol::ThreadGoal; + +use crate::accounting::BudgetLimitedGoalDisposition; +use crate::accounting::GoalAccountingState; +use crate::analytics::GoalAnalytics; +use crate::analytics::GoalEventAttribution; +use crate::events::GoalEventEmitter; +use crate::metrics::GoalMetrics; +use crate::steering::continuation_steering_item; +use crate::steering::objective_updated_steering_item; +use crate::tool::protocol_goal_from_state; +use tokio::sync::Semaphore; +use tokio::sync::SemaphorePermit; + +#[derive(Clone)] +pub struct GoalRuntimeHandle { + inner: Arc, +} + +pub(crate) struct GoalRuntimeConfig { + pub(crate) analytics: GoalAnalytics, + pub(crate) enabled: bool, + pub(crate) tools_available_for_thread: bool, + pub(crate) tools_visible_for_thread: bool, + pub(crate) root_accounting_state: Option>, +} + +pub(crate) enum ActiveGoalStopReason { + TurnError, + UsageLimit, + ExecutionUnavailable { expected_goal_id: String }, + EmptyResponse, +} + +struct GoalRuntimeInner { + thread_id: ThreadId, + state_dbs: Arc, + analytics: GoalAnalytics, + event_emitter: GoalEventEmitter, + metrics: GoalMetrics, + thread_manager: Weak, + accounting_state: Arc, + root_accounting_state: Option>, + enabled: AtomicBool, + tools_available_for_thread: bool, + tools_visible_for_thread: bool, + goal_state_lock: Semaphore, +} + +pub(crate) struct AccountedGoalProgress { + pub(crate) goal: ThreadGoal, + pub(crate) goal_id: String, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct PreviousGoalSnapshot { + pub goal_id: String, + pub status: codex_state::ThreadGoalStatus, + pub objective: String, +} + +impl From<&codex_state::ThreadGoal> for PreviousGoalSnapshot { + fn from(goal: &codex_state::ThreadGoal) -> Self { + Self { + goal_id: goal.goal_id.clone(), + status: goal.status, + objective: goal.objective.clone(), + } + } +} + +impl std::fmt::Debug for GoalRuntimeHandle { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("GoalRuntimeHandle").finish_non_exhaustive() + } +} + +impl GoalRuntimeHandle { + pub(crate) fn new( + thread_id: ThreadId, + state_dbs: Arc, + event_emitter: GoalEventEmitter, + metrics: GoalMetrics, + thread_manager: Weak, + accounting_state: Arc, + config: GoalRuntimeConfig, + ) -> Self { + Self { + inner: Arc::new(GoalRuntimeInner { + thread_id, + state_dbs, + analytics: config.analytics, + event_emitter, + metrics, + thread_manager, + accounting_state, + root_accounting_state: config.root_accounting_state, + enabled: AtomicBool::new(config.enabled), + tools_available_for_thread: config.tools_available_for_thread, + tools_visible_for_thread: config.tools_visible_for_thread, + goal_state_lock: Semaphore::new(/*permits*/ 1), + }), + } + } + + pub(crate) fn set_enabled(&self, enabled: bool) { + self.inner.enabled.store(enabled, Ordering::Relaxed); + } + + pub(crate) fn is_enabled(&self) -> bool { + self.inner.enabled.load(Ordering::Relaxed) + } + + pub(crate) fn tools_visible(&self) -> bool { + self.is_enabled() && self.inner.tools_visible_for_thread + } + + pub(crate) fn tools_available(&self) -> bool { + self.is_enabled() && self.inner.tools_available_for_thread + } + + pub(crate) fn thread_id(&self) -> ThreadId { + self.inner.thread_id + } + + pub(crate) fn accounting_state(&self) -> Arc { + Arc::clone(&self.inner.accounting_state) + } + + pub(crate) fn root_accounting_state(&self) -> Option> { + self.inner.root_accounting_state.clone() + } + + pub(crate) async fn clear_pending_turn_start_options(&self) { + let Some(thread_manager) = self.inner.thread_manager.upgrade() else { + return; + }; + let Ok(thread) = thread_manager.get_thread(self.inner.thread_id).await else { + return; + }; + thread.thread_extension_data().remove::(); + } + + pub(crate) async fn goal_state_permit(&self) -> Result, String> { + self.inner + .goal_state_lock + .acquire() + .await + .map_err(|err| err.to_string()) + } + + pub async fn prepare_external_goal_mutation(&self) -> Result<(), String> { + if !self.is_enabled() { + return Ok(()); + } + // Invalidate the old turn before the persisted objective/status changes. + self.inner.accounting_state.reset_empty_responses(); + + if let Some(turn_id) = self.inner.accounting_state.current_turn_id() { + self.account_active_goal_progress( + turn_id.as_str(), + &format!("{turn_id}:external-goal-mutation"), + codex_state::GoalAccountingMode::ActiveOnly, + BudgetLimitedGoalDisposition::ClearActive, + ) + .await?; + return Ok(()); + } + + self.account_idle_goal_progress( + &format!("{}:external-goal-mutation", self.inner.thread_id), + codex_state::GoalAccountingMode::ActiveOnly, + BudgetLimitedGoalDisposition::ClearActive, + ) + .await?; + Ok(()) + } + + pub async fn apply_external_goal_set( + &self, + goal: codex_state::ThreadGoal, + previous_goal: Option, + ) -> Result<(), String> { + if !self.is_enabled() { + return Ok(()); + } + + self.inner.accounting_state.reset_empty_responses(); + let replaced_existing_goal = previous_goal + .as_ref() + .is_some_and(|previous_goal| previous_goal.goal_id != goal.goal_id); + if previous_goal.is_none() || replaced_existing_goal { + self.inner.metrics.record_created(); + self.inner + .analytics + .created(&goal, GoalEventAttribution::NoTurn); + } + let previous_status = previous_goal + .as_ref() + .and_then(|previous_goal| (!replaced_existing_goal).then_some(previous_goal.status)); + self.inner + .metrics + .record_resumed_if_status_changed(previous_status, goal.status); + self.inner + .metrics + .record_terminal_if_status_changed(previous_status, &goal); + self.inner + .analytics + .status_changed(&goal, previous_status, GoalEventAttribution::NoTurn); + let objective_changed = previous_goal.as_ref().is_some_and(|previous_goal| { + !replaced_existing_goal && previous_goal.objective != goal.objective + }); + match goal.status { + codex_state::ThreadGoalStatus::Active => { + if self.inner.accounting_state.current_turn_id().is_some() { + let _ = self + .inner + .accounting_state + .mark_current_turn_goal_active(goal.goal_id.clone()); + } else { + self.inner + .accounting_state + .mark_idle_goal_active(goal.goal_id.clone()); + } + if objective_changed { + let item = objective_updated_steering_item(&protocol_goal_from_state(goal)); + self.inject_active_turn_steering(item).await; + } + self.continue_if_idle().await?; + } + codex_state::ThreadGoalStatus::BudgetLimited => { + if self.inner.accounting_state.current_turn_id().is_none() { + self.inner.accounting_state.clear_active_goal(); + } + } + codex_state::ThreadGoalStatus::Paused + | codex_state::ThreadGoalStatus::Blocked + | codex_state::ThreadGoalStatus::UsageLimited + | codex_state::ThreadGoalStatus::Complete => { + self.inner.accounting_state.clear_active_goal(); + } + } + Ok(()) + } + + pub async fn apply_external_goal_clear( + &self, + goal: codex_state::ThreadGoal, + ) -> Result<(), String> { + if !self.is_enabled() { + return Ok(()); + } + + self.inner.analytics.cleared(&goal); + self.inner.accounting_state.clear_active_goal(); + Ok(()) + } + + pub async fn usage_limit_active_goal_for_turn(&self, turn_id: &str) -> Result<(), String> { + self.stop_active_goal_for_turn(turn_id, ActiveGoalStopReason::UsageLimit) + .await + } + + /// Accounts the ending turn and stops its active goal after an error or repeated empty output. + pub(crate) async fn stop_active_goal_for_turn( + &self, + turn_id: &str, + reason: ActiveGoalStopReason, + ) -> Result<(), String> { + if !self.is_enabled() { + return Ok(()); + } + + // Hold this through accounting and the status update so external goal + // mutations and idle continuation cannot interleave between them. + let _goal_state_permit = self.goal_state_permit().await?; + let Some(accounting_goal_id) = self + .inner + .accounting_state + .current_active_goal_id_for_turn(turn_id) + else { + return Ok(()); + }; + if let ActiveGoalStopReason::ExecutionUnavailable { expected_goal_id } = &reason + && accounting_goal_id != *expected_goal_id + { + return Ok(()); + } + + let (event_name, status, expected_goal_id) = match reason { + ActiveGoalStopReason::TurnError => { + ("turn-error", codex_state::ThreadGoalStatus::Blocked, None) + } + ActiveGoalStopReason::UsageLimit => ( + "usage-limit", + codex_state::ThreadGoalStatus::UsageLimited, + None, + ), + ActiveGoalStopReason::EmptyResponse => { + let Some(expected_goal_id) = + self.inner.accounting_state.empty_response_goal(turn_id) + else { + return Ok(()); + }; + if accounting_goal_id != expected_goal_id { + return Ok(()); + } + ( + "empty-response", + codex_state::ThreadGoalStatus::Blocked, + Some(expected_goal_id), + ) + } + ActiveGoalStopReason::ExecutionUnavailable { expected_goal_id } => ( + "execution-unavailable", + codex_state::ThreadGoalStatus::Blocked, + Some(expected_goal_id), + ), + }; + self.account_active_goal_progress( + turn_id, + &format!("{turn_id}:{event_name}-progress"), + codex_state::GoalAccountingMode::ActiveOnly, + BudgetLimitedGoalDisposition::ClearActive, + ) + .await?; + + let Some(active_goal) = self + .inner + .state_dbs + .thread_goals() + .get_thread_goal(self.thread_id()) + .await + .map_err(|err| err.to_string())? + else { + self.inner.accounting_state.clear_active_goal(); + return Ok(()); + }; + if expected_goal_id + .as_ref() + .is_some_and(|expected_goal_id| active_goal.goal_id != *expected_goal_id) + { + return Ok(()); + } + let can_stop = active_goal.status == codex_state::ThreadGoalStatus::Active + || (active_goal.status == codex_state::ThreadGoalStatus::BudgetLimited + && status == codex_state::ThreadGoalStatus::UsageLimited); + if !can_stop { + self.inner.accounting_state.clear_active_goal(); + return Ok(()); + } + let previous_status = Some(active_goal.status); + let Some(goal) = self + .inner + .state_dbs + .thread_goals() + .update_thread_goal( + self.thread_id(), + codex_state::GoalUpdate { + objective: None, + status: Some(status), + token_budget: None, + expected_goal_id: Some(active_goal.goal_id), + }, + ) + .await + .map_err(|err| err.to_string())? + else { + return Ok(()); + }; + self.inner + .metrics + .record_terminal_if_status_changed(previous_status, &goal); + self.inner.analytics.status_changed( + &goal, + previous_status, + GoalEventAttribution::Turn(turn_id), + ); + self.inner.accounting_state.clear_active_goal(); + let goal = protocol_goal_from_state(goal); + self.inner.event_emitter.thread_goal_updated( + format!("{turn_id}:{event_name}"), + Some(turn_id.to_string()), + goal, + ); + Ok(()) + } + + pub async fn restore_after_resume(&self) -> Result<(), String> { + if !self.is_enabled() { + return Ok(()); + } + + let goal = self + .inner + .state_dbs + .thread_goals() + .get_thread_goal(self.thread_id()) + .await + .map_err(|err| err.to_string())?; + match goal { + Some(goal) if goal.status == codex_state::ThreadGoalStatus::Active => { + self.inner + .accounting_state + .mark_idle_goal_active(goal.goal_id); + self.inner.metrics.record_resumed(); + } + Some(_) | None => self.inner.accounting_state.clear_active_goal(), + } + Ok(()) + } + + pub(crate) async fn continue_if_idle(&self) -> Result<(), String> { + if !self.tools_available() { + self.inner.accounting_state.clear_active_goal(); + return Ok(()); + } + // Hold this through the read/start window so external set/clear cannot + // change the goal after we read it but before the continuation launches. + let _goal_state_permit = self.goal_state_permit().await?; + + if self + .inner + .state_dbs + .thread_goals() + .has_thread_goal_continuation_deferral(self.thread_id()) + .await + .map_err(|err| err.to_string())? + { + return Ok(()); + } + + let Some(thread_manager) = self.inner.thread_manager.upgrade() else { + tracing::debug!("skipping goal continuation because thread manager is unavailable"); + return Ok(()); + }; + let Ok(thread) = thread_manager.get_thread(self.inner.thread_id).await else { + tracing::debug!("skipping goal continuation because live thread is unavailable"); + return Ok(()); + }; + + let Some(goal) = self + .inner + .state_dbs + .thread_goals() + .get_thread_goal(self.thread_id()) + .await + .map_err(|err| err.to_string())? + else { + self.inner.accounting_state.clear_active_goal(); + return Ok(()); + }; + if goal.status != codex_state::ThreadGoalStatus::Active { + self.inner.accounting_state.clear_active_goal(); + return Ok(()); + } + let start_options = thread + .thread_extension_data() + .get::() + .map(|options| options.as_ref().clone()) + .unwrap_or_default(); + let item = continuation_steering_item( + &protocol_goal_from_state(goal), + thread.config().await.update_plan_enabled, + ); + + match thread + .start_turn_if_idle( + TurnInputRequest::new(TurnInput::ResponseItem(item)).on_start(TurnStartOptions { + turn_trigger: Some("goal".to_string()), + ..start_options + }), + ) + .await + { + Ok(StartIfIdleSubmission::Started { turn_id }) => { + // Turn-stop evaluation takes the same permit, so even a fast response + // cannot finish before this host-admitted continuation is identified. + self.inner.accounting_state.mark_goal_continuation(turn_id); + } + Ok(StartIfIdleSubmission::NotSubmitted { reason }) => { + tracing::debug!( + ?reason, + "skipping goal continuation because automatic idle work was rejected" + ); + } + Err(error) => { + tracing::debug!( + %error, + "skipping goal continuation because turn input submission failed" + ); + } + } + + let current_turn_is_goal_active = self + .inner + .accounting_state + .current_turn_id() + .is_some_and(|turn_id| { + self.inner + .accounting_state + .current_active_goal_id_for_turn(turn_id.as_str()) + .is_some() + }); + if !current_turn_is_goal_active { + self.inner + .accounting_state + .reset_idle_progress_baseline_and_clear_active_goal(); + } + Ok(()) + } + + pub(crate) async fn inject_active_turn_steering(&self, item: ResponseItem) { + let Some(thread_manager) = self.inner.thread_manager.upgrade() else { + tracing::debug!("skipping goal steering because thread manager is unavailable"); + return; + }; + let Ok(thread) = thread_manager.get_thread(self.inner.thread_id).await else { + tracing::debug!("skipping goal steering because live thread is unavailable"); + return; + }; + if thread.inject_if_running(vec![item]).await.is_err() { + tracing::debug!("skipping goal steering because no turn is active"); + } + } + + pub(crate) async fn account_active_goal_progress( + &self, + turn_id: &str, + event_id: &str, + mode: codex_state::GoalAccountingMode, + budget_limited_goal_disposition: BudgetLimitedGoalDisposition, + ) -> Result, String> { + let accounting = self.accounting_state(); + let _accounting_permit = accounting + .progress_accounting_permit() + .await + .map_err(|err| err.to_string())?; + let Some(snapshot) = accounting.progress_snapshot(turn_id) else { + return Ok(None); + }; + let previous_status = self + .current_goal_status_for_metrics(Some(snapshot.expected_goal_id.as_str())) + .await?; + let outcome = self + .inner + .state_dbs + .thread_goals() + .account_thread_goal_usage( + self.thread_id(), + snapshot.time_delta_seconds, + snapshot.token_delta, + mode, + Some(snapshot.expected_goal_id.as_str()), + ) + .await + .map_err(|err| err.to_string())?; + Ok(match outcome { + codex_state::GoalAccountingOutcome::Updated(goal) => { + let goal_id = goal.goal_id.clone(); + self.inner + .metrics + .record_terminal_if_status_changed(previous_status, &goal); + self.inner + .analytics + .usage_accounted(&goal, GoalEventAttribution::Turn(turn_id)); + self.inner.analytics.status_changed( + &goal, + previous_status, + GoalEventAttribution::Turn(turn_id), + ); + accounting.mark_progress_accounted_for_status( + turn_id, + &snapshot, + goal.status, + budget_limited_goal_disposition, + ); + let goal = protocol_goal_from_state(goal); + self.inner.event_emitter.thread_goal_updated( + event_id.to_string(), + Some(turn_id.to_string()), + goal.clone(), + ); + Some(AccountedGoalProgress { goal, goal_id }) + } + codex_state::GoalAccountingOutcome::Unchanged(_) => None, + }) + } + + async fn account_idle_goal_progress( + &self, + event_id: &str, + mode: codex_state::GoalAccountingMode, + budget_limited_goal_disposition: BudgetLimitedGoalDisposition, + ) -> Result, String> { + let accounting = self.accounting_state(); + let _accounting_permit = accounting + .progress_accounting_permit() + .await + .map_err(|err| err.to_string())?; + let Some(snapshot) = accounting.idle_progress_snapshot() else { + return Ok(None); + }; + let previous_status = self + .current_goal_status_for_metrics(Some(snapshot.expected_goal_id.as_str())) + .await?; + let outcome = self + .inner + .state_dbs + .thread_goals() + .account_thread_goal_usage( + self.thread_id(), + snapshot.time_delta_seconds, + snapshot.token_delta, + mode, + Some(snapshot.expected_goal_id.as_str()), + ) + .await + .map_err(|err| err.to_string())?; + Ok(match outcome { + codex_state::GoalAccountingOutcome::Updated(goal) => { + let goal_id = goal.goal_id.clone(); + self.inner + .metrics + .record_terminal_if_status_changed(previous_status, &goal); + self.inner + .analytics + .usage_accounted(&goal, GoalEventAttribution::NoTurn); + self.inner.analytics.status_changed( + &goal, + previous_status, + GoalEventAttribution::NoTurn, + ); + accounting.mark_idle_progress_accounted_for_status( + &snapshot, + goal.status, + budget_limited_goal_disposition, + ); + let goal = protocol_goal_from_state(goal); + self.inner.event_emitter.thread_goal_updated( + event_id.to_string(), + /*turn_id*/ None, + goal.clone(), + ); + Some(AccountedGoalProgress { goal, goal_id }) + } + codex_state::GoalAccountingOutcome::Unchanged(_) => { + accounting.reset_idle_progress_baseline_and_clear_active_goal(); + None + } + }) + } + + async fn current_goal_status_for_metrics( + &self, + expected_goal_id: Option<&str>, + ) -> Result, String> { + let goal = self + .inner + .state_dbs + .thread_goals() + .get_thread_goal(self.thread_id()) + .await + .map_err(|err| err.to_string())?; + Ok(goal.and_then(|goal| { + expected_goal_id + .is_none_or(|expected_goal_id| goal.goal_id == expected_goal_id) + .then_some(goal.status) + })) + } +} diff --git a/codex-rs/ext/goal/src/spec.rs b/codex-rs/ext/goal/src/spec.rs new file mode 100644 index 0000000000000000000000000000000000000000..6d1d7a9230083a554bda372341337108a546c90a --- /dev/null +++ b/codex-rs/ext/goal/src/spec.rs @@ -0,0 +1,94 @@ +//! Responses API tool definitions for persisted thread goals. + +use codex_tools::JsonSchema; +use codex_tools::ResponsesApiTool; +use codex_tools::ToolSpec; +use serde_json::json; +use std::collections::BTreeMap; + +pub const GET_GOAL_TOOL_NAME: &str = "get_goal"; +pub const CREATE_GOAL_TOOL_NAME: &str = "create_goal"; +pub const UPDATE_GOAL_TOOL_NAME: &str = "update_goal"; + +pub fn create_get_goal_tool() -> ToolSpec { + ToolSpec::Function(ResponsesApiTool { + name: GET_GOAL_TOOL_NAME.to_string(), + description: "Get the current goal for this thread, including status, budgets, token and elapsed-time usage, and remaining token budget." + .to_string(), + strict: false, + defer_loading: None, + parameters: JsonSchema::object(BTreeMap::new(), Some(Vec::new()), Some(false.into())), + output_schema: None, + }) +} + +pub fn create_create_goal_tool() -> ToolSpec { + let properties = BTreeMap::from([ + ( + "objective".to_string(), + JsonSchema::string(Some( + "Required. The concrete objective to start pursuing. This starts a new active goal when no goal exists or replaces the current goal when it is complete." + .to_string(), + )), + ), + ( + "token_budget".to_string(), + JsonSchema::integer(Some( + "Positive token budget for the new goal. Omit unless explicitly requested." + .to_string(), + )), + ), + ]); + + ToolSpec::Function(ResponsesApiTool { + name: CREATE_GOAL_TOOL_NAME.to_string(), + description: format!( + r#"Create a goal only when explicitly requested by the user or system/developer instructions; do not infer goals from ordinary tasks. +Set token_budget only when an explicit token budget is requested. Fails if an unfinished goal exists; use {UPDATE_GOAL_TOOL_NAME} only for status."# + ), + strict: false, + defer_loading: None, + parameters: JsonSchema::object( + properties, + /*required*/ Some(vec!["objective".to_string()]), + Some(false.into()), + ), + output_schema: None, + }) +} + +pub fn create_update_goal_tool() -> ToolSpec { + let properties = BTreeMap::from([( + "status".to_string(), + JsonSchema::string_enum( + vec![json!("complete"), json!("blocked"), json!("paused")], + Some( + "Required. `paused` requires an explicit user request. Set to `complete` only when the objective is achieved and no required work remains. Set to `blocked` only after the same blocking condition has recurred for at least three consecutive goal turns and the agent is at an impasse. After a previously blocked goal is resumed, the resumed run starts a fresh blocked audit." + .to_string(), + ), + ), + )]); + + ToolSpec::Function(ResponsesApiTool { + name: UPDATE_GOAL_TOOL_NAME.to_string(), + description: r#"Update the existing goal. +Set status to `paused` only at the user's explicit request to pause this goal, never on your own initiative. Ask if unclear; a later resume revokes that request. Report the returned status and stop goal work. Budget limits take precedence over pausing. +Set status to `complete` only when the objective has actually been achieved and no required work remains. +Set status to `blocked` only when the same blocking condition has repeated for at least three consecutive goal turns, counting the original/user-triggered turn and any automatic continuations, and the agent cannot make meaningful progress without user input or an external-state change. +If the user resumes a goal that was previously marked `blocked`, treat the resumed run as a fresh blocked audit. If the same blocking condition then repeats for at least three consecutive resumed goal turns, set status to `blocked` again. +Once the blocked threshold is satisfied, do not keep reporting that you are still blocked while leaving the goal active; set status to `blocked`. +Do not use `blocked` merely because the work is hard, slow, uncertain, incomplete, or would benefit from clarification. +Do not mark a goal complete merely because its budget is nearly exhausted or because you are stopping work. +You cannot use this tool to resume, budget-limit, or usage-limit a goal; those status changes are controlled by the user or system. +When marking a budgeted goal achieved with status `complete`, report the final token usage from the tool result to the user."# + .to_string(), + strict: false, + defer_loading: None, + parameters: JsonSchema::object( + properties, + /*required*/ Some(vec!["status".to_string()]), + Some(false.into()), + ), + output_schema: None, + }) +} diff --git a/codex-rs/ext/goal/src/steering.rs b/codex-rs/ext/goal/src/steering.rs new file mode 100644 index 0000000000000000000000000000000000000000..88de89d09e177ab3419157563579fca03f4954e0 --- /dev/null +++ b/codex-rs/ext/goal/src/steering.rs @@ -0,0 +1,145 @@ +use codex_core::context::ContextualUserFragment; +use codex_core::context::InternalContextSource; +use codex_core::context::InternalModelContextFragment; +use codex_core::context::without_update_plan_instructions; +use codex_protocol::models::ResponseItem; +use codex_protocol::protocol::ThreadGoal; +use codex_utils_template::Template; +use std::sync::LazyLock; + +static CONTINUATION_PROMPT_TEMPLATE: LazyLock