Add files using upload-large-folder tool
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- codex-rs/agent-roles/src/agent_role_config.rs +209 -0
- codex-rs/agent-roles/src/discovery.rs +40 -0
- codex-rs/agent-roles/src/lib.rs +8 -0
- codex-rs/agent-roles/src/loader.rs +335 -0
- codex-rs/app-server-protocol-noop-macros/src/lib.rs +20 -0
- codex-rs/cloud-tasks-mock-client/src/lib.rs +3 -0
- codex-rs/cloud-tasks-mock-client/src/mock.rs +267 -0
- codex-rs/code-mode-protocol/src/description.rs +993 -0
- codex-rs/code-mode-protocol/src/grpc/codex.code_mode.v1.proto +269 -0
- codex-rs/code-mode-protocol/src/grpc/mod.rs +7 -0
- codex-rs/code-mode-protocol/src/host/codec.rs +170 -0
- codex-rs/code-mode-protocol/src/host/codec_tests.rs +137 -0
- codex-rs/code-mode-protocol/src/host/error.rs +19 -0
- codex-rs/code-mode-protocol/src/host/host_tests.rs +847 -0
- codex-rs/code-mode-protocol/src/host/message.rs +264 -0
- codex-rs/code-mode-protocol/src/host/mod.rs +62 -0
- codex-rs/code-mode-protocol/src/host/payload.rs +492 -0
- codex-rs/code-mode-protocol/src/host/types.rs +248 -0
- codex-rs/code-mode-protocol/src/json_schema_types.rs +538 -0
- codex-rs/code-mode-protocol/src/json_schema_types_tests.rs +200 -0
- codex-rs/code-mode-protocol/src/lib.rs +52 -0
- codex-rs/code-mode-protocol/src/response.rs +29 -0
- codex-rs/code-mode-protocol/src/runtime.rs +183 -0
- codex-rs/code-mode-protocol/src/runtime_tests.rs +107 -0
- codex-rs/code-mode-protocol/src/session.rs +200 -0
- codex-rs/code-mode-protocol/src/session_tests.rs +19 -0
- codex-rs/collaboration-mode-templates/src/lib.rs +2 -0
- codex-rs/collaboration-mode-templates/templates/default.md +19 -0
- codex-rs/collaboration-mode-templates/templates/plan.md +128 -0
- codex-rs/exec-server/src/arg0_exec_helper.rs +31 -0
- codex-rs/exec-server/src/capability_discovery.rs +510 -0
- codex-rs/exec-server/src/capability_discovery_cache.rs +246 -0
- codex-rs/exec-server/src/client.rs +0 -0
- codex-rs/exec-server/src/client/accepted.rs +266 -0
- codex-rs/exec-server/src/client/accepted_tests.rs +34 -0
- codex-rs/exec-server/src/client/http_client.rs +26 -0
- codex-rs/exec-server/src/client/http_response_body_stream.rs +446 -0
- codex-rs/exec-server/src/client/network_policy_audit.rs +81 -0
- codex-rs/exec-server/src/client/route_aware_http_client.rs +378 -0
- codex-rs/exec-server/src/client/rpc_http_client.rs +92 -0
- codex-rs/exec-server/src/client/tests/network_policy_tests.rs +544 -0
- codex-rs/exec-server/src/client_api.rs +187 -0
- codex-rs/exec-server/src/client_recovery.rs +900 -0
- codex-rs/exec-server/src/client_recovery_tests.rs +262 -0
- codex-rs/exec-server/src/client_refresh.rs +269 -0
- codex-rs/exec-server/src/client_refresh_tests.rs +621 -0
- codex-rs/exec-server/src/client_telemetry.rs +23 -0
- codex-rs/exec-server/src/client_transport.rs +814 -0
- codex-rs/exec-server/src/client_transport_tests.rs +581 -0
- codex-rs/exec-server/src/connection.rs +1042 -0
codex-rs/agent-roles/src/agent_role_config.rs
ADDED
|
@@ -0,0 +1,209 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use codex_config::config_toml::ConfigToml;
|
| 2 |
+
use codex_utils_absolute_path::AbsolutePathBufGuard;
|
| 3 |
+
use serde::Deserialize;
|
| 4 |
+
use std::collections::BTreeSet;
|
| 5 |
+
use std::path::Path;
|
| 6 |
+
use std::path::PathBuf;
|
| 7 |
+
use toml::Value as TomlValue;
|
| 8 |
+
|
| 9 |
+
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
| 10 |
+
pub struct AgentRoleConfig {
|
| 11 |
+
/// Human-facing role documentation used in spawn tool guidance.
|
| 12 |
+
/// Required for loaded user-defined roles after deprecated/new metadata precedence resolves.
|
| 13 |
+
pub description: Option<String>,
|
| 14 |
+
/// Path to a role-specific config layer.
|
| 15 |
+
pub config_file: Option<PathBuf>,
|
| 16 |
+
/// Candidate nicknames for agents spawned with this role.
|
| 17 |
+
pub nickname_candidates: Option<Vec<String>>,
|
| 18 |
+
}
|
| 19 |
+
|
| 20 |
+
#[derive(Deserialize, Debug, Clone, Default, PartialEq)]
|
| 21 |
+
#[serde(deny_unknown_fields)]
|
| 22 |
+
struct RawAgentRoleFileToml {
|
| 23 |
+
name: Option<String>,
|
| 24 |
+
description: Option<String>,
|
| 25 |
+
nickname_candidates: Option<Vec<String>>,
|
| 26 |
+
#[serde(flatten)]
|
| 27 |
+
config: ConfigToml,
|
| 28 |
+
}
|
| 29 |
+
|
| 30 |
+
#[derive(Debug, Clone, PartialEq)]
|
| 31 |
+
pub struct ResolvedAgentRoleFile {
|
| 32 |
+
pub role_name: String,
|
| 33 |
+
pub description: Option<String>,
|
| 34 |
+
pub nickname_candidates: Option<Vec<String>>,
|
| 35 |
+
pub config: TomlValue,
|
| 36 |
+
}
|
| 37 |
+
|
| 38 |
+
pub fn parse_agent_role_file_contents(
|
| 39 |
+
contents: &str,
|
| 40 |
+
role_file_label: &Path,
|
| 41 |
+
config_base_dir: &Path,
|
| 42 |
+
role_name_hint: Option<&str>,
|
| 43 |
+
) -> std::io::Result<ResolvedAgentRoleFile> {
|
| 44 |
+
let role_file_toml: TomlValue = toml::from_str(contents).map_err(|err| {
|
| 45 |
+
std::io::Error::new(
|
| 46 |
+
std::io::ErrorKind::InvalidData,
|
| 47 |
+
format!(
|
| 48 |
+
"failed to parse agent role file at {}: {err}",
|
| 49 |
+
role_file_label.display()
|
| 50 |
+
),
|
| 51 |
+
)
|
| 52 |
+
})?;
|
| 53 |
+
let _guard = AbsolutePathBufGuard::new(config_base_dir);
|
| 54 |
+
let parsed: RawAgentRoleFileToml = role_file_toml.clone().try_into().map_err(|err| {
|
| 55 |
+
std::io::Error::new(
|
| 56 |
+
std::io::ErrorKind::InvalidData,
|
| 57 |
+
format!(
|
| 58 |
+
"failed to deserialize agent role file at {}: {err}",
|
| 59 |
+
role_file_label.display()
|
| 60 |
+
),
|
| 61 |
+
)
|
| 62 |
+
})?;
|
| 63 |
+
let description = normalize_agent_role_description(
|
| 64 |
+
&format!("agent role file {}.description", role_file_label.display()),
|
| 65 |
+
parsed.description.as_deref(),
|
| 66 |
+
)?;
|
| 67 |
+
validate_agent_role_file_developer_instructions(
|
| 68 |
+
role_file_label,
|
| 69 |
+
parsed.config.developer_instructions.as_deref(),
|
| 70 |
+
role_name_hint.is_none(),
|
| 71 |
+
)?;
|
| 72 |
+
|
| 73 |
+
let role_name = parsed
|
| 74 |
+
.name
|
| 75 |
+
.as_deref()
|
| 76 |
+
.map(str::trim)
|
| 77 |
+
.filter(|name| !name.is_empty())
|
| 78 |
+
.map(ToOwned::to_owned)
|
| 79 |
+
.or_else(|| role_name_hint.map(ToOwned::to_owned))
|
| 80 |
+
.ok_or_else(|| {
|
| 81 |
+
std::io::Error::new(
|
| 82 |
+
std::io::ErrorKind::InvalidInput,
|
| 83 |
+
format!(
|
| 84 |
+
"agent role file at {} must define a non-empty `name`",
|
| 85 |
+
role_file_label.display()
|
| 86 |
+
),
|
| 87 |
+
)
|
| 88 |
+
})?;
|
| 89 |
+
|
| 90 |
+
let nickname_candidates = normalize_agent_role_nickname_candidates(
|
| 91 |
+
&format!(
|
| 92 |
+
"agent role file {}.nickname_candidates",
|
| 93 |
+
role_file_label.display()
|
| 94 |
+
),
|
| 95 |
+
parsed.nickname_candidates.as_deref(),
|
| 96 |
+
)?;
|
| 97 |
+
|
| 98 |
+
let mut config = role_file_toml;
|
| 99 |
+
let Some(config_table) = config.as_table_mut() else {
|
| 100 |
+
return Err(std::io::Error::new(
|
| 101 |
+
std::io::ErrorKind::InvalidData,
|
| 102 |
+
format!(
|
| 103 |
+
"agent role file at {} must contain a TOML table",
|
| 104 |
+
role_file_label.display()
|
| 105 |
+
),
|
| 106 |
+
));
|
| 107 |
+
};
|
| 108 |
+
config_table.remove("name");
|
| 109 |
+
config_table.remove("description");
|
| 110 |
+
config_table.remove("nickname_candidates");
|
| 111 |
+
|
| 112 |
+
Ok(ResolvedAgentRoleFile {
|
| 113 |
+
role_name,
|
| 114 |
+
description,
|
| 115 |
+
nickname_candidates,
|
| 116 |
+
config,
|
| 117 |
+
})
|
| 118 |
+
}
|
| 119 |
+
|
| 120 |
+
pub(crate) fn normalize_agent_role_description(
|
| 121 |
+
field_label: &str,
|
| 122 |
+
description: Option<&str>,
|
| 123 |
+
) -> std::io::Result<Option<String>> {
|
| 124 |
+
match description.map(str::trim) {
|
| 125 |
+
Some("") => Err(std::io::Error::new(
|
| 126 |
+
std::io::ErrorKind::InvalidInput,
|
| 127 |
+
format!("{field_label} cannot be blank"),
|
| 128 |
+
)),
|
| 129 |
+
Some(description) => Ok(Some(description.to_string())),
|
| 130 |
+
None => Ok(None),
|
| 131 |
+
}
|
| 132 |
+
}
|
| 133 |
+
|
| 134 |
+
fn validate_agent_role_file_developer_instructions(
|
| 135 |
+
role_file_label: &Path,
|
| 136 |
+
developer_instructions: Option<&str>,
|
| 137 |
+
require_present: bool,
|
| 138 |
+
) -> std::io::Result<()> {
|
| 139 |
+
match developer_instructions.map(str::trim) {
|
| 140 |
+
Some("") => Err(std::io::Error::new(
|
| 141 |
+
std::io::ErrorKind::InvalidInput,
|
| 142 |
+
format!(
|
| 143 |
+
"agent role file at {}.developer_instructions cannot be blank",
|
| 144 |
+
role_file_label.display()
|
| 145 |
+
),
|
| 146 |
+
)),
|
| 147 |
+
Some(_) => Ok(()),
|
| 148 |
+
None if require_present => Err(std::io::Error::new(
|
| 149 |
+
std::io::ErrorKind::InvalidInput,
|
| 150 |
+
format!(
|
| 151 |
+
"agent role file at {} must define `developer_instructions`",
|
| 152 |
+
role_file_label.display()
|
| 153 |
+
),
|
| 154 |
+
)),
|
| 155 |
+
None => Ok(()),
|
| 156 |
+
}
|
| 157 |
+
}
|
| 158 |
+
|
| 159 |
+
pub(crate) fn normalize_agent_role_nickname_candidates(
|
| 160 |
+
field_label: &str,
|
| 161 |
+
nickname_candidates: Option<&[String]>,
|
| 162 |
+
) -> std::io::Result<Option<Vec<String>>> {
|
| 163 |
+
let Some(nickname_candidates) = nickname_candidates else {
|
| 164 |
+
return Ok(None);
|
| 165 |
+
};
|
| 166 |
+
|
| 167 |
+
if nickname_candidates.is_empty() {
|
| 168 |
+
return Err(std::io::Error::new(
|
| 169 |
+
std::io::ErrorKind::InvalidInput,
|
| 170 |
+
format!("{field_label} must contain at least one name"),
|
| 171 |
+
));
|
| 172 |
+
}
|
| 173 |
+
|
| 174 |
+
let mut normalized_candidates = Vec::with_capacity(nickname_candidates.len());
|
| 175 |
+
let mut seen_candidates = BTreeSet::new();
|
| 176 |
+
|
| 177 |
+
for nickname in nickname_candidates {
|
| 178 |
+
let normalized_nickname = nickname.trim();
|
| 179 |
+
if normalized_nickname.is_empty() {
|
| 180 |
+
return Err(std::io::Error::new(
|
| 181 |
+
std::io::ErrorKind::InvalidInput,
|
| 182 |
+
format!("{field_label} cannot contain blank names"),
|
| 183 |
+
));
|
| 184 |
+
}
|
| 185 |
+
|
| 186 |
+
if !seen_candidates.insert(normalized_nickname.to_owned()) {
|
| 187 |
+
return Err(std::io::Error::new(
|
| 188 |
+
std::io::ErrorKind::InvalidInput,
|
| 189 |
+
format!("{field_label} cannot contain duplicates"),
|
| 190 |
+
));
|
| 191 |
+
}
|
| 192 |
+
|
| 193 |
+
if !normalized_nickname
|
| 194 |
+
.chars()
|
| 195 |
+
.all(|c| c.is_ascii_alphanumeric() || matches!(c, ' ' | '-' | '_'))
|
| 196 |
+
{
|
| 197 |
+
return Err(std::io::Error::new(
|
| 198 |
+
std::io::ErrorKind::InvalidInput,
|
| 199 |
+
format!(
|
| 200 |
+
"{field_label} may only contain ASCII letters, digits, spaces, hyphens, and underscores"
|
| 201 |
+
),
|
| 202 |
+
));
|
| 203 |
+
}
|
| 204 |
+
|
| 205 |
+
normalized_candidates.push(normalized_nickname.to_owned());
|
| 206 |
+
}
|
| 207 |
+
|
| 208 |
+
Ok(Some(normalized_candidates))
|
| 209 |
+
}
|
codex-rs/agent-roles/src/discovery.rs
ADDED
|
@@ -0,0 +1,40 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use codex_file_system::ExecutorFileSystem;
|
| 2 |
+
use codex_utils_absolute_path::AbsolutePathBuf;
|
| 3 |
+
use codex_utils_path_uri::PathUri;
|
| 4 |
+
use std::io;
|
| 5 |
+
use std::io::ErrorKind;
|
| 6 |
+
|
| 7 |
+
pub(crate) async fn collect_agent_role_files(
|
| 8 |
+
fs: &dyn ExecutorFileSystem,
|
| 9 |
+
dir: &AbsolutePathBuf,
|
| 10 |
+
) -> io::Result<Vec<AbsolutePathBuf>> {
|
| 11 |
+
let mut files = Vec::new();
|
| 12 |
+
let mut dirs = vec![dir.clone()];
|
| 13 |
+
while let Some(dir) = dirs.pop() {
|
| 14 |
+
let dir_uri = PathUri::from_abs_path(&dir);
|
| 15 |
+
let entries = match fs.read_directory(&dir_uri, /*sandbox*/ None).await {
|
| 16 |
+
Ok(entries) => entries,
|
| 17 |
+
Err(err) if err.kind() == ErrorKind::NotFound => continue,
|
| 18 |
+
Err(err) => return Err(err),
|
| 19 |
+
};
|
| 20 |
+
|
| 21 |
+
for entry in entries {
|
| 22 |
+
let path = dir.join(entry.file_name);
|
| 23 |
+
if entry.is_directory {
|
| 24 |
+
dirs.push(path);
|
| 25 |
+
continue;
|
| 26 |
+
}
|
| 27 |
+
if entry.is_file
|
| 28 |
+
&& path
|
| 29 |
+
.as_path()
|
| 30 |
+
.extension()
|
| 31 |
+
.is_some_and(|extension| extension == "toml")
|
| 32 |
+
{
|
| 33 |
+
files.push(path);
|
| 34 |
+
}
|
| 35 |
+
}
|
| 36 |
+
}
|
| 37 |
+
|
| 38 |
+
files.sort();
|
| 39 |
+
Ok(files)
|
| 40 |
+
}
|
codex-rs/agent-roles/src/lib.rs
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
mod agent_role_config;
|
| 2 |
+
mod discovery;
|
| 3 |
+
mod loader;
|
| 4 |
+
|
| 5 |
+
pub use agent_role_config::AgentRoleConfig;
|
| 6 |
+
pub use agent_role_config::ResolvedAgentRoleFile;
|
| 7 |
+
pub use agent_role_config::parse_agent_role_file_contents;
|
| 8 |
+
pub use loader::load_agent_roles;
|
codex-rs/agent-roles/src/loader.rs
ADDED
|
@@ -0,0 +1,335 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use crate::AgentRoleConfig;
|
| 2 |
+
use crate::ResolvedAgentRoleFile;
|
| 3 |
+
use crate::agent_role_config::normalize_agent_role_description;
|
| 4 |
+
use crate::agent_role_config::normalize_agent_role_nickname_candidates;
|
| 5 |
+
use crate::discovery::collect_agent_role_files;
|
| 6 |
+
use crate::parse_agent_role_file_contents;
|
| 7 |
+
use codex_config::ConfigLayerStack;
|
| 8 |
+
use codex_config::config_toml::AgentRoleToml;
|
| 9 |
+
use codex_config::config_toml::AgentsToml;
|
| 10 |
+
use codex_config::config_toml::ConfigToml;
|
| 11 |
+
use codex_file_system::ExecutorFileSystem;
|
| 12 |
+
use codex_file_system::GetMetadataOptions;
|
| 13 |
+
use codex_file_system::ReadFileOptions;
|
| 14 |
+
use codex_utils_absolute_path::AbsolutePathBuf;
|
| 15 |
+
use codex_utils_absolute_path::AbsolutePathBufGuard;
|
| 16 |
+
use codex_utils_path_uri::PathUri;
|
| 17 |
+
use std::collections::BTreeMap;
|
| 18 |
+
use std::collections::BTreeSet;
|
| 19 |
+
use std::path::Path;
|
| 20 |
+
use std::path::PathBuf;
|
| 21 |
+
use toml::Value as TomlValue;
|
| 22 |
+
|
| 23 |
+
pub async fn load_agent_roles(
|
| 24 |
+
fs: &dyn ExecutorFileSystem,
|
| 25 |
+
cfg: &ConfigToml,
|
| 26 |
+
config_layer_stack: &ConfigLayerStack,
|
| 27 |
+
startup_warnings: &mut Vec<String>,
|
| 28 |
+
) -> std::io::Result<BTreeMap<String, AgentRoleConfig>> {
|
| 29 |
+
let mut layers = config_layer_stack.layers_low_to_high().peekable();
|
| 30 |
+
if layers.peek().is_none() {
|
| 31 |
+
return load_agent_roles_without_layers(fs, cfg).await;
|
| 32 |
+
}
|
| 33 |
+
|
| 34 |
+
let mut roles: BTreeMap<String, AgentRoleConfig> = BTreeMap::new();
|
| 35 |
+
for layer in layers {
|
| 36 |
+
let mut layer_roles: BTreeMap<String, AgentRoleConfig> = BTreeMap::new();
|
| 37 |
+
let mut declared_role_files = BTreeSet::new();
|
| 38 |
+
let config_folder = layer.config_folder();
|
| 39 |
+
let agents_toml = match agents_toml_from_layer(&layer.config, config_folder.as_deref()) {
|
| 40 |
+
Ok(agents_toml) => agents_toml,
|
| 41 |
+
Err(err) => {
|
| 42 |
+
push_agent_role_warning(startup_warnings, err);
|
| 43 |
+
None
|
| 44 |
+
}
|
| 45 |
+
};
|
| 46 |
+
if let Some(agents_toml) = agents_toml {
|
| 47 |
+
for (declared_role_name, role_toml) in &agents_toml.roles {
|
| 48 |
+
let (role_name, role) =
|
| 49 |
+
match read_declared_role(fs, declared_role_name, role_toml).await {
|
| 50 |
+
Ok(role) => role,
|
| 51 |
+
Err(err) => {
|
| 52 |
+
push_agent_role_warning(startup_warnings, err);
|
| 53 |
+
continue;
|
| 54 |
+
}
|
| 55 |
+
};
|
| 56 |
+
if let Some(config_file) = role.config_file.clone() {
|
| 57 |
+
declared_role_files.insert(config_file);
|
| 58 |
+
}
|
| 59 |
+
if layer_roles.contains_key(&role_name) {
|
| 60 |
+
push_agent_role_warning(
|
| 61 |
+
startup_warnings,
|
| 62 |
+
std::io::Error::new(
|
| 63 |
+
std::io::ErrorKind::InvalidInput,
|
| 64 |
+
format!(
|
| 65 |
+
"duplicate agent role name `{role_name}` declared in the same config layer"
|
| 66 |
+
),
|
| 67 |
+
),
|
| 68 |
+
);
|
| 69 |
+
continue;
|
| 70 |
+
}
|
| 71 |
+
layer_roles.insert(role_name, role);
|
| 72 |
+
}
|
| 73 |
+
}
|
| 74 |
+
|
| 75 |
+
if let Some(config_folder) = layer.config_folder() {
|
| 76 |
+
for (role_name, role) in discover_agent_roles_in_dir(
|
| 77 |
+
fs,
|
| 78 |
+
&config_folder.join("agents"),
|
| 79 |
+
&declared_role_files,
|
| 80 |
+
startup_warnings,
|
| 81 |
+
)
|
| 82 |
+
.await?
|
| 83 |
+
{
|
| 84 |
+
if layer_roles.contains_key(&role_name) {
|
| 85 |
+
push_agent_role_warning(
|
| 86 |
+
startup_warnings,
|
| 87 |
+
std::io::Error::new(
|
| 88 |
+
std::io::ErrorKind::InvalidInput,
|
| 89 |
+
format!(
|
| 90 |
+
"duplicate agent role name `{role_name}` declared in the same config layer"
|
| 91 |
+
),
|
| 92 |
+
),
|
| 93 |
+
);
|
| 94 |
+
continue;
|
| 95 |
+
}
|
| 96 |
+
layer_roles.insert(role_name, role);
|
| 97 |
+
}
|
| 98 |
+
}
|
| 99 |
+
|
| 100 |
+
for (role_name, role) in layer_roles {
|
| 101 |
+
let mut merged_role = role;
|
| 102 |
+
if let Some(existing_role) = roles.get(&role_name) {
|
| 103 |
+
merge_missing_role_fields(&mut merged_role, existing_role);
|
| 104 |
+
}
|
| 105 |
+
if let Err(err) = validate_required_agent_role_description(
|
| 106 |
+
&role_name,
|
| 107 |
+
merged_role.description.as_deref(),
|
| 108 |
+
) {
|
| 109 |
+
push_agent_role_warning(startup_warnings, err);
|
| 110 |
+
continue;
|
| 111 |
+
}
|
| 112 |
+
roles.insert(role_name, merged_role);
|
| 113 |
+
}
|
| 114 |
+
}
|
| 115 |
+
|
| 116 |
+
Ok(roles)
|
| 117 |
+
}
|
| 118 |
+
|
| 119 |
+
fn push_agent_role_warning(startup_warnings: &mut Vec<String>, err: std::io::Error) {
|
| 120 |
+
let message = format!("Ignoring malformed agent role definition: {err}");
|
| 121 |
+
tracing::warn!("{message}");
|
| 122 |
+
startup_warnings.push(message);
|
| 123 |
+
}
|
| 124 |
+
|
| 125 |
+
async fn load_agent_roles_without_layers(
|
| 126 |
+
fs: &dyn ExecutorFileSystem,
|
| 127 |
+
cfg: &ConfigToml,
|
| 128 |
+
) -> std::io::Result<BTreeMap<String, AgentRoleConfig>> {
|
| 129 |
+
let mut roles = BTreeMap::new();
|
| 130 |
+
if let Some(agents_toml) = cfg.agents.as_ref() {
|
| 131 |
+
for (declared_role_name, role_toml) in &agents_toml.roles {
|
| 132 |
+
let (role_name, role) = read_declared_role(fs, declared_role_name, role_toml).await?;
|
| 133 |
+
validate_required_agent_role_description(&role_name, role.description.as_deref())?;
|
| 134 |
+
|
| 135 |
+
if roles.insert(role_name.clone(), role).is_some() {
|
| 136 |
+
return Err(std::io::Error::new(
|
| 137 |
+
std::io::ErrorKind::InvalidInput,
|
| 138 |
+
format!("duplicate agent role name `{role_name}` declared in config"),
|
| 139 |
+
));
|
| 140 |
+
}
|
| 141 |
+
}
|
| 142 |
+
}
|
| 143 |
+
|
| 144 |
+
Ok(roles)
|
| 145 |
+
}
|
| 146 |
+
|
| 147 |
+
async fn read_declared_role(
|
| 148 |
+
fs: &dyn ExecutorFileSystem,
|
| 149 |
+
declared_role_name: &str,
|
| 150 |
+
role_toml: &AgentRoleToml,
|
| 151 |
+
) -> std::io::Result<(String, AgentRoleConfig)> {
|
| 152 |
+
let mut role = agent_role_config_from_toml(fs, declared_role_name, role_toml).await?;
|
| 153 |
+
let mut role_name = declared_role_name.to_string();
|
| 154 |
+
if let Some(config_file) = role.config_file.as_deref() {
|
| 155 |
+
let config_file = AbsolutePathBuf::from_absolute_path(config_file)?;
|
| 156 |
+
let parsed_file =
|
| 157 |
+
read_resolved_agent_role_file(fs, &config_file, Some(declared_role_name)).await?;
|
| 158 |
+
role_name = parsed_file.role_name;
|
| 159 |
+
role.description = parsed_file.description.or(role.description);
|
| 160 |
+
role.nickname_candidates = parsed_file.nickname_candidates.or(role.nickname_candidates);
|
| 161 |
+
}
|
| 162 |
+
|
| 163 |
+
Ok((role_name, role))
|
| 164 |
+
}
|
| 165 |
+
|
| 166 |
+
fn merge_missing_role_fields(role: &mut AgentRoleConfig, fallback: &AgentRoleConfig) {
|
| 167 |
+
role.description = role.description.clone().or(fallback.description.clone());
|
| 168 |
+
role.config_file = role.config_file.clone().or(fallback.config_file.clone());
|
| 169 |
+
role.nickname_candidates = role
|
| 170 |
+
.nickname_candidates
|
| 171 |
+
.clone()
|
| 172 |
+
.or(fallback.nickname_candidates.clone());
|
| 173 |
+
}
|
| 174 |
+
|
| 175 |
+
fn agents_toml_from_layer(
|
| 176 |
+
layer_toml: &TomlValue,
|
| 177 |
+
config_base_dir: Option<&Path>,
|
| 178 |
+
) -> std::io::Result<Option<AgentsToml>> {
|
| 179 |
+
let Some(agents_toml) = layer_toml.get("agents") else {
|
| 180 |
+
return Ok(None);
|
| 181 |
+
};
|
| 182 |
+
|
| 183 |
+
// AbsolutePathBufGuard resolves relative paths while it remains in scope.
|
| 184 |
+
let _guard = config_base_dir.map(AbsolutePathBufGuard::new);
|
| 185 |
+
agents_toml
|
| 186 |
+
.clone()
|
| 187 |
+
.try_into()
|
| 188 |
+
.map(Some)
|
| 189 |
+
.map_err(|err| std::io::Error::new(std::io::ErrorKind::InvalidData, err))
|
| 190 |
+
}
|
| 191 |
+
|
| 192 |
+
async fn agent_role_config_from_toml(
|
| 193 |
+
fs: &dyn ExecutorFileSystem,
|
| 194 |
+
role_name: &str,
|
| 195 |
+
role: &AgentRoleToml,
|
| 196 |
+
) -> std::io::Result<AgentRoleConfig> {
|
| 197 |
+
let config_file = role
|
| 198 |
+
.config_file
|
| 199 |
+
.as_ref()
|
| 200 |
+
.map(AbsolutePathBuf::from_absolute_path)
|
| 201 |
+
.transpose()?;
|
| 202 |
+
validate_agent_role_config_file(fs, role_name, config_file.as_ref()).await?;
|
| 203 |
+
let description = normalize_agent_role_description(
|
| 204 |
+
&format!("agents.{role_name}.description"),
|
| 205 |
+
role.description.as_deref(),
|
| 206 |
+
)?;
|
| 207 |
+
let nickname_candidates = normalize_agent_role_nickname_candidates(
|
| 208 |
+
&format!("agents.{role_name}.nickname_candidates"),
|
| 209 |
+
role.nickname_candidates.as_deref(),
|
| 210 |
+
)?;
|
| 211 |
+
|
| 212 |
+
Ok(AgentRoleConfig {
|
| 213 |
+
description,
|
| 214 |
+
config_file: config_file.map(AbsolutePathBuf::into_path_buf),
|
| 215 |
+
nickname_candidates,
|
| 216 |
+
})
|
| 217 |
+
}
|
| 218 |
+
|
| 219 |
+
async fn read_resolved_agent_role_file(
|
| 220 |
+
fs: &dyn ExecutorFileSystem,
|
| 221 |
+
path: &AbsolutePathBuf,
|
| 222 |
+
role_name_hint: Option<&str>,
|
| 223 |
+
) -> std::io::Result<ResolvedAgentRoleFile> {
|
| 224 |
+
let path_uri = PathUri::from_abs_path(path);
|
| 225 |
+
let contents = fs
|
| 226 |
+
.read_file_text(&path_uri, ReadFileOptions::default(), /*sandbox*/ None)
|
| 227 |
+
.await?;
|
| 228 |
+
let config_base_dir = path.parent().unwrap_or_else(|| path.clone());
|
| 229 |
+
parse_agent_role_file_contents(
|
| 230 |
+
&contents,
|
| 231 |
+
path.as_path(),
|
| 232 |
+
config_base_dir.as_path(),
|
| 233 |
+
role_name_hint,
|
| 234 |
+
)
|
| 235 |
+
}
|
| 236 |
+
|
| 237 |
+
fn validate_required_agent_role_description(
|
| 238 |
+
role_name: &str,
|
| 239 |
+
description: Option<&str>,
|
| 240 |
+
) -> std::io::Result<()> {
|
| 241 |
+
if description.is_some() {
|
| 242 |
+
Ok(())
|
| 243 |
+
} else {
|
| 244 |
+
Err(std::io::Error::new(
|
| 245 |
+
std::io::ErrorKind::InvalidInput,
|
| 246 |
+
format!("agent role `{role_name}` must define a description"),
|
| 247 |
+
))
|
| 248 |
+
}
|
| 249 |
+
}
|
| 250 |
+
|
| 251 |
+
async fn validate_agent_role_config_file(
|
| 252 |
+
fs: &dyn ExecutorFileSystem,
|
| 253 |
+
role_name: &str,
|
| 254 |
+
config_file: Option<&AbsolutePathBuf>,
|
| 255 |
+
) -> std::io::Result<()> {
|
| 256 |
+
let Some(config_file) = config_file else {
|
| 257 |
+
return Ok(());
|
| 258 |
+
};
|
| 259 |
+
|
| 260 |
+
let config_file_uri = PathUri::from_abs_path(config_file);
|
| 261 |
+
let metadata = fs
|
| 262 |
+
.get_metadata(
|
| 263 |
+
&config_file_uri,
|
| 264 |
+
GetMetadataOptions::default(),
|
| 265 |
+
/*sandbox*/ None,
|
| 266 |
+
)
|
| 267 |
+
.await
|
| 268 |
+
.map_err(|e| {
|
| 269 |
+
std::io::Error::new(
|
| 270 |
+
std::io::ErrorKind::InvalidInput,
|
| 271 |
+
format!(
|
| 272 |
+
"agents.{role_name}.config_file must point to an existing file at {}: {e}",
|
| 273 |
+
config_file.as_path().display()
|
| 274 |
+
),
|
| 275 |
+
)
|
| 276 |
+
})?;
|
| 277 |
+
if metadata.is_file {
|
| 278 |
+
Ok(())
|
| 279 |
+
} else {
|
| 280 |
+
Err(std::io::Error::new(
|
| 281 |
+
std::io::ErrorKind::InvalidInput,
|
| 282 |
+
format!(
|
| 283 |
+
"agents.{role_name}.config_file must point to a file: {}",
|
| 284 |
+
config_file.as_path().display()
|
| 285 |
+
),
|
| 286 |
+
))
|
| 287 |
+
}
|
| 288 |
+
}
|
| 289 |
+
|
| 290 |
+
async fn discover_agent_roles_in_dir(
|
| 291 |
+
fs: &dyn ExecutorFileSystem,
|
| 292 |
+
agents_dir: &AbsolutePathBuf,
|
| 293 |
+
declared_role_files: &BTreeSet<PathBuf>,
|
| 294 |
+
startup_warnings: &mut Vec<String>,
|
| 295 |
+
) -> std::io::Result<BTreeMap<String, AgentRoleConfig>> {
|
| 296 |
+
let mut roles = BTreeMap::new();
|
| 297 |
+
|
| 298 |
+
for agent_file in collect_agent_role_files(fs, agents_dir).await? {
|
| 299 |
+
if declared_role_files.contains(agent_file.as_path()) {
|
| 300 |
+
continue;
|
| 301 |
+
}
|
| 302 |
+
let parsed_file =
|
| 303 |
+
match read_resolved_agent_role_file(fs, &agent_file, /*role_name_hint*/ None).await {
|
| 304 |
+
Ok(parsed_file) => parsed_file,
|
| 305 |
+
Err(err) => {
|
| 306 |
+
push_agent_role_warning(startup_warnings, err);
|
| 307 |
+
continue;
|
| 308 |
+
}
|
| 309 |
+
};
|
| 310 |
+
let role_name = parsed_file.role_name;
|
| 311 |
+
if roles.contains_key(&role_name) {
|
| 312 |
+
push_agent_role_warning(
|
| 313 |
+
startup_warnings,
|
| 314 |
+
std::io::Error::new(
|
| 315 |
+
std::io::ErrorKind::InvalidInput,
|
| 316 |
+
format!(
|
| 317 |
+
"duplicate agent role name `{role_name}` discovered in {}",
|
| 318 |
+
agents_dir.as_path().display()
|
| 319 |
+
),
|
| 320 |
+
),
|
| 321 |
+
);
|
| 322 |
+
continue;
|
| 323 |
+
}
|
| 324 |
+
roles.insert(
|
| 325 |
+
role_name,
|
| 326 |
+
AgentRoleConfig {
|
| 327 |
+
description: parsed_file.description,
|
| 328 |
+
config_file: Some(agent_file.to_path_buf()),
|
| 329 |
+
nickname_candidates: parsed_file.nickname_candidates,
|
| 330 |
+
},
|
| 331 |
+
);
|
| 332 |
+
}
|
| 333 |
+
|
| 334 |
+
Ok(roles)
|
| 335 |
+
}
|
codex-rs/app-server-protocol-noop-macros/src/lib.rs
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! No-op schema derives for production app-server protocol builds.
|
| 2 |
+
//!
|
| 3 |
+
//! The real `ts-rs` and `schemars` derives are only needed when regenerating
|
| 4 |
+
//! the vendored protocol exports. Normal builds retain the annotations so the
|
| 5 |
+
//! protocol definitions stay readable, but use these derives to avoid
|
| 6 |
+
//! generating implementations that cannot be reached at runtime.
|
| 7 |
+
|
| 8 |
+
use proc_macro::TokenStream;
|
| 9 |
+
|
| 10 |
+
/// Accepts `#[schemars(...)]` helper attributes without generating an impl.
|
| 11 |
+
#[proc_macro_derive(JsonSchema, attributes(schemars))]
|
| 12 |
+
pub fn derive_json_schema(_input: TokenStream) -> TokenStream {
|
| 13 |
+
TokenStream::new()
|
| 14 |
+
}
|
| 15 |
+
|
| 16 |
+
/// Accepts `#[ts(...)]` helper attributes without generating an impl.
|
| 17 |
+
#[proc_macro_derive(TS, attributes(ts))]
|
| 18 |
+
pub fn derive_ts(_input: TokenStream) -> TokenStream {
|
| 19 |
+
TokenStream::new()
|
| 20 |
+
}
|
codex-rs/cloud-tasks-mock-client/src/lib.rs
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
mod mock;
|
| 2 |
+
|
| 3 |
+
pub use mock::MockClient;
|
codex-rs/cloud-tasks-mock-client/src/mock.rs
ADDED
|
@@ -0,0 +1,267 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use chrono::Utc;
|
| 2 |
+
use codex_cloud_tasks_client::ApplyOutcome;
|
| 3 |
+
use codex_cloud_tasks_client::ApplyStatus;
|
| 4 |
+
use codex_cloud_tasks_client::AttemptStatus;
|
| 5 |
+
use codex_cloud_tasks_client::CloudBackend;
|
| 6 |
+
use codex_cloud_tasks_client::CloudBackendFuture;
|
| 7 |
+
use codex_cloud_tasks_client::CloudTaskError;
|
| 8 |
+
use codex_cloud_tasks_client::CreatedTask;
|
| 9 |
+
use codex_cloud_tasks_client::DiffSummary;
|
| 10 |
+
use codex_cloud_tasks_client::Result;
|
| 11 |
+
use codex_cloud_tasks_client::TaskId;
|
| 12 |
+
use codex_cloud_tasks_client::TaskListPage;
|
| 13 |
+
use codex_cloud_tasks_client::TaskStatus;
|
| 14 |
+
use codex_cloud_tasks_client::TaskSummary;
|
| 15 |
+
use codex_cloud_tasks_client::TaskText;
|
| 16 |
+
use codex_cloud_tasks_client::TurnAttempt;
|
| 17 |
+
|
| 18 |
+
#[derive(Clone, Default)]
|
| 19 |
+
pub struct MockClient;
|
| 20 |
+
|
| 21 |
+
impl MockClient {
|
| 22 |
+
async fn list_tasks(
|
| 23 |
+
&self,
|
| 24 |
+
_env: Option<&str>,
|
| 25 |
+
_limit: Option<i64>,
|
| 26 |
+
_cursor: Option<&str>,
|
| 27 |
+
) -> Result<TaskListPage> {
|
| 28 |
+
// Slightly vary content by env to aid tests that rely on the mock
|
| 29 |
+
let rows = match _env {
|
| 30 |
+
Some("env-A") => vec![("T-2000", "A: First", TaskStatus::Ready)],
|
| 31 |
+
Some("env-B") => vec![
|
| 32 |
+
("T-3000", "B: One", TaskStatus::Ready),
|
| 33 |
+
("T-3001", "B: Two", TaskStatus::Pending),
|
| 34 |
+
],
|
| 35 |
+
_ => vec![
|
| 36 |
+
("T-1000", "Update README formatting", TaskStatus::Ready),
|
| 37 |
+
("T-1001", "Fix clippy warnings in core", TaskStatus::Pending),
|
| 38 |
+
("T-1002", "Add contributing guide", TaskStatus::Ready),
|
| 39 |
+
],
|
| 40 |
+
};
|
| 41 |
+
let environment_id = _env.map(str::to_string);
|
| 42 |
+
let environment_label = match _env {
|
| 43 |
+
Some("env-A") => Some("Env A".to_string()),
|
| 44 |
+
Some("env-B") => Some("Env B".to_string()),
|
| 45 |
+
Some(other) => Some(other.to_string()),
|
| 46 |
+
None => Some("Global".to_string()),
|
| 47 |
+
};
|
| 48 |
+
let mut out = Vec::new();
|
| 49 |
+
for (id_str, title, status) in rows {
|
| 50 |
+
let id = TaskId(id_str.to_string());
|
| 51 |
+
let diff = mock_diff_for(&id);
|
| 52 |
+
let (a, d) = count_from_unified(&diff);
|
| 53 |
+
out.push(TaskSummary {
|
| 54 |
+
id,
|
| 55 |
+
title: title.to_string(),
|
| 56 |
+
status,
|
| 57 |
+
updated_at: Utc::now(),
|
| 58 |
+
environment_id: environment_id.clone(),
|
| 59 |
+
environment_label: environment_label.clone(),
|
| 60 |
+
summary: DiffSummary {
|
| 61 |
+
files_changed: 1,
|
| 62 |
+
lines_added: a,
|
| 63 |
+
lines_removed: d,
|
| 64 |
+
},
|
| 65 |
+
is_review: false,
|
| 66 |
+
attempt_total: Some(if id_str == "T-1000" { 2 } else { 1 }),
|
| 67 |
+
});
|
| 68 |
+
}
|
| 69 |
+
Ok(TaskListPage {
|
| 70 |
+
tasks: out,
|
| 71 |
+
cursor: None,
|
| 72 |
+
})
|
| 73 |
+
}
|
| 74 |
+
|
| 75 |
+
async fn get_task_summary(&self, id: TaskId) -> Result<TaskSummary> {
|
| 76 |
+
let tasks = self
|
| 77 |
+
.list_tasks(/*env*/ None, /*limit*/ None, /*cursor*/ None)
|
| 78 |
+
.await?
|
| 79 |
+
.tasks;
|
| 80 |
+
tasks
|
| 81 |
+
.into_iter()
|
| 82 |
+
.find(|t| t.id == id)
|
| 83 |
+
.ok_or_else(|| CloudTaskError::Msg(format!("Task {} not found (mock)", id.0)))
|
| 84 |
+
}
|
| 85 |
+
|
| 86 |
+
async fn get_task_diff(&self, id: TaskId) -> Result<Option<String>> {
|
| 87 |
+
Ok(Some(mock_diff_for(&id)))
|
| 88 |
+
}
|
| 89 |
+
|
| 90 |
+
async fn get_task_messages(&self, _id: TaskId) -> Result<Vec<String>> {
|
| 91 |
+
Ok(vec![
|
| 92 |
+
"Mock assistant output: this task contains no diff.".to_string(),
|
| 93 |
+
])
|
| 94 |
+
}
|
| 95 |
+
|
| 96 |
+
async fn get_task_text(&self, _id: TaskId) -> Result<TaskText> {
|
| 97 |
+
Ok(TaskText {
|
| 98 |
+
prompt: Some("Why is there no diff?".to_string()),
|
| 99 |
+
messages: vec!["Mock assistant output: this task contains no diff.".to_string()],
|
| 100 |
+
turn_id: Some("mock-turn".to_string()),
|
| 101 |
+
sibling_turn_ids: Vec::new(),
|
| 102 |
+
attempt_placement: Some(0),
|
| 103 |
+
attempt_status: AttemptStatus::Completed,
|
| 104 |
+
})
|
| 105 |
+
}
|
| 106 |
+
|
| 107 |
+
async fn apply_task(&self, id: TaskId, _diff_override: Option<String>) -> Result<ApplyOutcome> {
|
| 108 |
+
Ok(ApplyOutcome {
|
| 109 |
+
applied: true,
|
| 110 |
+
status: ApplyStatus::Success,
|
| 111 |
+
message: format!("Applied task {} locally (mock)", id.0),
|
| 112 |
+
skipped_paths: Vec::new(),
|
| 113 |
+
conflict_paths: Vec::new(),
|
| 114 |
+
})
|
| 115 |
+
}
|
| 116 |
+
|
| 117 |
+
async fn apply_task_preflight(
|
| 118 |
+
&self,
|
| 119 |
+
id: TaskId,
|
| 120 |
+
_diff_override: Option<String>,
|
| 121 |
+
) -> Result<ApplyOutcome> {
|
| 122 |
+
Ok(ApplyOutcome {
|
| 123 |
+
applied: false,
|
| 124 |
+
status: ApplyStatus::Success,
|
| 125 |
+
message: format!("Preflight passed for task {} (mock)", id.0),
|
| 126 |
+
skipped_paths: Vec::new(),
|
| 127 |
+
conflict_paths: Vec::new(),
|
| 128 |
+
})
|
| 129 |
+
}
|
| 130 |
+
|
| 131 |
+
async fn list_sibling_attempts(
|
| 132 |
+
&self,
|
| 133 |
+
task: TaskId,
|
| 134 |
+
_turn_id: String,
|
| 135 |
+
) -> Result<Vec<TurnAttempt>> {
|
| 136 |
+
if task.0 == "T-1000" {
|
| 137 |
+
return Ok(vec![TurnAttempt {
|
| 138 |
+
turn_id: "T-1000-attempt-2".to_string(),
|
| 139 |
+
attempt_placement: Some(1),
|
| 140 |
+
created_at: Some(Utc::now()),
|
| 141 |
+
status: AttemptStatus::Completed,
|
| 142 |
+
diff: Some(mock_diff_for(&task)),
|
| 143 |
+
messages: vec!["Mock alternate attempt".to_string()],
|
| 144 |
+
}]);
|
| 145 |
+
}
|
| 146 |
+
Ok(Vec::new())
|
| 147 |
+
}
|
| 148 |
+
|
| 149 |
+
async fn create_task(
|
| 150 |
+
&self,
|
| 151 |
+
env_id: &str,
|
| 152 |
+
prompt: &str,
|
| 153 |
+
git_ref: &str,
|
| 154 |
+
qa_mode: bool,
|
| 155 |
+
best_of_n: usize,
|
| 156 |
+
) -> Result<CreatedTask> {
|
| 157 |
+
let _ = (env_id, prompt, git_ref, qa_mode, best_of_n);
|
| 158 |
+
let id = format!("task_local_{}", chrono::Utc::now().timestamp_millis());
|
| 159 |
+
Ok(CreatedTask { id: TaskId(id) })
|
| 160 |
+
}
|
| 161 |
+
}
|
| 162 |
+
|
| 163 |
+
impl CloudBackend for MockClient {
|
| 164 |
+
fn list_tasks<'a>(
|
| 165 |
+
&'a self,
|
| 166 |
+
env: Option<&'a str>,
|
| 167 |
+
limit: Option<i64>,
|
| 168 |
+
cursor: Option<&'a str>,
|
| 169 |
+
) -> CloudBackendFuture<'a, TaskListPage> {
|
| 170 |
+
Box::pin(MockClient::list_tasks(self, env, limit, cursor))
|
| 171 |
+
}
|
| 172 |
+
|
| 173 |
+
fn get_task_summary(&self, id: TaskId) -> CloudBackendFuture<'_, TaskSummary> {
|
| 174 |
+
Box::pin(MockClient::get_task_summary(self, id))
|
| 175 |
+
}
|
| 176 |
+
|
| 177 |
+
fn get_task_diff(&self, id: TaskId) -> CloudBackendFuture<'_, Option<String>> {
|
| 178 |
+
Box::pin(MockClient::get_task_diff(self, id))
|
| 179 |
+
}
|
| 180 |
+
|
| 181 |
+
fn get_task_messages(&self, id: TaskId) -> CloudBackendFuture<'_, Vec<String>> {
|
| 182 |
+
Box::pin(MockClient::get_task_messages(self, id))
|
| 183 |
+
}
|
| 184 |
+
|
| 185 |
+
fn get_task_text(&self, id: TaskId) -> CloudBackendFuture<'_, TaskText> {
|
| 186 |
+
Box::pin(MockClient::get_task_text(self, id))
|
| 187 |
+
}
|
| 188 |
+
|
| 189 |
+
fn apply_task(
|
| 190 |
+
&self,
|
| 191 |
+
id: TaskId,
|
| 192 |
+
diff_override: Option<String>,
|
| 193 |
+
) -> CloudBackendFuture<'_, ApplyOutcome> {
|
| 194 |
+
Box::pin(MockClient::apply_task(self, id, diff_override))
|
| 195 |
+
}
|
| 196 |
+
|
| 197 |
+
fn apply_task_preflight(
|
| 198 |
+
&self,
|
| 199 |
+
id: TaskId,
|
| 200 |
+
diff_override: Option<String>,
|
| 201 |
+
) -> CloudBackendFuture<'_, ApplyOutcome> {
|
| 202 |
+
Box::pin(MockClient::apply_task_preflight(self, id, diff_override))
|
| 203 |
+
}
|
| 204 |
+
|
| 205 |
+
fn list_sibling_attempts(
|
| 206 |
+
&self,
|
| 207 |
+
task: TaskId,
|
| 208 |
+
turn_id: String,
|
| 209 |
+
) -> CloudBackendFuture<'_, Vec<TurnAttempt>> {
|
| 210 |
+
Box::pin(MockClient::list_sibling_attempts(self, task, turn_id))
|
| 211 |
+
}
|
| 212 |
+
|
| 213 |
+
fn create_task<'a>(
|
| 214 |
+
&'a self,
|
| 215 |
+
env_id: &'a str,
|
| 216 |
+
prompt: &'a str,
|
| 217 |
+
git_ref: &'a str,
|
| 218 |
+
qa_mode: bool,
|
| 219 |
+
best_of_n: usize,
|
| 220 |
+
) -> CloudBackendFuture<'a, CreatedTask> {
|
| 221 |
+
Box::pin(MockClient::create_task(
|
| 222 |
+
self, env_id, prompt, git_ref, qa_mode, best_of_n,
|
| 223 |
+
))
|
| 224 |
+
}
|
| 225 |
+
}
|
| 226 |
+
|
| 227 |
+
fn mock_diff_for(id: &TaskId) -> String {
|
| 228 |
+
match id.0.as_str() {
|
| 229 |
+
"T-1000" => {
|
| 230 |
+
"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()
|
| 231 |
+
}
|
| 232 |
+
"T-1001" => {
|
| 233 |
+
"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()
|
| 234 |
+
}
|
| 235 |
+
_ => {
|
| 236 |
+
"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()
|
| 237 |
+
}
|
| 238 |
+
}
|
| 239 |
+
}
|
| 240 |
+
|
| 241 |
+
fn count_from_unified(diff: &str) -> (usize, usize) {
|
| 242 |
+
if let Ok(patch) = diffy::Patch::from_str(diff) {
|
| 243 |
+
patch
|
| 244 |
+
.hunks()
|
| 245 |
+
.iter()
|
| 246 |
+
.flat_map(diffy::Hunk::lines)
|
| 247 |
+
.fold((0, 0), |(a, d), l| match l {
|
| 248 |
+
diffy::Line::Insert(_) => (a + 1, d),
|
| 249 |
+
diffy::Line::Delete(_) => (a, d + 1),
|
| 250 |
+
_ => (a, d),
|
| 251 |
+
})
|
| 252 |
+
} else {
|
| 253 |
+
let mut a = 0;
|
| 254 |
+
let mut d = 0;
|
| 255 |
+
for l in diff.lines() {
|
| 256 |
+
if l.starts_with("+++") || l.starts_with("---") || l.starts_with("@@") {
|
| 257 |
+
continue;
|
| 258 |
+
}
|
| 259 |
+
match l.as_bytes().first() {
|
| 260 |
+
Some(b'+') => a += 1,
|
| 261 |
+
Some(b'-') => d += 1,
|
| 262 |
+
_ => {}
|
| 263 |
+
}
|
| 264 |
+
}
|
| 265 |
+
(a, d)
|
| 266 |
+
}
|
| 267 |
+
}
|
codex-rs/code-mode-protocol/src/description.rs
ADDED
|
@@ -0,0 +1,993 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use codex_protocol::ToolName;
|
| 2 |
+
use serde::Deserialize;
|
| 3 |
+
use serde::Serialize;
|
| 4 |
+
use serde_json::Value as JsonValue;
|
| 5 |
+
use std::collections::BTreeMap;
|
| 6 |
+
|
| 7 |
+
use crate::PUBLIC_TOOL_NAME;
|
| 8 |
+
use crate::json_schema_types::render_json_schema_to_typescript;
|
| 9 |
+
|
| 10 |
+
const MAX_JS_SAFE_INTEGER: u64 = (1_u64 << 53) - 1;
|
| 11 |
+
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`.
|
| 12 |
+
To find one, filter `ALL_TOOLS` by `name` and `description`."#;
|
| 13 |
+
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."#;
|
| 14 |
+
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])`."#;
|
| 15 |
+
const EXEC_DESCRIPTION_TEMPLATE: &str = r#"Run JavaScript code to orchestrate/compose tool calls
|
| 16 |
+
- Evaluates the provided JavaScript code in a fresh V8 isolate as an async module.
|
| 17 |
+
- 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(...)`.
|
| 18 |
+
- Nested tool methods take either a string or an object as their input argument.
|
| 19 |
+
- Nested tools return either an object or a string, based on the description.
|
| 20 |
+
- Runs raw JavaScript -- no Node, no file system, no network access, no console.
|
| 21 |
+
- Accepts raw JavaScript source text, not JSON, quoted strings, or markdown code fences.
|
| 22 |
+
- You may optionally start the tool input with a first-line pragma like `// @exec: {"yield_time_ms": 10000, "max_output_tokens": 1000}`.
|
| 23 |
+
- `yield_time_ms` asks `exec` to yield early if the script is still running. Defaults to 10000 ms.
|
| 24 |
+
- `max_output_tokens` sets the token budget for direct `exec` results. Defaults to 10000 tokens.
|
| 25 |
+
- When the JS code is fully evaluated, the isolate's lifetime ends and unawaited promises are silently discarded.
|
| 26 |
+
|
| 27 |
+
- Global helpers:
|
| 28 |
+
- `exit()`: Immediately ends the current script successfully (like an early return from the top level).
|
| 29 |
+
- `text(value: string | number | boolean | undefined | null)`: Appends a text item. Non-string values are stringified with `JSON.stringify(...)` when possible.
|
| 30 |
+
- `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.
|
| 31 |
+
- `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])`.
|
| 32 |
+
- `generatedImage(result: { image_url: string; output_hint?: string })`: Appends an image-generation result and its optional output hint. HTTP(S) URLs are not supported.
|
| 33 |
+
- `store(key: string, value: any)`: stores a serializable value under a string key for later `exec` calls in the same session.
|
| 34 |
+
- `load(key: string)`: returns the stored value for a string key, or `undefined` if it is missing.
|
| 35 |
+
- `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(...)`.
|
| 36 |
+
- `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.
|
| 37 |
+
- `clearTimeout(timeoutId?: number)`: cancels a timeout created by `setTimeout`.
|
| 38 |
+
- `ALL_TOOLS`: metadata for the enabled nested tools as `{ name, description }` entries.
|
| 39 |
+
- `yield_control()`: yields the accumulated output to the model immediately while the script keeps running."#;
|
| 40 |
+
const WAIT_DESCRIPTION_TEMPLATE: &str = r#"- Use `wait` only after `exec` returns `Script running with cell ID ...`.
|
| 41 |
+
- `cell_id` identifies the running `exec` cell to resume.
|
| 42 |
+
- `yield_time_ms` controls how long to wait for more output before yielding again. Defaults to 10000 ms.
|
| 43 |
+
- `max_tokens` limits how much new output this wait call returns. Defaults to 10000 tokens.
|
| 44 |
+
- `terminate: true` stops the running cell; false or omitted waits for output.
|
| 45 |
+
- `wait` returns only the new output since the last yield, or the final completion or termination result for that cell.
|
| 46 |
+
- If the cell is still running, `wait` may yield again with the same `cell_id`.
|
| 47 |
+
- If the cell has already finished, `wait` returns the completed result and closes the cell."#;
|
| 48 |
+
// Based off of https://modelcontextprotocol.io/specification/draft/schema#calltoolresult
|
| 49 |
+
const MCP_TYPESCRIPT_PREAMBLE: &str = r#"type Role = "user" | "assistant";
|
| 50 |
+
type MetaObject = Record<string, unknown>;
|
| 51 |
+
type Annotations = {
|
| 52 |
+
audience?: Role[];
|
| 53 |
+
priority?: number;
|
| 54 |
+
lastModified?: string;
|
| 55 |
+
};
|
| 56 |
+
type Icon = {
|
| 57 |
+
src: string;
|
| 58 |
+
mimeType?: string;
|
| 59 |
+
sizes?: string[];
|
| 60 |
+
theme?: "light" | "dark";
|
| 61 |
+
};
|
| 62 |
+
type TextResourceContents = {
|
| 63 |
+
uri: string;
|
| 64 |
+
mimeType?: string;
|
| 65 |
+
_meta?: MetaObject;
|
| 66 |
+
text: string;
|
| 67 |
+
};
|
| 68 |
+
type BlobResourceContents = {
|
| 69 |
+
uri: string;
|
| 70 |
+
mimeType?: string;
|
| 71 |
+
_meta?: MetaObject;
|
| 72 |
+
blob: string;
|
| 73 |
+
};
|
| 74 |
+
type TextContent = {
|
| 75 |
+
type: "text";
|
| 76 |
+
text: string;
|
| 77 |
+
annotations?: Annotations;
|
| 78 |
+
_meta?: MetaObject;
|
| 79 |
+
};
|
| 80 |
+
type ImageContent = {
|
| 81 |
+
type: "image";
|
| 82 |
+
data: string;
|
| 83 |
+
mimeType: string;
|
| 84 |
+
annotations?: Annotations;
|
| 85 |
+
_meta?: MetaObject;
|
| 86 |
+
};
|
| 87 |
+
type AudioContent = {
|
| 88 |
+
type: "audio";
|
| 89 |
+
data: string;
|
| 90 |
+
mimeType: string;
|
| 91 |
+
annotations?: Annotations;
|
| 92 |
+
_meta?: MetaObject;
|
| 93 |
+
};
|
| 94 |
+
type ResourceLink = {
|
| 95 |
+
icons?: Icon[];
|
| 96 |
+
name: string;
|
| 97 |
+
title?: string;
|
| 98 |
+
uri: string;
|
| 99 |
+
description?: string;
|
| 100 |
+
mimeType?: string;
|
| 101 |
+
annotations?: Annotations;
|
| 102 |
+
size?: number;
|
| 103 |
+
_meta?: MetaObject;
|
| 104 |
+
type: "resource_link";
|
| 105 |
+
};
|
| 106 |
+
type EmbeddedResource = {
|
| 107 |
+
type: "resource";
|
| 108 |
+
resource: TextResourceContents | BlobResourceContents;
|
| 109 |
+
annotations?: Annotations;
|
| 110 |
+
_meta?: MetaObject;
|
| 111 |
+
};
|
| 112 |
+
type ContentBlock =
|
| 113 |
+
| TextContent
|
| 114 |
+
| ImageContent
|
| 115 |
+
| AudioContent
|
| 116 |
+
| ResourceLink
|
| 117 |
+
| EmbeddedResource;
|
| 118 |
+
type CallToolResult<TStructured = { [key: string]: unknown }> = {
|
| 119 |
+
_meta?: MetaObject;
|
| 120 |
+
content: ContentBlock[];
|
| 121 |
+
isError?: boolean;
|
| 122 |
+
structuredContent?: TStructured;
|
| 123 |
+
[key: string]: unknown;
|
| 124 |
+
};"#;
|
| 125 |
+
|
| 126 |
+
pub const CODE_MODE_PRAGMA_PREFIX: &str = "// @exec:";
|
| 127 |
+
|
| 128 |
+
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
| 129 |
+
#[serde(rename_all = "snake_case")]
|
| 130 |
+
pub enum CodeModeToolKind {
|
| 131 |
+
Function,
|
| 132 |
+
Freeform,
|
| 133 |
+
}
|
| 134 |
+
|
| 135 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 136 |
+
pub struct ToolDefinition {
|
| 137 |
+
pub name: String,
|
| 138 |
+
pub tool_name: ToolName,
|
| 139 |
+
pub description: String,
|
| 140 |
+
pub kind: CodeModeToolKind,
|
| 141 |
+
pub input_schema: Option<JsonValue>,
|
| 142 |
+
pub output_schema: Option<JsonValue>,
|
| 143 |
+
}
|
| 144 |
+
|
| 145 |
+
#[derive(Clone, Debug, Eq, PartialEq)]
|
| 146 |
+
pub struct ToolNamespaceDescription {
|
| 147 |
+
pub name: String,
|
| 148 |
+
pub description: String,
|
| 149 |
+
}
|
| 150 |
+
|
| 151 |
+
#[derive(Debug, Default, Deserialize, PartialEq, Eq)]
|
| 152 |
+
#[serde(deny_unknown_fields)]
|
| 153 |
+
struct CodeModeExecPragma {
|
| 154 |
+
#[serde(default)]
|
| 155 |
+
yield_time_ms: Option<u64>,
|
| 156 |
+
#[serde(default)]
|
| 157 |
+
max_output_tokens: Option<usize>,
|
| 158 |
+
}
|
| 159 |
+
|
| 160 |
+
#[derive(Debug, PartialEq, Eq)]
|
| 161 |
+
pub struct ParsedExecSource {
|
| 162 |
+
pub code: String,
|
| 163 |
+
pub yield_time_ms: Option<u64>,
|
| 164 |
+
pub max_output_tokens: Option<usize>,
|
| 165 |
+
}
|
| 166 |
+
|
| 167 |
+
pub fn parse_exec_source(input: &str) -> Result<ParsedExecSource, String> {
|
| 168 |
+
if input.trim().is_empty() {
|
| 169 |
+
return Err(
|
| 170 |
+
"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(),
|
| 171 |
+
);
|
| 172 |
+
}
|
| 173 |
+
|
| 174 |
+
let mut args = ParsedExecSource {
|
| 175 |
+
code: input.to_string(),
|
| 176 |
+
yield_time_ms: None,
|
| 177 |
+
max_output_tokens: None,
|
| 178 |
+
};
|
| 179 |
+
|
| 180 |
+
let mut lines = input.splitn(2, '\n');
|
| 181 |
+
let first_line = lines.next().unwrap_or_default();
|
| 182 |
+
let rest = lines.next().unwrap_or_default();
|
| 183 |
+
let trimmed = first_line.trim_start();
|
| 184 |
+
let Some(pragma) = trimmed.strip_prefix(CODE_MODE_PRAGMA_PREFIX) else {
|
| 185 |
+
return Ok(args);
|
| 186 |
+
};
|
| 187 |
+
|
| 188 |
+
if rest.trim().is_empty() {
|
| 189 |
+
return Err(
|
| 190 |
+
"exec pragma must be followed by JavaScript source on subsequent lines".to_string(),
|
| 191 |
+
);
|
| 192 |
+
}
|
| 193 |
+
|
| 194 |
+
let directive = pragma.trim();
|
| 195 |
+
if directive.is_empty() {
|
| 196 |
+
return Err(
|
| 197 |
+
"exec pragma must be a JSON object with supported fields `yield_time_ms` and `max_output_tokens`"
|
| 198 |
+
.to_string(),
|
| 199 |
+
);
|
| 200 |
+
}
|
| 201 |
+
|
| 202 |
+
let value: serde_json::Value = serde_json::from_str(directive).map_err(|err| {
|
| 203 |
+
format!(
|
| 204 |
+
"exec pragma must be valid JSON with supported fields `yield_time_ms` and `max_output_tokens`: {err}"
|
| 205 |
+
)
|
| 206 |
+
})?;
|
| 207 |
+
let object = value.as_object().ok_or_else(|| {
|
| 208 |
+
"exec pragma must be a JSON object with supported fields `yield_time_ms` and `max_output_tokens`"
|
| 209 |
+
.to_string()
|
| 210 |
+
})?;
|
| 211 |
+
for key in object.keys() {
|
| 212 |
+
match key.as_str() {
|
| 213 |
+
"yield_time_ms" | "max_output_tokens" => {}
|
| 214 |
+
_ => {
|
| 215 |
+
return Err(format!(
|
| 216 |
+
"exec pragma only supports `yield_time_ms` and `max_output_tokens`; got `{key}`"
|
| 217 |
+
));
|
| 218 |
+
}
|
| 219 |
+
}
|
| 220 |
+
}
|
| 221 |
+
|
| 222 |
+
let pragma: CodeModeExecPragma = serde_json::from_value(value).map_err(|err| {
|
| 223 |
+
format!(
|
| 224 |
+
"exec pragma fields `yield_time_ms` and `max_output_tokens` must be non-negative safe integers: {err}"
|
| 225 |
+
)
|
| 226 |
+
})?;
|
| 227 |
+
if pragma
|
| 228 |
+
.yield_time_ms
|
| 229 |
+
.is_some_and(|yield_time_ms| yield_time_ms > MAX_JS_SAFE_INTEGER)
|
| 230 |
+
{
|
| 231 |
+
return Err(
|
| 232 |
+
"exec pragma field `yield_time_ms` must be a non-negative safe integer".to_string(),
|
| 233 |
+
);
|
| 234 |
+
}
|
| 235 |
+
if pragma.max_output_tokens.is_some_and(|max_output_tokens| {
|
| 236 |
+
u64::try_from(max_output_tokens)
|
| 237 |
+
.map(|max_output_tokens| max_output_tokens > MAX_JS_SAFE_INTEGER)
|
| 238 |
+
.unwrap_or(true)
|
| 239 |
+
}) {
|
| 240 |
+
return Err(
|
| 241 |
+
"exec pragma field `max_output_tokens` must be a non-negative safe integer".to_string(),
|
| 242 |
+
);
|
| 243 |
+
}
|
| 244 |
+
|
| 245 |
+
args.code = rest.to_string();
|
| 246 |
+
args.yield_time_ms = pragma.yield_time_ms;
|
| 247 |
+
args.max_output_tokens = pragma.max_output_tokens;
|
| 248 |
+
Ok(args)
|
| 249 |
+
}
|
| 250 |
+
|
| 251 |
+
pub fn is_code_mode_nested_tool(tool_name: &str) -> bool {
|
| 252 |
+
tool_name != crate::PUBLIC_TOOL_NAME && tool_name != crate::WAIT_TOOL_NAME
|
| 253 |
+
}
|
| 254 |
+
|
| 255 |
+
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
| 256 |
+
pub enum ImageDetailVisibility {
|
| 257 |
+
Visible,
|
| 258 |
+
Hidden,
|
| 259 |
+
}
|
| 260 |
+
|
| 261 |
+
pub fn build_exec_tool_description(
|
| 262 |
+
enabled_tools: &[ToolDefinition],
|
| 263 |
+
deferred_tools: &[ToolDefinition],
|
| 264 |
+
namespace_descriptions: &BTreeMap<String, ToolNamespaceDescription>,
|
| 265 |
+
default_exec_yield_time_ms: u64,
|
| 266 |
+
code_mode_only: bool,
|
| 267 |
+
image_detail_visibility: ImageDetailVisibility,
|
| 268 |
+
) -> String {
|
| 269 |
+
let mut sections = Vec::new();
|
| 270 |
+
sections.push(EXEC_DESCRIPTION_TEMPLATE.replace(
|
| 271 |
+
"Defaults to 10000 ms.",
|
| 272 |
+
&format!("Defaults to {default_exec_yield_time_ms} ms."),
|
| 273 |
+
));
|
| 274 |
+
if image_detail_visibility == ImageDetailVisibility::Hidden {
|
| 275 |
+
sections[0] = sections[0].replace(
|
| 276 |
+
LEGACY_IMAGE_HELPER_DESCRIPTION,
|
| 277 |
+
UNIFIED_IMAGE_HELPER_DESCRIPTION,
|
| 278 |
+
);
|
| 279 |
+
}
|
| 280 |
+
if !deferred_tools.is_empty() {
|
| 281 |
+
sections.push(DEFERRED_NESTED_TOOLS_GUIDANCE.to_string());
|
| 282 |
+
}
|
| 283 |
+
if !code_mode_only {
|
| 284 |
+
return sections.join("\n\n");
|
| 285 |
+
}
|
| 286 |
+
|
| 287 |
+
let has_mcp_tools = enabled_tools
|
| 288 |
+
.iter()
|
| 289 |
+
.chain(deferred_tools)
|
| 290 |
+
.any(|tool| mcp_structured_content_schema(tool.output_schema.as_ref()).is_some());
|
| 291 |
+
if has_mcp_tools {
|
| 292 |
+
sections.push(format!(
|
| 293 |
+
"Shared MCP Types:\n```ts\n{MCP_TYPESCRIPT_PREAMBLE}\n```"
|
| 294 |
+
));
|
| 295 |
+
}
|
| 296 |
+
|
| 297 |
+
if !enabled_tools.is_empty() {
|
| 298 |
+
let mut current_namespace: Option<&str> = None;
|
| 299 |
+
let mut nested_tool_sections = Vec::with_capacity(enabled_tools.len());
|
| 300 |
+
|
| 301 |
+
for tool in enabled_tools {
|
| 302 |
+
let name = tool.name.as_str();
|
| 303 |
+
let nested_description = render_code_mode_sample_for_definition(tool);
|
| 304 |
+
let namespace_description = tool
|
| 305 |
+
.tool_name
|
| 306 |
+
.namespace
|
| 307 |
+
.as_ref()
|
| 308 |
+
.and_then(|namespace| namespace_descriptions.get(namespace));
|
| 309 |
+
let next_namespace = namespace_description
|
| 310 |
+
.map(|namespace_description| namespace_description.name.as_str());
|
| 311 |
+
if next_namespace != current_namespace {
|
| 312 |
+
if let Some(namespace_description) = namespace_description {
|
| 313 |
+
let namespace_description_text = namespace_description.description.trim();
|
| 314 |
+
if !namespace_description_text.is_empty() {
|
| 315 |
+
nested_tool_sections.push(format!(
|
| 316 |
+
"## {}\n{namespace_description_text}",
|
| 317 |
+
namespace_description.name
|
| 318 |
+
));
|
| 319 |
+
}
|
| 320 |
+
}
|
| 321 |
+
current_namespace = next_namespace;
|
| 322 |
+
}
|
| 323 |
+
|
| 324 |
+
let global_name = normalize_code_mode_identifier(name);
|
| 325 |
+
let nested_description = nested_description.trim();
|
| 326 |
+
if nested_description.is_empty() {
|
| 327 |
+
nested_tool_sections.push(render_tool_heading(&global_name, name));
|
| 328 |
+
} else {
|
| 329 |
+
nested_tool_sections.push(format!(
|
| 330 |
+
"{}\n{nested_description}",
|
| 331 |
+
render_tool_heading(&global_name, name)
|
| 332 |
+
));
|
| 333 |
+
}
|
| 334 |
+
}
|
| 335 |
+
|
| 336 |
+
sections.push(nested_tool_sections.join("\n\n"));
|
| 337 |
+
}
|
| 338 |
+
|
| 339 |
+
sections.join("\n\n")
|
| 340 |
+
}
|
| 341 |
+
|
| 342 |
+
pub fn build_wait_tool_description() -> &'static str {
|
| 343 |
+
WAIT_DESCRIPTION_TEMPLATE
|
| 344 |
+
}
|
| 345 |
+
|
| 346 |
+
pub fn normalize_code_mode_identifier(tool_key: &str) -> String {
|
| 347 |
+
let mut identifier = String::new();
|
| 348 |
+
|
| 349 |
+
for (index, ch) in tool_key.chars().enumerate() {
|
| 350 |
+
let is_valid = if index == 0 {
|
| 351 |
+
ch == '_' || ch == '$' || ch.is_ascii_alphabetic()
|
| 352 |
+
} else {
|
| 353 |
+
ch == '_' || ch == '$' || ch.is_ascii_alphanumeric()
|
| 354 |
+
};
|
| 355 |
+
|
| 356 |
+
if is_valid {
|
| 357 |
+
identifier.push(ch);
|
| 358 |
+
} else {
|
| 359 |
+
identifier.push('_');
|
| 360 |
+
}
|
| 361 |
+
}
|
| 362 |
+
|
| 363 |
+
if identifier.is_empty() {
|
| 364 |
+
"_".to_string()
|
| 365 |
+
} else {
|
| 366 |
+
identifier
|
| 367 |
+
}
|
| 368 |
+
}
|
| 369 |
+
|
| 370 |
+
pub fn augment_tool_definition(mut definition: ToolDefinition) -> ToolDefinition {
|
| 371 |
+
if definition.name != PUBLIC_TOOL_NAME {
|
| 372 |
+
definition.description = render_code_mode_sample_for_definition(&definition);
|
| 373 |
+
}
|
| 374 |
+
definition
|
| 375 |
+
}
|
| 376 |
+
|
| 377 |
+
pub fn enabled_tool_metadata(definition: &ToolDefinition) -> EnabledToolMetadata {
|
| 378 |
+
EnabledToolMetadata {
|
| 379 |
+
tool_name: definition.tool_name.clone(),
|
| 380 |
+
global_name: normalize_code_mode_identifier(&definition.name),
|
| 381 |
+
description: definition.description.clone(),
|
| 382 |
+
kind: definition.kind,
|
| 383 |
+
}
|
| 384 |
+
}
|
| 385 |
+
|
| 386 |
+
#[derive(Clone, Debug, Eq, PartialEq, Serialize)]
|
| 387 |
+
pub struct EnabledToolMetadata {
|
| 388 |
+
pub tool_name: ToolName,
|
| 389 |
+
pub global_name: String,
|
| 390 |
+
pub description: String,
|
| 391 |
+
pub kind: CodeModeToolKind,
|
| 392 |
+
}
|
| 393 |
+
|
| 394 |
+
pub fn render_code_mode_sample(
|
| 395 |
+
description: &str,
|
| 396 |
+
tool_name: &str,
|
| 397 |
+
input_name: &str,
|
| 398 |
+
input_type: String,
|
| 399 |
+
output_type: String,
|
| 400 |
+
) -> String {
|
| 401 |
+
let declaration = format!(
|
| 402 |
+
"declare const tools: {{ {} }};",
|
| 403 |
+
render_code_mode_tool_declaration(tool_name, input_name, &input_type, &output_type)
|
| 404 |
+
);
|
| 405 |
+
format!("{description}\n\nexec tool declaration:\n```ts\n{declaration}\n```")
|
| 406 |
+
}
|
| 407 |
+
|
| 408 |
+
fn render_code_mode_sample_for_definition(definition: &ToolDefinition) -> String {
|
| 409 |
+
let input_name = match definition.kind {
|
| 410 |
+
CodeModeToolKind::Function => "args",
|
| 411 |
+
CodeModeToolKind::Freeform => "input",
|
| 412 |
+
};
|
| 413 |
+
let input_type = match definition.kind {
|
| 414 |
+
CodeModeToolKind::Function => definition
|
| 415 |
+
.input_schema
|
| 416 |
+
.as_ref()
|
| 417 |
+
.map(render_json_schema_to_typescript)
|
| 418 |
+
.unwrap_or_else(|| "unknown".to_string()),
|
| 419 |
+
CodeModeToolKind::Freeform => "string".to_string(),
|
| 420 |
+
};
|
| 421 |
+
let output_type = if let Some(structured_content_schema) =
|
| 422 |
+
mcp_structured_content_schema(definition.output_schema.as_ref())
|
| 423 |
+
{
|
| 424 |
+
let structured_content_type = render_json_schema_to_typescript(structured_content_schema);
|
| 425 |
+
if structured_content_type == "unknown" {
|
| 426 |
+
"CallToolResult".to_string()
|
| 427 |
+
} else {
|
| 428 |
+
format!("CallToolResult<{structured_content_type}>")
|
| 429 |
+
}
|
| 430 |
+
} else {
|
| 431 |
+
definition
|
| 432 |
+
.output_schema
|
| 433 |
+
.as_ref()
|
| 434 |
+
.map(render_json_schema_to_typescript)
|
| 435 |
+
.unwrap_or_else(|| "unknown".to_string())
|
| 436 |
+
};
|
| 437 |
+
render_code_mode_sample(
|
| 438 |
+
&definition.description,
|
| 439 |
+
&definition.name,
|
| 440 |
+
input_name,
|
| 441 |
+
input_type,
|
| 442 |
+
output_type,
|
| 443 |
+
)
|
| 444 |
+
}
|
| 445 |
+
|
| 446 |
+
fn render_code_mode_tool_declaration(
|
| 447 |
+
tool_name: &str,
|
| 448 |
+
input_name: &str,
|
| 449 |
+
input_type: &str,
|
| 450 |
+
output_type: &str,
|
| 451 |
+
) -> String {
|
| 452 |
+
let tool_name = normalize_code_mode_identifier(tool_name);
|
| 453 |
+
format!("{tool_name}({input_name}: {input_type}): Promise<{output_type}>;")
|
| 454 |
+
}
|
| 455 |
+
|
| 456 |
+
fn render_tool_heading(global_name: &str, raw_name: &str) -> String {
|
| 457 |
+
if global_name == raw_name {
|
| 458 |
+
format!("### `{global_name}`")
|
| 459 |
+
} else {
|
| 460 |
+
format!("### `{global_name}` (`{raw_name}`)")
|
| 461 |
+
}
|
| 462 |
+
}
|
| 463 |
+
|
| 464 |
+
fn mcp_structured_content_schema(output_schema: Option<&JsonValue>) -> Option<&JsonValue> {
|
| 465 |
+
let output_schema = output_schema?;
|
| 466 |
+
let properties = output_schema
|
| 467 |
+
.get("properties")
|
| 468 |
+
.and_then(JsonValue::as_object)?;
|
| 469 |
+
let content_schema = properties.get("content").and_then(JsonValue::as_object)?;
|
| 470 |
+
if content_schema.get("type").and_then(JsonValue::as_str) != Some("array") {
|
| 471 |
+
return None;
|
| 472 |
+
}
|
| 473 |
+
|
| 474 |
+
if content_schema
|
| 475 |
+
.get("items")
|
| 476 |
+
.and_then(JsonValue::as_object)
|
| 477 |
+
.is_none_or(|items| items.get("type").and_then(JsonValue::as_str) != Some("object"))
|
| 478 |
+
{
|
| 479 |
+
return None;
|
| 480 |
+
}
|
| 481 |
+
|
| 482 |
+
if properties
|
| 483 |
+
.get("isError")
|
| 484 |
+
.and_then(JsonValue::as_object)
|
| 485 |
+
.is_none_or(|schema| schema.get("type").and_then(JsonValue::as_str) != Some("boolean"))
|
| 486 |
+
{
|
| 487 |
+
return None;
|
| 488 |
+
}
|
| 489 |
+
|
| 490 |
+
if properties
|
| 491 |
+
.get("_meta")
|
| 492 |
+
.and_then(JsonValue::as_object)
|
| 493 |
+
.is_none_or(|schema| schema.get("type").and_then(JsonValue::as_str) != Some("object"))
|
| 494 |
+
{
|
| 495 |
+
return None;
|
| 496 |
+
}
|
| 497 |
+
|
| 498 |
+
Some(
|
| 499 |
+
properties
|
| 500 |
+
.get("structuredContent")
|
| 501 |
+
.unwrap_or(&JsonValue::Bool(true)),
|
| 502 |
+
)
|
| 503 |
+
}
|
| 504 |
+
|
| 505 |
+
#[cfg(test)]
|
| 506 |
+
mod tests {
|
| 507 |
+
use super::CodeModeToolKind;
|
| 508 |
+
use super::ImageDetailVisibility;
|
| 509 |
+
use super::ParsedExecSource;
|
| 510 |
+
use super::ToolDefinition;
|
| 511 |
+
use super::ToolNamespaceDescription;
|
| 512 |
+
use super::augment_tool_definition;
|
| 513 |
+
use super::build_exec_tool_description;
|
| 514 |
+
use super::normalize_code_mode_identifier;
|
| 515 |
+
use super::parse_exec_source;
|
| 516 |
+
use codex_protocol::ToolName;
|
| 517 |
+
use pretty_assertions::assert_eq;
|
| 518 |
+
use serde_json::Value as JsonValue;
|
| 519 |
+
use serde_json::json;
|
| 520 |
+
use std::collections::BTreeMap;
|
| 521 |
+
|
| 522 |
+
fn mcp_call_tool_result_schema(structured_content_schema: JsonValue) -> JsonValue {
|
| 523 |
+
json!({
|
| 524 |
+
"type": "object",
|
| 525 |
+
"properties": {
|
| 526 |
+
"content": {
|
| 527 |
+
"type": "array",
|
| 528 |
+
"items": {
|
| 529 |
+
"type": "object"
|
| 530 |
+
}
|
| 531 |
+
},
|
| 532 |
+
"structuredContent": structured_content_schema,
|
| 533 |
+
"isError": { "type": "boolean" },
|
| 534 |
+
"_meta": { "type": "object" }
|
| 535 |
+
},
|
| 536 |
+
"required": ["content"],
|
| 537 |
+
"additionalProperties": false
|
| 538 |
+
})
|
| 539 |
+
}
|
| 540 |
+
|
| 541 |
+
#[test]
|
| 542 |
+
fn parse_exec_source_without_pragma() {
|
| 543 |
+
assert_eq!(
|
| 544 |
+
parse_exec_source("text('hi')").unwrap(),
|
| 545 |
+
ParsedExecSource {
|
| 546 |
+
code: "text('hi')".to_string(),
|
| 547 |
+
yield_time_ms: None,
|
| 548 |
+
max_output_tokens: None,
|
| 549 |
+
}
|
| 550 |
+
);
|
| 551 |
+
}
|
| 552 |
+
|
| 553 |
+
#[test]
|
| 554 |
+
fn parse_exec_source_with_pragma() {
|
| 555 |
+
assert_eq!(
|
| 556 |
+
parse_exec_source("// @exec: {\"yield_time_ms\": 10}\ntext('hi')").unwrap(),
|
| 557 |
+
ParsedExecSource {
|
| 558 |
+
code: "text('hi')".to_string(),
|
| 559 |
+
yield_time_ms: Some(10),
|
| 560 |
+
max_output_tokens: None,
|
| 561 |
+
}
|
| 562 |
+
);
|
| 563 |
+
}
|
| 564 |
+
|
| 565 |
+
#[test]
|
| 566 |
+
fn normalize_identifier_rewrites_invalid_characters() {
|
| 567 |
+
assert_eq!(
|
| 568 |
+
"mcp__ologs__get_profile",
|
| 569 |
+
normalize_code_mode_identifier("mcp__ologs__get_profile")
|
| 570 |
+
);
|
| 571 |
+
assert_eq!(
|
| 572 |
+
"hidden_dynamic_tool",
|
| 573 |
+
normalize_code_mode_identifier("hidden-dynamic-tool")
|
| 574 |
+
);
|
| 575 |
+
}
|
| 576 |
+
|
| 577 |
+
#[test]
|
| 578 |
+
fn augment_tool_definition_appends_typed_declaration() {
|
| 579 |
+
let definition = ToolDefinition {
|
| 580 |
+
name: "hidden_dynamic_tool".to_string(),
|
| 581 |
+
tool_name: ToolName::plain("hidden_dynamic_tool"),
|
| 582 |
+
description: "Test tool".to_string(),
|
| 583 |
+
kind: CodeModeToolKind::Function,
|
| 584 |
+
input_schema: Some(json!({
|
| 585 |
+
"type": "object",
|
| 586 |
+
"properties": { "city": { "type": "string" } },
|
| 587 |
+
"required": ["city"],
|
| 588 |
+
"additionalProperties": false
|
| 589 |
+
})),
|
| 590 |
+
output_schema: Some(json!({
|
| 591 |
+
"type": "object",
|
| 592 |
+
"properties": { "ok": { "type": "boolean" } },
|
| 593 |
+
"required": ["ok"]
|
| 594 |
+
})),
|
| 595 |
+
};
|
| 596 |
+
|
| 597 |
+
let description = augment_tool_definition(definition).description;
|
| 598 |
+
assert!(description.contains("declare const tools"));
|
| 599 |
+
assert!(
|
| 600 |
+
description.contains(
|
| 601 |
+
"hidden_dynamic_tool(args: { city: string; }): Promise<{ ok: boolean; }>;"
|
| 602 |
+
)
|
| 603 |
+
);
|
| 604 |
+
}
|
| 605 |
+
|
| 606 |
+
#[test]
|
| 607 |
+
fn augment_tool_definition_includes_property_descriptions_as_comments() {
|
| 608 |
+
let definition = ToolDefinition {
|
| 609 |
+
name: "weather_tool".to_string(),
|
| 610 |
+
tool_name: ToolName::plain("weather_tool"),
|
| 611 |
+
description: "Weather tool".to_string(),
|
| 612 |
+
kind: CodeModeToolKind::Function,
|
| 613 |
+
input_schema: Some(json!({
|
| 614 |
+
"type": "object",
|
| 615 |
+
"properties": {
|
| 616 |
+
"weather": {
|
| 617 |
+
"type": "array",
|
| 618 |
+
"description": "look up weather for a given list of locations",
|
| 619 |
+
"items": {
|
| 620 |
+
"type": "object",
|
| 621 |
+
"properties": {
|
| 622 |
+
"location": { "type": "string" }
|
| 623 |
+
},
|
| 624 |
+
"required": ["location"]
|
| 625 |
+
}
|
| 626 |
+
}
|
| 627 |
+
},
|
| 628 |
+
"required": ["weather"]
|
| 629 |
+
})),
|
| 630 |
+
output_schema: Some(json!({
|
| 631 |
+
"type": "object",
|
| 632 |
+
"properties": {
|
| 633 |
+
"forecast": {
|
| 634 |
+
"type": "string",
|
| 635 |
+
"description": "human readable weather forecast"
|
| 636 |
+
}
|
| 637 |
+
},
|
| 638 |
+
"required": ["forecast"]
|
| 639 |
+
})),
|
| 640 |
+
};
|
| 641 |
+
|
| 642 |
+
let description = augment_tool_definition(definition).description;
|
| 643 |
+
assert!(description.contains(
|
| 644 |
+
r#"weather_tool(args: {
|
| 645 |
+
// look up weather for a given list of locations
|
| 646 |
+
weather: Array<{ location: string; }>;
|
| 647 |
+
}): Promise<{
|
| 648 |
+
// human readable weather forecast
|
| 649 |
+
forecast: string;
|
| 650 |
+
}>;"#
|
| 651 |
+
));
|
| 652 |
+
}
|
| 653 |
+
|
| 654 |
+
#[test]
|
| 655 |
+
fn code_mode_types_structured_content_result_refs() {
|
| 656 |
+
let definition = ToolDefinition {
|
| 657 |
+
name: "mcp__sample__search".to_string(),
|
| 658 |
+
tool_name: ToolName::namespaced("mcp__sample__", "search"),
|
| 659 |
+
description: "Search".to_string(),
|
| 660 |
+
kind: CodeModeToolKind::Function,
|
| 661 |
+
input_schema: Some(json!({
|
| 662 |
+
"type": "object",
|
| 663 |
+
"properties": {},
|
| 664 |
+
"additionalProperties": false
|
| 665 |
+
})),
|
| 666 |
+
output_schema: Some(mcp_call_tool_result_schema(json!({
|
| 667 |
+
"type": "object",
|
| 668 |
+
"properties": {
|
| 669 |
+
"results": {
|
| 670 |
+
"type": "array",
|
| 671 |
+
"items": { "$ref": "#/definitions/Result~1item~0v1" }
|
| 672 |
+
}
|
| 673 |
+
},
|
| 674 |
+
"required": ["results"],
|
| 675 |
+
"additionalProperties": false,
|
| 676 |
+
"definitions": {
|
| 677 |
+
"Result/item~v1": {
|
| 678 |
+
"type": "object",
|
| 679 |
+
"properties": {
|
| 680 |
+
"id": { "type": "string" },
|
| 681 |
+
"score": { "type": "number" }
|
| 682 |
+
},
|
| 683 |
+
"required": ["id", "score"],
|
| 684 |
+
"additionalProperties": false
|
| 685 |
+
}
|
| 686 |
+
}
|
| 687 |
+
}))),
|
| 688 |
+
};
|
| 689 |
+
|
| 690 |
+
let description = augment_tool_definition(definition).description;
|
| 691 |
+
assert!(description.contains(
|
| 692 |
+
"mcp__sample__search(args: {}): Promise<CallToolResult<{ results: Array<{ id: string; score: number; }>; }>>;"
|
| 693 |
+
));
|
| 694 |
+
}
|
| 695 |
+
|
| 696 |
+
#[test]
|
| 697 |
+
fn code_mode_only_description_includes_nested_tools() {
|
| 698 |
+
let description = build_exec_tool_description(
|
| 699 |
+
&[ToolDefinition {
|
| 700 |
+
name: "foo".to_string(),
|
| 701 |
+
tool_name: ToolName::plain("foo"),
|
| 702 |
+
description: "bar".to_string(),
|
| 703 |
+
kind: CodeModeToolKind::Function,
|
| 704 |
+
input_schema: None,
|
| 705 |
+
output_schema: None,
|
| 706 |
+
}],
|
| 707 |
+
&[],
|
| 708 |
+
&BTreeMap::new(),
|
| 709 |
+
crate::DEFAULT_EXEC_YIELD_TIME_MS,
|
| 710 |
+
/*code_mode_only*/ true,
|
| 711 |
+
ImageDetailVisibility::Visible,
|
| 712 |
+
);
|
| 713 |
+
assert!(description.contains(
|
| 714 |
+
"### `foo`
|
| 715 |
+
bar"
|
| 716 |
+
));
|
| 717 |
+
assert!(!description.contains("do not attempt to use any other tools directly"));
|
| 718 |
+
}
|
| 719 |
+
|
| 720 |
+
#[test]
|
| 721 |
+
fn exec_description_mentions_timeout_helpers() {
|
| 722 |
+
let description = build_exec_tool_description(
|
| 723 |
+
&[],
|
| 724 |
+
&[],
|
| 725 |
+
&BTreeMap::new(),
|
| 726 |
+
crate::DEFAULT_EXEC_YIELD_TIME_MS,
|
| 727 |
+
/*code_mode_only*/ false,
|
| 728 |
+
ImageDetailVisibility::Visible,
|
| 729 |
+
);
|
| 730 |
+
assert!(description.contains("`audio(audioUrlOrItem:"));
|
| 731 |
+
assert!(description.contains("`setTimeout(callback: () => void, delayMs?: number)`"));
|
| 732 |
+
assert!(description.contains("`clearTimeout(timeoutId?: number)`"));
|
| 733 |
+
}
|
| 734 |
+
|
| 735 |
+
#[test]
|
| 736 |
+
fn code_mode_only_description_groups_namespace_instructions_once() {
|
| 737 |
+
let namespace_descriptions = BTreeMap::from([(
|
| 738 |
+
"mcp__sample__".to_string(),
|
| 739 |
+
ToolNamespaceDescription {
|
| 740 |
+
name: "mcp__sample".to_string(),
|
| 741 |
+
description: "Shared namespace guidance.".to_string(),
|
| 742 |
+
},
|
| 743 |
+
)]);
|
| 744 |
+
let description = build_exec_tool_description(
|
| 745 |
+
&[
|
| 746 |
+
ToolDefinition {
|
| 747 |
+
name: "mcp__sample__alpha".to_string(),
|
| 748 |
+
tool_name: ToolName::namespaced("mcp__sample__", "alpha"),
|
| 749 |
+
description: "First tool".to_string(),
|
| 750 |
+
kind: CodeModeToolKind::Function,
|
| 751 |
+
input_schema: Some(json!({
|
| 752 |
+
"type": "object",
|
| 753 |
+
"properties": {},
|
| 754 |
+
"additionalProperties": false
|
| 755 |
+
})),
|
| 756 |
+
output_schema: Some(mcp_call_tool_result_schema(json!({
|
| 757 |
+
"type": "object",
|
| 758 |
+
"properties": {},
|
| 759 |
+
"additionalProperties": false
|
| 760 |
+
}))),
|
| 761 |
+
},
|
| 762 |
+
ToolDefinition {
|
| 763 |
+
name: "mcp__sample__beta".to_string(),
|
| 764 |
+
tool_name: ToolName::namespaced("mcp__sample__", "beta"),
|
| 765 |
+
description: "Second tool".to_string(),
|
| 766 |
+
kind: CodeModeToolKind::Function,
|
| 767 |
+
input_schema: Some(json!({
|
| 768 |
+
"type": "object",
|
| 769 |
+
"properties": {},
|
| 770 |
+
"additionalProperties": false
|
| 771 |
+
})),
|
| 772 |
+
output_schema: Some(mcp_call_tool_result_schema(json!({
|
| 773 |
+
"type": "object",
|
| 774 |
+
"properties": {},
|
| 775 |
+
"additionalProperties": false
|
| 776 |
+
}))),
|
| 777 |
+
},
|
| 778 |
+
],
|
| 779 |
+
&[],
|
| 780 |
+
&namespace_descriptions,
|
| 781 |
+
crate::DEFAULT_EXEC_YIELD_TIME_MS,
|
| 782 |
+
/*code_mode_only*/ true,
|
| 783 |
+
ImageDetailVisibility::Visible,
|
| 784 |
+
);
|
| 785 |
+
assert_eq!(description.matches("## mcp__sample").count(), 1);
|
| 786 |
+
assert!(description.contains("## mcp__sample\nShared namespace guidance."));
|
| 787 |
+
assert!(description.contains(
|
| 788 |
+
"declare const tools: { mcp__sample__alpha(args: {}): Promise<CallToolResult<{}>>; };"
|
| 789 |
+
));
|
| 790 |
+
assert!(description.contains(
|
| 791 |
+
"declare const tools: { mcp__sample__beta(args: {}): Promise<CallToolResult<{}>>; };"
|
| 792 |
+
));
|
| 793 |
+
}
|
| 794 |
+
|
| 795 |
+
#[test]
|
| 796 |
+
fn code_mode_only_description_omits_empty_namespace_sections() {
|
| 797 |
+
let namespace_descriptions = BTreeMap::from([(
|
| 798 |
+
"mcp__sample__".to_string(),
|
| 799 |
+
ToolNamespaceDescription {
|
| 800 |
+
name: "mcp__sample".to_string(),
|
| 801 |
+
description: String::new(),
|
| 802 |
+
},
|
| 803 |
+
)]);
|
| 804 |
+
let description = build_exec_tool_description(
|
| 805 |
+
&[ToolDefinition {
|
| 806 |
+
name: "mcp__sample__alpha".to_string(),
|
| 807 |
+
tool_name: ToolName::namespaced("mcp__sample__", "alpha"),
|
| 808 |
+
description: "First tool".to_string(),
|
| 809 |
+
kind: CodeModeToolKind::Function,
|
| 810 |
+
input_schema: Some(json!({
|
| 811 |
+
"type": "object",
|
| 812 |
+
"properties": {},
|
| 813 |
+
"additionalProperties": false
|
| 814 |
+
})),
|
| 815 |
+
output_schema: Some(mcp_call_tool_result_schema(json!({
|
| 816 |
+
"type": "object",
|
| 817 |
+
"properties": {},
|
| 818 |
+
"additionalProperties": false
|
| 819 |
+
}))),
|
| 820 |
+
}],
|
| 821 |
+
&[],
|
| 822 |
+
&namespace_descriptions,
|
| 823 |
+
crate::DEFAULT_EXEC_YIELD_TIME_MS,
|
| 824 |
+
/*code_mode_only*/ true,
|
| 825 |
+
ImageDetailVisibility::Visible,
|
| 826 |
+
);
|
| 827 |
+
|
| 828 |
+
assert!(!description.contains("## mcp__sample"));
|
| 829 |
+
assert!(description.contains("### `mcp__sample__alpha`"));
|
| 830 |
+
}
|
| 831 |
+
|
| 832 |
+
#[test]
|
| 833 |
+
fn code_mode_only_description_renders_shared_mcp_types_once() {
|
| 834 |
+
let first_tool = augment_tool_definition(ToolDefinition {
|
| 835 |
+
name: "mcp__sample__alpha".to_string(),
|
| 836 |
+
tool_name: ToolName::namespaced("mcp__sample__", "alpha"),
|
| 837 |
+
description: "First tool".to_string(),
|
| 838 |
+
kind: CodeModeToolKind::Function,
|
| 839 |
+
input_schema: Some(json!({
|
| 840 |
+
"type": "object",
|
| 841 |
+
"properties": {},
|
| 842 |
+
"additionalProperties": false
|
| 843 |
+
})),
|
| 844 |
+
output_schema: Some(json!({
|
| 845 |
+
"type": "object",
|
| 846 |
+
"properties": {
|
| 847 |
+
"content": {
|
| 848 |
+
"type": "array",
|
| 849 |
+
"items": {
|
| 850 |
+
"type": "object"
|
| 851 |
+
}
|
| 852 |
+
},
|
| 853 |
+
"structuredContent": {
|
| 854 |
+
"type": "object",
|
| 855 |
+
"properties": {
|
| 856 |
+
"echo": { "type": "string" }
|
| 857 |
+
},
|
| 858 |
+
"required": ["echo"],
|
| 859 |
+
"additionalProperties": false
|
| 860 |
+
},
|
| 861 |
+
"isError": { "type": "boolean" },
|
| 862 |
+
"_meta": { "type": "object" }
|
| 863 |
+
},
|
| 864 |
+
"required": ["content"],
|
| 865 |
+
"additionalProperties": false
|
| 866 |
+
})),
|
| 867 |
+
});
|
| 868 |
+
let second_tool = augment_tool_definition(ToolDefinition {
|
| 869 |
+
name: "mcp__sample__beta".to_string(),
|
| 870 |
+
tool_name: ToolName::namespaced("mcp__sample__", "beta"),
|
| 871 |
+
description: "Second tool".to_string(),
|
| 872 |
+
kind: CodeModeToolKind::Function,
|
| 873 |
+
input_schema: Some(json!({
|
| 874 |
+
"type": "object",
|
| 875 |
+
"properties": {},
|
| 876 |
+
"additionalProperties": false
|
| 877 |
+
})),
|
| 878 |
+
output_schema: Some(json!({
|
| 879 |
+
"type": "object",
|
| 880 |
+
"properties": {
|
| 881 |
+
"content": {
|
| 882 |
+
"type": "array",
|
| 883 |
+
"items": {
|
| 884 |
+
"type": "object"
|
| 885 |
+
}
|
| 886 |
+
},
|
| 887 |
+
"structuredContent": {
|
| 888 |
+
"type": "object",
|
| 889 |
+
"properties": {
|
| 890 |
+
"count": { "type": "integer" }
|
| 891 |
+
},
|
| 892 |
+
"required": ["count"],
|
| 893 |
+
"additionalProperties": false
|
| 894 |
+
},
|
| 895 |
+
"isError": { "type": "boolean" },
|
| 896 |
+
"_meta": { "type": "object" }
|
| 897 |
+
},
|
| 898 |
+
"required": ["content"],
|
| 899 |
+
"additionalProperties": false
|
| 900 |
+
})),
|
| 901 |
+
});
|
| 902 |
+
|
| 903 |
+
let description = build_exec_tool_description(
|
| 904 |
+
&[
|
| 905 |
+
ToolDefinition {
|
| 906 |
+
name: first_tool.name,
|
| 907 |
+
tool_name: first_tool.tool_name,
|
| 908 |
+
description: "First tool".to_string(),
|
| 909 |
+
kind: first_tool.kind,
|
| 910 |
+
input_schema: first_tool.input_schema,
|
| 911 |
+
output_schema: first_tool.output_schema,
|
| 912 |
+
},
|
| 913 |
+
ToolDefinition {
|
| 914 |
+
name: second_tool.name,
|
| 915 |
+
tool_name: second_tool.tool_name,
|
| 916 |
+
description: "Second tool".to_string(),
|
| 917 |
+
kind: second_tool.kind,
|
| 918 |
+
input_schema: second_tool.input_schema,
|
| 919 |
+
output_schema: second_tool.output_schema,
|
| 920 |
+
},
|
| 921 |
+
],
|
| 922 |
+
&[],
|
| 923 |
+
&BTreeMap::new(),
|
| 924 |
+
crate::DEFAULT_EXEC_YIELD_TIME_MS,
|
| 925 |
+
/*code_mode_only*/ true,
|
| 926 |
+
ImageDetailVisibility::Visible,
|
| 927 |
+
);
|
| 928 |
+
|
| 929 |
+
assert_eq!(
|
| 930 |
+
description
|
| 931 |
+
.matches("type CallToolResult<TStructured = { [key: string]: unknown }>")
|
| 932 |
+
.count(),
|
| 933 |
+
1
|
| 934 |
+
);
|
| 935 |
+
assert_eq!(description.matches("Shared MCP Types:").count(), 1);
|
| 936 |
+
}
|
| 937 |
+
|
| 938 |
+
#[test]
|
| 939 |
+
fn code_mode_only_description_renders_shared_mcp_types_for_deferred_tools() {
|
| 940 |
+
let deferred_tool = ToolDefinition {
|
| 941 |
+
name: "mcp__sample__alpha".to_string(),
|
| 942 |
+
tool_name: ToolName::namespaced("mcp__sample__", "alpha"),
|
| 943 |
+
description: "Deferred tool".to_string(),
|
| 944 |
+
kind: CodeModeToolKind::Function,
|
| 945 |
+
input_schema: Some(json!({
|
| 946 |
+
"type": "object",
|
| 947 |
+
"properties": {},
|
| 948 |
+
"additionalProperties": false
|
| 949 |
+
})),
|
| 950 |
+
output_schema: Some(mcp_call_tool_result_schema(json!({
|
| 951 |
+
"type": "object",
|
| 952 |
+
"properties": {},
|
| 953 |
+
"additionalProperties": false
|
| 954 |
+
}))),
|
| 955 |
+
};
|
| 956 |
+
|
| 957 |
+
let description = build_exec_tool_description(
|
| 958 |
+
&[],
|
| 959 |
+
&[deferred_tool],
|
| 960 |
+
&BTreeMap::new(),
|
| 961 |
+
crate::DEFAULT_EXEC_YIELD_TIME_MS,
|
| 962 |
+
/*code_mode_only*/ true,
|
| 963 |
+
ImageDetailVisibility::Visible,
|
| 964 |
+
);
|
| 965 |
+
|
| 966 |
+
assert!(description.contains("Some deferred nested tools may be omitted"));
|
| 967 |
+
assert!(description.contains("Shared MCP Types:"));
|
| 968 |
+
assert!(!description.contains("### `mcp__sample__alpha`"));
|
| 969 |
+
}
|
| 970 |
+
|
| 971 |
+
#[test]
|
| 972 |
+
fn exec_description_mentions_deferred_nested_tools_when_available() {
|
| 973 |
+
let description = build_exec_tool_description(
|
| 974 |
+
&[],
|
| 975 |
+
&[ToolDefinition {
|
| 976 |
+
name: "deferred_tool".to_string(),
|
| 977 |
+
tool_name: ToolName::plain("deferred_tool"),
|
| 978 |
+
description: "Deferred tool".to_string(),
|
| 979 |
+
kind: CodeModeToolKind::Function,
|
| 980 |
+
input_schema: None,
|
| 981 |
+
output_schema: None,
|
| 982 |
+
}],
|
| 983 |
+
&BTreeMap::new(),
|
| 984 |
+
crate::DEFAULT_EXEC_YIELD_TIME_MS,
|
| 985 |
+
/*code_mode_only*/ false,
|
| 986 |
+
ImageDetailVisibility::Visible,
|
| 987 |
+
);
|
| 988 |
+
|
| 989 |
+
assert!(description.contains("Some deferred nested tools may be omitted"));
|
| 990 |
+
assert!(description.contains("filter `ALL_TOOLS` by `name` and `description`"));
|
| 991 |
+
assert!(!description.contains("do not print the full `ALL_TOOLS` array"));
|
| 992 |
+
}
|
| 993 |
+
}
|
codex-rs/code-mode-protocol/src/grpc/codex.code_mode.v1.proto
ADDED
|
@@ -0,0 +1,269 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
syntax = "proto3";
|
| 2 |
+
|
| 3 |
+
package codex.code_mode.v1;
|
| 4 |
+
|
| 5 |
+
// Hosts stateful JavaScript execution and delegates nested tool calls to the
|
| 6 |
+
// session owner. Large tool inputs and outputs stay off the session event
|
| 7 |
+
// stream so independent HTTP/2 streams can make progress concurrently.
|
| 8 |
+
service CodeModeHost {
|
| 9 |
+
// Opens a session lease. The first event is always SessionOpened; dropping
|
| 10 |
+
// this stream closes the session and terminates its active cells.
|
| 11 |
+
rpc OpenSession(OpenSessionRequest) returns (stream SessionEvent);
|
| 12 |
+
rpc CloseSession(CloseSessionRequest) returns (CloseSessionResponse);
|
| 13 |
+
|
| 14 |
+
// Each subscription owns an independent stream of matching invocations. An
|
| 15 |
+
// empty tool_names filter matches every tool. Each invocation is routed to
|
| 16 |
+
// exactly one matching subscription, even when filters overlap.
|
| 17 |
+
rpc SubscribeToToolCalls(SubscribeToToolCallsRequest)
|
| 18 |
+
returns (stream ToolCall);
|
| 19 |
+
|
| 20 |
+
// Each result receives its own HTTP/2 stream, preventing a large response
|
| 21 |
+
// from blocking unrelated tool completions or session control events.
|
| 22 |
+
rpc CompleteToolCall(CompleteToolCallRequest)
|
| 23 |
+
returns (CompleteToolCallResponse);
|
| 24 |
+
rpc AcknowledgeNotification(AcknowledgeNotificationRequest)
|
| 25 |
+
returns (AcknowledgeNotificationResponse);
|
| 26 |
+
|
| 27 |
+
// Emits ExecutionStarted immediately, followed by one ExecutionOutcome when
|
| 28 |
+
// the execution yields, completes, or is terminated.
|
| 29 |
+
rpc Execute(ExecuteRequest) returns (stream ExecuteEvent);
|
| 30 |
+
rpc Wait(WaitRequest) returns (WaitResponse);
|
| 31 |
+
|
| 32 |
+
// Acknowledges that a canceled wait has retired before another wait starts.
|
| 33 |
+
rpc CancelWait(CancelWaitRequest) returns (CancelWaitResponse);
|
| 34 |
+
rpc Terminate(TerminateRequest) returns (WaitResponse);
|
| 35 |
+
}
|
| 36 |
+
|
| 37 |
+
message OpenSessionRequest {
|
| 38 |
+
optional SessionCellExecutionLimits cell_execution_limits = 1;
|
| 39 |
+
}
|
| 40 |
+
|
| 41 |
+
message SessionCellExecutionLimits {
|
| 42 |
+
optional uint64 max_yield_time_ms = 1;
|
| 43 |
+
optional uint64 max_heap_size_bytes = 2;
|
| 44 |
+
}
|
| 45 |
+
|
| 46 |
+
message SessionEvent {
|
| 47 |
+
oneof event {
|
| 48 |
+
SessionOpened opened = 1;
|
| 49 |
+
ToolCallCancelled tool_call_cancelled = 2;
|
| 50 |
+
Notification notification = 3;
|
| 51 |
+
NotificationCancelled notification_cancelled = 4;
|
| 52 |
+
CellClosed cell_closed = 5;
|
| 53 |
+
}
|
| 54 |
+
}
|
| 55 |
+
|
| 56 |
+
message SessionOpened {
|
| 57 |
+
string session_id = 1;
|
| 58 |
+
}
|
| 59 |
+
|
| 60 |
+
message CloseSessionRequest {
|
| 61 |
+
string session_id = 1;
|
| 62 |
+
}
|
| 63 |
+
|
| 64 |
+
message CloseSessionResponse {}
|
| 65 |
+
|
| 66 |
+
message SubscribeToToolCallsRequest {
|
| 67 |
+
string session_id = 1;
|
| 68 |
+
repeated ToolName tool_names = 2;
|
| 69 |
+
}
|
| 70 |
+
|
| 71 |
+
message ToolCall {
|
| 72 |
+
string session_id = 1;
|
| 73 |
+
|
| 74 |
+
// Correlates callbacks with Execute before ExecutionStarted is received.
|
| 75 |
+
string execution_id = 2;
|
| 76 |
+
string cell_id = 3;
|
| 77 |
+
string invocation_id = 4;
|
| 78 |
+
string runtime_tool_call_id = 5;
|
| 79 |
+
ToolName tool_name = 6;
|
| 80 |
+
ToolKind tool_kind = 7;
|
| 81 |
+
optional bytes input_json = 8;
|
| 82 |
+
|
| 83 |
+
// Starts at one and increases independently for each execution.
|
| 84 |
+
uint64 sequence = 9;
|
| 85 |
+
|
| 86 |
+
// The host tool invocation span's W3C parent context for this streamed callback.
|
| 87 |
+
// gRPC metadata is fixed when the subscription stream opens, so each call
|
| 88 |
+
// carries its own context in the message.
|
| 89 |
+
optional string traceparent = 10;
|
| 90 |
+
}
|
| 91 |
+
|
| 92 |
+
message CompleteToolCallRequest {
|
| 93 |
+
string session_id = 1;
|
| 94 |
+
string invocation_id = 2;
|
| 95 |
+
|
| 96 |
+
oneof outcome {
|
| 97 |
+
ToolCallSucceeded succeeded = 3;
|
| 98 |
+
ToolCallFailed failed = 4;
|
| 99 |
+
}
|
| 100 |
+
}
|
| 101 |
+
|
| 102 |
+
message ToolCallSucceeded {
|
| 103 |
+
bytes output_json = 1;
|
| 104 |
+
}
|
| 105 |
+
|
| 106 |
+
message ToolCallFailed {
|
| 107 |
+
string message = 1;
|
| 108 |
+
}
|
| 109 |
+
|
| 110 |
+
message CompleteToolCallResponse {}
|
| 111 |
+
|
| 112 |
+
message ToolCallCancelled {
|
| 113 |
+
string invocation_id = 1;
|
| 114 |
+
|
| 115 |
+
// Cancellation can arrive before the corresponding ToolCall because session
|
| 116 |
+
// control events and tool subscriptions use independent HTTP/2 streams.
|
| 117 |
+
}
|
| 118 |
+
|
| 119 |
+
message Notification {
|
| 120 |
+
string notification_id = 1;
|
| 121 |
+
string execution_id = 2;
|
| 122 |
+
string cell_id = 3;
|
| 123 |
+
string call_id = 4;
|
| 124 |
+
string text = 5;
|
| 125 |
+
}
|
| 126 |
+
|
| 127 |
+
message NotificationCancelled {
|
| 128 |
+
string notification_id = 1;
|
| 129 |
+
}
|
| 130 |
+
|
| 131 |
+
message AcknowledgeNotificationRequest {
|
| 132 |
+
string session_id = 1;
|
| 133 |
+
string notification_id = 2;
|
| 134 |
+
}
|
| 135 |
+
|
| 136 |
+
message AcknowledgeNotificationResponse {}
|
| 137 |
+
|
| 138 |
+
message CellClosed {
|
| 139 |
+
string execution_id = 1;
|
| 140 |
+
string cell_id = 2;
|
| 141 |
+
|
| 142 |
+
// Last tool-call sequence issued before closure. Clients may retire the cell
|
| 143 |
+
// immediately and reject tool calls delivered after its closure.
|
| 144 |
+
uint64 final_tool_call_sequence = 3;
|
| 145 |
+
}
|
| 146 |
+
|
| 147 |
+
message ExecuteRequest {
|
| 148 |
+
string session_id = 1;
|
| 149 |
+
|
| 150 |
+
// Chosen by the client so callbacks can be correlated before cell admission.
|
| 151 |
+
string execution_id = 2;
|
| 152 |
+
string tool_call_id = 3;
|
| 153 |
+
string source = 4;
|
| 154 |
+
repeated ToolDefinition enabled_tools = 5;
|
| 155 |
+
optional uint64 yield_time_ms = 6;
|
| 156 |
+
optional uint64 max_output_tokens = 7;
|
| 157 |
+
}
|
| 158 |
+
|
| 159 |
+
message ExecuteEvent {
|
| 160 |
+
oneof event {
|
| 161 |
+
ExecutionStarted started = 1;
|
| 162 |
+
ExecutionOutcome outcome = 2;
|
| 163 |
+
}
|
| 164 |
+
}
|
| 165 |
+
|
| 166 |
+
message ExecutionStarted {
|
| 167 |
+
string execution_id = 1;
|
| 168 |
+
string cell_id = 2;
|
| 169 |
+
}
|
| 170 |
+
|
| 171 |
+
message WaitRequest {
|
| 172 |
+
string session_id = 1;
|
| 173 |
+
string cell_id = 2;
|
| 174 |
+
string wait_id = 3;
|
| 175 |
+
uint64 yield_time_ms = 4;
|
| 176 |
+
}
|
| 177 |
+
|
| 178 |
+
message WaitResponse {
|
| 179 |
+
oneof state {
|
| 180 |
+
ExecutionOutcome live_cell = 1;
|
| 181 |
+
ExecutionOutcome missing_cell = 2;
|
| 182 |
+
}
|
| 183 |
+
}
|
| 184 |
+
|
| 185 |
+
message CancelWaitRequest {
|
| 186 |
+
string session_id = 1;
|
| 187 |
+
string wait_id = 2;
|
| 188 |
+
}
|
| 189 |
+
|
| 190 |
+
message CancelWaitResponse {}
|
| 191 |
+
|
| 192 |
+
message TerminateRequest {
|
| 193 |
+
string session_id = 1;
|
| 194 |
+
string cell_id = 2;
|
| 195 |
+
}
|
| 196 |
+
|
| 197 |
+
message ExecutionOutcome {
|
| 198 |
+
string cell_id = 1;
|
| 199 |
+
repeated ContentItem content_items = 2;
|
| 200 |
+
|
| 201 |
+
// Elapsed monotonic time from receipt of this Execute, Wait, or Terminate
|
| 202 |
+
// request until its outcome is ready, before serialization and delivery.
|
| 203 |
+
// Includes nested tool waits, but not background time between requests.
|
| 204 |
+
// Always supplied by the host; zero is a valid measurement.
|
| 205 |
+
uint64 code_mode_host_duration_ns = 6;
|
| 206 |
+
|
| 207 |
+
oneof outcome {
|
| 208 |
+
ExecutionYielded yielded = 3;
|
| 209 |
+
ExecutionTerminated terminated = 4;
|
| 210 |
+
ExecutionCompleted completed = 5;
|
| 211 |
+
}
|
| 212 |
+
}
|
| 213 |
+
|
| 214 |
+
message ExecutionYielded {}
|
| 215 |
+
|
| 216 |
+
message ExecutionTerminated {}
|
| 217 |
+
|
| 218 |
+
message ExecutionCompleted {
|
| 219 |
+
optional string error_text = 1;
|
| 220 |
+
}
|
| 221 |
+
|
| 222 |
+
message ToolDefinition {
|
| 223 |
+
string name = 1;
|
| 224 |
+
ToolName tool_name = 2;
|
| 225 |
+
string description = 3;
|
| 226 |
+
ToolKind kind = 4;
|
| 227 |
+
optional bytes input_schema_json = 5;
|
| 228 |
+
optional bytes output_schema_json = 6;
|
| 229 |
+
}
|
| 230 |
+
|
| 231 |
+
message ToolName {
|
| 232 |
+
string name = 1;
|
| 233 |
+
optional string namespace = 2;
|
| 234 |
+
}
|
| 235 |
+
|
| 236 |
+
enum ToolKind {
|
| 237 |
+
TOOL_KIND_UNSPECIFIED = 0;
|
| 238 |
+
TOOL_KIND_FUNCTION = 1;
|
| 239 |
+
TOOL_KIND_FREEFORM = 2;
|
| 240 |
+
}
|
| 241 |
+
|
| 242 |
+
message ContentItem {
|
| 243 |
+
oneof item {
|
| 244 |
+
TextContent text = 1;
|
| 245 |
+
ImageContent image = 2;
|
| 246 |
+
AudioContent audio = 3;
|
| 247 |
+
}
|
| 248 |
+
}
|
| 249 |
+
|
| 250 |
+
message TextContent {
|
| 251 |
+
string text = 1;
|
| 252 |
+
}
|
| 253 |
+
|
| 254 |
+
message ImageContent {
|
| 255 |
+
string image_url = 1;
|
| 256 |
+
optional ImageDetail detail = 2;
|
| 257 |
+
}
|
| 258 |
+
|
| 259 |
+
message AudioContent {
|
| 260 |
+
string audio_url = 1;
|
| 261 |
+
}
|
| 262 |
+
|
| 263 |
+
enum ImageDetail {
|
| 264 |
+
IMAGE_DETAIL_UNSPECIFIED = 0;
|
| 265 |
+
IMAGE_DETAIL_AUTO = 1;
|
| 266 |
+
IMAGE_DETAIL_LOW = 2;
|
| 267 |
+
IMAGE_DETAIL_HIGH = 3;
|
| 268 |
+
IMAGE_DETAIL_ORIGINAL = 4;
|
| 269 |
+
}
|
codex-rs/code-mode-protocol/src/grpc/mod.rs
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#[cfg(codex_bazel)]
|
| 2 |
+
pub use code_mode_proto::codex::code_mode::v1::*;
|
| 3 |
+
|
| 4 |
+
#[cfg(not(codex_bazel))]
|
| 5 |
+
tonic::include_proto!("codex.code_mode.v1");
|
| 6 |
+
|
| 7 |
+
pub const MAX_IDENTIFIER_BYTES: usize = 256;
|
codex-rs/code-mode-protocol/src/host/codec.rs
ADDED
|
@@ -0,0 +1,170 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::io;
|
| 2 |
+
use std::mem::size_of;
|
| 3 |
+
|
| 4 |
+
use serde::Serialize;
|
| 5 |
+
use serde::de::DeserializeOwned;
|
| 6 |
+
use tokio::io::AsyncRead;
|
| 7 |
+
use tokio::io::AsyncReadExt;
|
| 8 |
+
use tokio::io::AsyncWrite;
|
| 9 |
+
use tokio::io::AsyncWriteExt;
|
| 10 |
+
|
| 11 |
+
/// Maximum JSON payload size accepted for one code-mode host frame.
|
| 12 |
+
pub const MAX_FRAME_BYTES: usize = 64 * 1024 * 1024;
|
| 13 |
+
|
| 14 |
+
/// A serialized IPC frame that has already passed the payload size limit.
|
| 15 |
+
#[derive(Clone, Debug)]
|
| 16 |
+
pub struct EncodedFrame {
|
| 17 |
+
payload: Vec<u8>,
|
| 18 |
+
}
|
| 19 |
+
|
| 20 |
+
impl EncodedFrame {
|
| 21 |
+
pub fn encode<T>(message: &T) -> io::Result<Self>
|
| 22 |
+
where
|
| 23 |
+
T: Serialize,
|
| 24 |
+
{
|
| 25 |
+
let payload = serde_json::to_vec(message).map_err(|err| {
|
| 26 |
+
io::Error::new(
|
| 27 |
+
io::ErrorKind::InvalidData,
|
| 28 |
+
format!("failed to encode code-mode IPC frame: {err}"),
|
| 29 |
+
)
|
| 30 |
+
})?;
|
| 31 |
+
if payload.len() > MAX_FRAME_BYTES {
|
| 32 |
+
return Err(io::Error::new(
|
| 33 |
+
io::ErrorKind::InvalidData,
|
| 34 |
+
format!(
|
| 35 |
+
"code-mode IPC frame length {} exceeds {MAX_FRAME_BYTES} bytes",
|
| 36 |
+
payload.len()
|
| 37 |
+
),
|
| 38 |
+
));
|
| 39 |
+
}
|
| 40 |
+
Ok(Self { payload })
|
| 41 |
+
}
|
| 42 |
+
|
| 43 |
+
/// Returns the complete length-prefixed representation of this frame.
|
| 44 |
+
pub fn into_framed_bytes(self) -> Vec<u8> {
|
| 45 |
+
let mut bytes = Vec::with_capacity(size_of::<u32>() + self.payload.len());
|
| 46 |
+
bytes.extend_from_slice(&(self.payload.len() as u32).to_le_bytes());
|
| 47 |
+
bytes.extend_from_slice(&self.payload);
|
| 48 |
+
bytes
|
| 49 |
+
}
|
| 50 |
+
|
| 51 |
+
/// Decodes exactly one complete length-prefixed frame.
|
| 52 |
+
pub fn decode_framed<T>(bytes: &[u8]) -> io::Result<T>
|
| 53 |
+
where
|
| 54 |
+
T: DeserializeOwned,
|
| 55 |
+
{
|
| 56 |
+
let length_bytes: [u8; size_of::<u32>()] = bytes
|
| 57 |
+
.get(..size_of::<u32>())
|
| 58 |
+
.and_then(|length_bytes| length_bytes.try_into().ok())
|
| 59 |
+
.ok_or_else(|| {
|
| 60 |
+
io::Error::new(
|
| 61 |
+
io::ErrorKind::InvalidData,
|
| 62 |
+
"code-mode IPC frame is missing its length prefix",
|
| 63 |
+
)
|
| 64 |
+
})?;
|
| 65 |
+
let length = u32::from_le_bytes(length_bytes) as usize;
|
| 66 |
+
if length > MAX_FRAME_BYTES {
|
| 67 |
+
return Err(io::Error::new(
|
| 68 |
+
io::ErrorKind::InvalidData,
|
| 69 |
+
format!("code-mode IPC frame length {length} exceeds {MAX_FRAME_BYTES} bytes"),
|
| 70 |
+
));
|
| 71 |
+
}
|
| 72 |
+
|
| 73 |
+
let payload = &bytes[size_of::<u32>()..];
|
| 74 |
+
if payload.len() != length {
|
| 75 |
+
return Err(io::Error::new(
|
| 76 |
+
io::ErrorKind::InvalidData,
|
| 77 |
+
format!(
|
| 78 |
+
"code-mode IPC frame declares {length} payload bytes but contains {}",
|
| 79 |
+
payload.len()
|
| 80 |
+
),
|
| 81 |
+
));
|
| 82 |
+
}
|
| 83 |
+
|
| 84 |
+
serde_json::from_slice(payload).map_err(|err| {
|
| 85 |
+
io::Error::new(
|
| 86 |
+
io::ErrorKind::InvalidData,
|
| 87 |
+
format!("failed to decode code-mode IPC frame: {err}"),
|
| 88 |
+
)
|
| 89 |
+
})
|
| 90 |
+
}
|
| 91 |
+
}
|
| 92 |
+
|
| 93 |
+
/// Decodes JSON messages prefixed by a four-byte little-endian payload length.
|
| 94 |
+
pub struct FramedReader<R> {
|
| 95 |
+
reader: R,
|
| 96 |
+
}
|
| 97 |
+
|
| 98 |
+
impl<R> FramedReader<R>
|
| 99 |
+
where
|
| 100 |
+
R: AsyncRead + Unpin,
|
| 101 |
+
{
|
| 102 |
+
pub fn new(reader: R) -> Self {
|
| 103 |
+
Self { reader }
|
| 104 |
+
}
|
| 105 |
+
|
| 106 |
+
/// Reads the next frame, returning `None` only for EOF at a frame boundary.
|
| 107 |
+
pub async fn read<T>(&mut self) -> io::Result<Option<T>>
|
| 108 |
+
where
|
| 109 |
+
T: DeserializeOwned,
|
| 110 |
+
{
|
| 111 |
+
let mut length_bytes = [0_u8; size_of::<u32>()];
|
| 112 |
+
if self.reader.read(&mut length_bytes[..1]).await? == 0 {
|
| 113 |
+
return Ok(None);
|
| 114 |
+
}
|
| 115 |
+
self.reader.read_exact(&mut length_bytes[1..]).await?;
|
| 116 |
+
|
| 117 |
+
let length = u32::from_le_bytes(length_bytes) as usize;
|
| 118 |
+
if length > MAX_FRAME_BYTES {
|
| 119 |
+
return Err(io::Error::new(
|
| 120 |
+
io::ErrorKind::InvalidData,
|
| 121 |
+
format!("code-mode IPC frame length {length} exceeds {MAX_FRAME_BYTES} bytes"),
|
| 122 |
+
));
|
| 123 |
+
}
|
| 124 |
+
|
| 125 |
+
let mut payload = vec![0; length];
|
| 126 |
+
self.reader.read_exact(&mut payload).await?;
|
| 127 |
+
serde_json::from_slice(&payload).map(Some).map_err(|err| {
|
| 128 |
+
io::Error::new(
|
| 129 |
+
io::ErrorKind::InvalidData,
|
| 130 |
+
format!("failed to decode code-mode IPC frame: {err}"),
|
| 131 |
+
)
|
| 132 |
+
})
|
| 133 |
+
}
|
| 134 |
+
}
|
| 135 |
+
|
| 136 |
+
/// Encodes JSON messages with a four-byte little-endian payload length.
|
| 137 |
+
pub struct FramedWriter<W> {
|
| 138 |
+
writer: W,
|
| 139 |
+
}
|
| 140 |
+
|
| 141 |
+
impl<W> FramedWriter<W>
|
| 142 |
+
where
|
| 143 |
+
W: AsyncWrite + Unpin,
|
| 144 |
+
{
|
| 145 |
+
pub fn new(writer: W) -> Self {
|
| 146 |
+
Self { writer }
|
| 147 |
+
}
|
| 148 |
+
|
| 149 |
+
/// Writes and flushes one complete frame.
|
| 150 |
+
pub async fn write<T>(&mut self, message: &T) -> io::Result<()>
|
| 151 |
+
where
|
| 152 |
+
T: Serialize,
|
| 153 |
+
{
|
| 154 |
+
self.write_frame(&EncodedFrame::encode(message)?).await
|
| 155 |
+
}
|
| 156 |
+
|
| 157 |
+
/// Writes and flushes a frame encoded before it entered an I/O queue.
|
| 158 |
+
pub async fn write_frame(&mut self, frame: &EncodedFrame) -> io::Result<()> {
|
| 159 |
+
let length = u32::try_from(frame.payload.len()).map_err(|_| {
|
| 160 |
+
io::Error::new(
|
| 161 |
+
io::ErrorKind::InvalidData,
|
| 162 |
+
"code-mode IPC frame length exceeds u32",
|
| 163 |
+
)
|
| 164 |
+
})?;
|
| 165 |
+
|
| 166 |
+
self.writer.write_all(&length.to_le_bytes()).await?;
|
| 167 |
+
self.writer.write_all(&frame.payload).await?;
|
| 168 |
+
self.writer.flush().await
|
| 169 |
+
}
|
| 170 |
+
}
|
codex-rs/code-mode-protocol/src/host/codec_tests.rs
ADDED
|
@@ -0,0 +1,137 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use pretty_assertions::assert_eq;
|
| 2 |
+
use serde_json::json;
|
| 3 |
+
use tokio::io::AsyncReadExt;
|
| 4 |
+
use tokio::io::AsyncWriteExt;
|
| 5 |
+
|
| 6 |
+
use super::EncodedFrame;
|
| 7 |
+
use super::FramedReader;
|
| 8 |
+
use super::FramedWriter;
|
| 9 |
+
use super::MAX_FRAME_BYTES;
|
| 10 |
+
|
| 11 |
+
#[test]
|
| 12 |
+
fn complete_frame_round_trips_without_a_byte_stream() {
|
| 13 |
+
let value = json!({"type": "session/open", "sessionId": "session-1"});
|
| 14 |
+
let bytes = EncodedFrame::encode(&value)
|
| 15 |
+
.expect("encode frame")
|
| 16 |
+
.into_framed_bytes();
|
| 17 |
+
|
| 18 |
+
assert_eq!(
|
| 19 |
+
EncodedFrame::decode_framed::<serde_json::Value>(&bytes).expect("decode frame"),
|
| 20 |
+
value
|
| 21 |
+
);
|
| 22 |
+
}
|
| 23 |
+
|
| 24 |
+
#[test]
|
| 25 |
+
fn complete_frame_rejects_truncated_and_trailing_payloads() {
|
| 26 |
+
let value = json!({"value": 1});
|
| 27 |
+
let bytes = EncodedFrame::encode(&value)
|
| 28 |
+
.expect("encode frame")
|
| 29 |
+
.into_framed_bytes();
|
| 30 |
+
|
| 31 |
+
let truncated = &bytes[..bytes.len() - 1];
|
| 32 |
+
let truncated_error = EncodedFrame::decode_framed::<serde_json::Value>(truncated)
|
| 33 |
+
.expect_err("truncated frame should fail");
|
| 34 |
+
assert_eq!(truncated_error.kind(), std::io::ErrorKind::InvalidData);
|
| 35 |
+
|
| 36 |
+
let mut trailing = bytes;
|
| 37 |
+
trailing.push(0);
|
| 38 |
+
let trailing_error = EncodedFrame::decode_framed::<serde_json::Value>(&trailing)
|
| 39 |
+
.expect_err("frame with trailing bytes should fail");
|
| 40 |
+
assert_eq!(trailing_error.kind(), std::io::ErrorKind::InvalidData);
|
| 41 |
+
}
|
| 42 |
+
|
| 43 |
+
#[tokio::test]
|
| 44 |
+
async fn frame_wire_format_is_little_endian_length_prefixed_json() {
|
| 45 |
+
let (writer, mut reader) = tokio::io::duplex(/*max_buf_size*/ 128);
|
| 46 |
+
let write = tokio::spawn(async move {
|
| 47 |
+
FramedWriter::new(writer)
|
| 48 |
+
.write(&json!({"value": 1}))
|
| 49 |
+
.await
|
| 50 |
+
.expect("write frame");
|
| 51 |
+
});
|
| 52 |
+
|
| 53 |
+
let mut bytes = Vec::new();
|
| 54 |
+
reader.read_to_end(&mut bytes).await.expect("read bytes");
|
| 55 |
+
write.await.expect("writer task");
|
| 56 |
+
|
| 57 |
+
let payload = br#"{"value":1}"#;
|
| 58 |
+
let mut expected = (payload.len() as u32).to_le_bytes().to_vec();
|
| 59 |
+
expected.extend_from_slice(payload);
|
| 60 |
+
assert_eq!(bytes, expected);
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
+
#[tokio::test]
|
| 64 |
+
async fn fragmented_frame_round_trips() {
|
| 65 |
+
let value = json!({"type": "session/open", "sessionId": "session-1"});
|
| 66 |
+
let payload = serde_json::to_vec(&value).expect("serialize");
|
| 67 |
+
let mut bytes = (payload.len() as u32).to_le_bytes().to_vec();
|
| 68 |
+
bytes.extend(payload);
|
| 69 |
+
|
| 70 |
+
let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 128);
|
| 71 |
+
let write = tokio::spawn(async move {
|
| 72 |
+
for byte in bytes {
|
| 73 |
+
writer.write_all(&[byte]).await.expect("write byte");
|
| 74 |
+
tokio::task::yield_now().await;
|
| 75 |
+
}
|
| 76 |
+
});
|
| 77 |
+
|
| 78 |
+
assert_eq!(
|
| 79 |
+
FramedReader::new(reader)
|
| 80 |
+
.read::<serde_json::Value>()
|
| 81 |
+
.await
|
| 82 |
+
.expect("read frame"),
|
| 83 |
+
Some(value)
|
| 84 |
+
);
|
| 85 |
+
write.await.expect("writer task");
|
| 86 |
+
}
|
| 87 |
+
|
| 88 |
+
#[tokio::test]
|
| 89 |
+
async fn eof_is_clean_only_at_a_frame_boundary() {
|
| 90 |
+
let (writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16);
|
| 91 |
+
drop(writer);
|
| 92 |
+
assert_eq!(
|
| 93 |
+
FramedReader::new(reader)
|
| 94 |
+
.read::<serde_json::Value>()
|
| 95 |
+
.await
|
| 96 |
+
.expect("clean eof"),
|
| 97 |
+
None
|
| 98 |
+
);
|
| 99 |
+
|
| 100 |
+
let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16);
|
| 101 |
+
writer
|
| 102 |
+
.write_all(&[1, 0])
|
| 103 |
+
.await
|
| 104 |
+
.expect("write partial header");
|
| 105 |
+
drop(writer);
|
| 106 |
+
let err = FramedReader::new(reader)
|
| 107 |
+
.read::<serde_json::Value>()
|
| 108 |
+
.await
|
| 109 |
+
.expect_err("truncated header");
|
| 110 |
+
assert_eq!(err.kind(), std::io::ErrorKind::UnexpectedEof);
|
| 111 |
+
}
|
| 112 |
+
|
| 113 |
+
#[tokio::test]
|
| 114 |
+
async fn oversized_and_malformed_frames_are_rejected() {
|
| 115 |
+
let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16);
|
| 116 |
+
writer
|
| 117 |
+
.write_all(&((MAX_FRAME_BYTES as u32) + 1).to_le_bytes())
|
| 118 |
+
.await
|
| 119 |
+
.expect("write oversized header");
|
| 120 |
+
let err = FramedReader::new(reader)
|
| 121 |
+
.read::<serde_json::Value>()
|
| 122 |
+
.await
|
| 123 |
+
.expect_err("oversized frame");
|
| 124 |
+
assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
|
| 125 |
+
|
| 126 |
+
let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16);
|
| 127 |
+
writer
|
| 128 |
+
.write_all(&(1_u32).to_le_bytes())
|
| 129 |
+
.await
|
| 130 |
+
.expect("write length");
|
| 131 |
+
writer.write_all(b"{").await.expect("write malformed json");
|
| 132 |
+
let err = FramedReader::new(reader)
|
| 133 |
+
.read::<serde_json::Value>()
|
| 134 |
+
.await
|
| 135 |
+
.expect_err("malformed frame");
|
| 136 |
+
assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
|
| 137 |
+
}
|
codex-rs/code-mode-protocol/src/host/error.rs
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use serde::Deserialize;
|
| 2 |
+
use serde::Serialize;
|
| 3 |
+
|
| 4 |
+
use super::Capability;
|
| 5 |
+
use super::SupportedProtocolVersions;
|
| 6 |
+
|
| 7 |
+
/// Explains why connection negotiation was rejected before any session opened.
|
| 8 |
+
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
| 9 |
+
#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")]
|
| 10 |
+
pub enum HandshakeRejectReason {
|
| 11 |
+
#[serde(rename = "noCompatibleVersion")]
|
| 12 |
+
NoCompatibleVersion {
|
| 13 |
+
supported_versions: SupportedProtocolVersions,
|
| 14 |
+
},
|
| 15 |
+
#[serde(rename = "missingRequiredCapability")]
|
| 16 |
+
MissingRequiredCapability { capability: Capability },
|
| 17 |
+
#[serde(rename = "invalidHello")]
|
| 18 |
+
InvalidHello { message: String },
|
| 19 |
+
}
|
codex-rs/code-mode-protocol/src/host/host_tests.rs
ADDED
|
@@ -0,0 +1,847 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::fmt::Debug;
|
| 2 |
+
|
| 3 |
+
use pretty_assertions::assert_eq;
|
| 4 |
+
use serde::Serialize;
|
| 5 |
+
use serde::de::DeserializeOwned;
|
| 6 |
+
use serde_json::Value;
|
| 7 |
+
use serde_json::json;
|
| 8 |
+
|
| 9 |
+
use super::Capability;
|
| 10 |
+
use super::CapabilitySet;
|
| 11 |
+
use super::ClientHello;
|
| 12 |
+
use super::ClientToHost;
|
| 13 |
+
use super::DelegateRequest;
|
| 14 |
+
use super::DelegateRequestId;
|
| 15 |
+
use super::DelegateResponse;
|
| 16 |
+
use super::HandshakeRejectReason;
|
| 17 |
+
use super::HostHello;
|
| 18 |
+
use super::HostRequest;
|
| 19 |
+
use super::HostResponse;
|
| 20 |
+
use super::HostToClient;
|
| 21 |
+
use super::ProtocolVersion;
|
| 22 |
+
use super::RequestId;
|
| 23 |
+
use super::SessionId;
|
| 24 |
+
use super::SupportedProtocolVersions;
|
| 25 |
+
use super::WireCellId;
|
| 26 |
+
use super::WireContentItem;
|
| 27 |
+
use super::WireExecuteRequest;
|
| 28 |
+
use super::WireImageDetail;
|
| 29 |
+
use super::WireNestedToolCall;
|
| 30 |
+
use super::WireResult;
|
| 31 |
+
use super::WireRuntimeResponse;
|
| 32 |
+
use super::WireSessionCellExecutionLimits;
|
| 33 |
+
use super::WireToolDefinition;
|
| 34 |
+
use super::WireToolKind;
|
| 35 |
+
use super::WireToolName;
|
| 36 |
+
use super::WireWaitOutcome;
|
| 37 |
+
use super::WireWaitRequest;
|
| 38 |
+
use crate::CodeModeSessionCellExecutionLimits;
|
| 39 |
+
use crate::ExecuteRequest;
|
| 40 |
+
|
| 41 |
+
fn session_id() -> SessionId {
|
| 42 |
+
SessionId::new("session-1").expect("valid session ID")
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
fn cell_id(value: &str) -> WireCellId {
|
| 46 |
+
WireCellId::new(value)
|
| 47 |
+
}
|
| 48 |
+
|
| 49 |
+
fn request_id(value: i64) -> RequestId {
|
| 50 |
+
RequestId::new(value)
|
| 51 |
+
}
|
| 52 |
+
|
| 53 |
+
fn delegate_request_id(value: i64) -> DelegateRequestId {
|
| 54 |
+
DelegateRequestId::new(value)
|
| 55 |
+
}
|
| 56 |
+
|
| 57 |
+
fn capability(value: &str) -> Capability {
|
| 58 |
+
Capability::new(value).expect("valid capability")
|
| 59 |
+
}
|
| 60 |
+
|
| 61 |
+
fn supported_versions() -> SupportedProtocolVersions {
|
| 62 |
+
SupportedProtocolVersions::try_new([ProtocolVersion::V1])
|
| 63 |
+
.expect("nonempty unique protocol versions")
|
| 64 |
+
}
|
| 65 |
+
|
| 66 |
+
fn assert_wire_round_trip<T>(message: T, encoded: Value)
|
| 67 |
+
where
|
| 68 |
+
T: Debug + DeserializeOwned + PartialEq + Serialize,
|
| 69 |
+
{
|
| 70 |
+
assert_eq!(serde_json::to_value(&message).expect("serialize"), encoded);
|
| 71 |
+
assert_eq!(
|
| 72 |
+
serde_json::from_value::<T>(encoded).expect("deserialize"),
|
| 73 |
+
message
|
| 74 |
+
);
|
| 75 |
+
}
|
| 76 |
+
|
| 77 |
+
fn execute_request() -> WireExecuteRequest {
|
| 78 |
+
WireExecuteRequest {
|
| 79 |
+
tool_call_id: "call-1".to_string(),
|
| 80 |
+
enabled_tools: vec![
|
| 81 |
+
WireToolDefinition {
|
| 82 |
+
name: "function_tool".to_string(),
|
| 83 |
+
tool_name: WireToolName {
|
| 84 |
+
name: "function_tool".to_string(),
|
| 85 |
+
namespace: None,
|
| 86 |
+
},
|
| 87 |
+
description: "function tool".to_string(),
|
| 88 |
+
kind: WireToolKind::Function,
|
| 89 |
+
input_schema: Some(json!({ "type": "object" })),
|
| 90 |
+
output_schema: None,
|
| 91 |
+
},
|
| 92 |
+
WireToolDefinition {
|
| 93 |
+
name: "freeform_tool".to_string(),
|
| 94 |
+
tool_name: WireToolName {
|
| 95 |
+
name: "freeform_tool".to_string(),
|
| 96 |
+
namespace: Some("mcp__sample__".to_string()),
|
| 97 |
+
},
|
| 98 |
+
description: "freeform tool".to_string(),
|
| 99 |
+
kind: WireToolKind::Freeform,
|
| 100 |
+
input_schema: None,
|
| 101 |
+
output_schema: Some(json!({ "type": "string" })),
|
| 102 |
+
},
|
| 103 |
+
],
|
| 104 |
+
source: "text('hello');".to_string(),
|
| 105 |
+
yield_time_ms: Some(25),
|
| 106 |
+
max_output_tokens: Some(100),
|
| 107 |
+
}
|
| 108 |
+
}
|
| 109 |
+
|
| 110 |
+
fn content_items() -> Vec<WireContentItem> {
|
| 111 |
+
vec![
|
| 112 |
+
WireContentItem::InputText {
|
| 113 |
+
text: "hello".to_string(),
|
| 114 |
+
},
|
| 115 |
+
WireContentItem::InputImage {
|
| 116 |
+
image_url: "data:image/png;base64,none".to_string(),
|
| 117 |
+
detail: None,
|
| 118 |
+
},
|
| 119 |
+
WireContentItem::InputImage {
|
| 120 |
+
image_url: "data:image/png;base64,auto".to_string(),
|
| 121 |
+
detail: Some(WireImageDetail::Auto),
|
| 122 |
+
},
|
| 123 |
+
WireContentItem::InputImage {
|
| 124 |
+
image_url: "data:image/png;base64,low".to_string(),
|
| 125 |
+
detail: Some(WireImageDetail::Low),
|
| 126 |
+
},
|
| 127 |
+
WireContentItem::InputImage {
|
| 128 |
+
image_url: "data:image/png;base64,high".to_string(),
|
| 129 |
+
detail: Some(WireImageDetail::High),
|
| 130 |
+
},
|
| 131 |
+
WireContentItem::InputImage {
|
| 132 |
+
image_url: "data:image/png;base64,original".to_string(),
|
| 133 |
+
detail: Some(WireImageDetail::Original),
|
| 134 |
+
},
|
| 135 |
+
WireContentItem::InputAudio {
|
| 136 |
+
audio_url: "data:audio/wav;base64,YXVkaW8=".to_string(),
|
| 137 |
+
},
|
| 138 |
+
]
|
| 139 |
+
}
|
| 140 |
+
|
| 141 |
+
fn content_items_json() -> Value {
|
| 142 |
+
json!([
|
| 143 |
+
{ "type": "input_text", "text": "hello" },
|
| 144 |
+
{ "type": "input_image", "image_url": "data:image/png;base64,none" },
|
| 145 |
+
{
|
| 146 |
+
"type": "input_image",
|
| 147 |
+
"image_url": "data:image/png;base64,auto",
|
| 148 |
+
"detail": "auto",
|
| 149 |
+
},
|
| 150 |
+
{
|
| 151 |
+
"type": "input_image",
|
| 152 |
+
"image_url": "data:image/png;base64,low",
|
| 153 |
+
"detail": "low",
|
| 154 |
+
},
|
| 155 |
+
{
|
| 156 |
+
"type": "input_image",
|
| 157 |
+
"image_url": "data:image/png;base64,high",
|
| 158 |
+
"detail": "high",
|
| 159 |
+
},
|
| 160 |
+
{
|
| 161 |
+
"type": "input_image",
|
| 162 |
+
"image_url": "data:image/png;base64,original",
|
| 163 |
+
"detail": "original",
|
| 164 |
+
},
|
| 165 |
+
{
|
| 166 |
+
"type": "input_audio",
|
| 167 |
+
"audio_url": "data:audio/wav;base64,YXVkaW8=",
|
| 168 |
+
},
|
| 169 |
+
])
|
| 170 |
+
}
|
| 171 |
+
|
| 172 |
+
#[test]
|
| 173 |
+
fn handshake_v1_variants_are_pinned() {
|
| 174 |
+
assert_wire_round_trip(
|
| 175 |
+
ClientToHost::ClientHello(
|
| 176 |
+
ClientHello::new(
|
| 177 |
+
supported_versions(),
|
| 178 |
+
CapabilitySet::try_new([capability("required")]).expect("valid required set"),
|
| 179 |
+
CapabilitySet::try_new([capability("optional")]).expect("valid optional set"),
|
| 180 |
+
)
|
| 181 |
+
.expect("disjoint capabilities"),
|
| 182 |
+
),
|
| 183 |
+
json!({
|
| 184 |
+
"type": "connection/hello",
|
| 185 |
+
"supportedVersions": [1],
|
| 186 |
+
"requiredCapabilities": ["required"],
|
| 187 |
+
"optionalCapabilities": ["optional"],
|
| 188 |
+
}),
|
| 189 |
+
);
|
| 190 |
+
assert_wire_round_trip(
|
| 191 |
+
HostToClient::HostHello(HostHello::new(
|
| 192 |
+
ProtocolVersion::V1,
|
| 193 |
+
CapabilitySet::try_new([capability("required")]).expect("valid capabilities"),
|
| 194 |
+
)),
|
| 195 |
+
json!({
|
| 196 |
+
"type": "connection/ready",
|
| 197 |
+
"selectedVersion": 1,
|
| 198 |
+
"capabilities": ["required"],
|
| 199 |
+
}),
|
| 200 |
+
);
|
| 201 |
+
for (reason, encoded) in [
|
| 202 |
+
(
|
| 203 |
+
HandshakeRejectReason::NoCompatibleVersion {
|
| 204 |
+
supported_versions: supported_versions(),
|
| 205 |
+
},
|
| 206 |
+
json!({
|
| 207 |
+
"type": "connection/rejected",
|
| 208 |
+
"reason": {
|
| 209 |
+
"type": "noCompatibleVersion",
|
| 210 |
+
"supportedVersions": [1],
|
| 211 |
+
},
|
| 212 |
+
}),
|
| 213 |
+
),
|
| 214 |
+
(
|
| 215 |
+
HandshakeRejectReason::MissingRequiredCapability {
|
| 216 |
+
capability: capability("required"),
|
| 217 |
+
},
|
| 218 |
+
json!({
|
| 219 |
+
"type": "connection/rejected",
|
| 220 |
+
"reason": {
|
| 221 |
+
"type": "missingRequiredCapability",
|
| 222 |
+
"capability": "required",
|
| 223 |
+
},
|
| 224 |
+
}),
|
| 225 |
+
),
|
| 226 |
+
(
|
| 227 |
+
HandshakeRejectReason::InvalidHello {
|
| 228 |
+
message: "invalid hello".to_string(),
|
| 229 |
+
},
|
| 230 |
+
json!({
|
| 231 |
+
"type": "connection/rejected",
|
| 232 |
+
"reason": {
|
| 233 |
+
"type": "invalidHello",
|
| 234 |
+
"message": "invalid hello",
|
| 235 |
+
},
|
| 236 |
+
}),
|
| 237 |
+
),
|
| 238 |
+
] {
|
| 239 |
+
assert_wire_round_trip(HostToClient::HandshakeRejected { reason }, encoded);
|
| 240 |
+
}
|
| 241 |
+
}
|
| 242 |
+
|
| 243 |
+
#[test]
|
| 244 |
+
fn open_session_serializes_optional_cell_execution_limits() {
|
| 245 |
+
assert_wire_round_trip(
|
| 246 |
+
HostRequest::OpenSession {
|
| 247 |
+
session_id: session_id(),
|
| 248 |
+
cell_execution_limits: Some(WireSessionCellExecutionLimits {
|
| 249 |
+
max_yield_time_ms: Some(250),
|
| 250 |
+
max_heap_size_bytes: Some(16 * 1024 * 1024),
|
| 251 |
+
}),
|
| 252 |
+
},
|
| 253 |
+
json!({
|
| 254 |
+
"method": "session/open",
|
| 255 |
+
"sessionId": "session-1",
|
| 256 |
+
"cellExecutionLimits": {
|
| 257 |
+
"maxYieldTimeMs": 250,
|
| 258 |
+
"maxHeapSizeBytes": 16 * 1024 * 1024,
|
| 259 |
+
},
|
| 260 |
+
}),
|
| 261 |
+
);
|
| 262 |
+
}
|
| 263 |
+
|
| 264 |
+
#[test]
|
| 265 |
+
fn session_cell_execution_limits_convert_between_domain_and_wire() {
|
| 266 |
+
let domain_limits = CodeModeSessionCellExecutionLimits {
|
| 267 |
+
max_yield_time_ms: Some(250),
|
| 268 |
+
max_heap_size_bytes: Some(16_usize * 1024 * 1024),
|
| 269 |
+
};
|
| 270 |
+
let wire_limits = WireSessionCellExecutionLimits {
|
| 271 |
+
max_yield_time_ms: Some(250),
|
| 272 |
+
max_heap_size_bytes: Some(16_u64 * 1024 * 1024),
|
| 273 |
+
};
|
| 274 |
+
|
| 275 |
+
assert_eq!(
|
| 276 |
+
WireSessionCellExecutionLimits::try_from(domain_limits.clone())
|
| 277 |
+
.expect("domain limits convert to wire limits"),
|
| 278 |
+
wire_limits
|
| 279 |
+
);
|
| 280 |
+
assert_eq!(
|
| 281 |
+
CodeModeSessionCellExecutionLimits::try_from(wire_limits)
|
| 282 |
+
.expect("wire limits convert to domain limits"),
|
| 283 |
+
domain_limits
|
| 284 |
+
);
|
| 285 |
+
}
|
| 286 |
+
|
| 287 |
+
#[cfg(target_pointer_width = "32")]
|
| 288 |
+
#[test]
|
| 289 |
+
fn session_cell_execution_limits_reject_heap_sizes_that_exceed_usize() {
|
| 290 |
+
let wire_limits = WireSessionCellExecutionLimits {
|
| 291 |
+
max_yield_time_ms: None,
|
| 292 |
+
max_heap_size_bytes: Some(u64::from(u32::MAX) + 1),
|
| 293 |
+
};
|
| 294 |
+
|
| 295 |
+
assert!(CodeModeSessionCellExecutionLimits::try_from(wire_limits).is_err());
|
| 296 |
+
}
|
| 297 |
+
|
| 298 |
+
#[test]
|
| 299 |
+
fn client_to_host_v1_variants_are_pinned() {
|
| 300 |
+
let execute_request = execute_request();
|
| 301 |
+
for (id, request, encoded_request) in [
|
| 302 |
+
(
|
| 303 |
+
request_id(/*value*/ 1),
|
| 304 |
+
HostRequest::OpenSession {
|
| 305 |
+
session_id: session_id(),
|
| 306 |
+
cell_execution_limits: None,
|
| 307 |
+
},
|
| 308 |
+
json!({ "method": "session/open", "sessionId": "session-1" }),
|
| 309 |
+
),
|
| 310 |
+
(
|
| 311 |
+
request_id(/*value*/ 2),
|
| 312 |
+
HostRequest::Execute {
|
| 313 |
+
session_id: session_id(),
|
| 314 |
+
request: execute_request,
|
| 315 |
+
},
|
| 316 |
+
json!({
|
| 317 |
+
"method": "session/execute",
|
| 318 |
+
"sessionId": "session-1",
|
| 319 |
+
"request": {
|
| 320 |
+
"tool_call_id": "call-1",
|
| 321 |
+
"enabled_tools": [
|
| 322 |
+
{
|
| 323 |
+
"name": "function_tool",
|
| 324 |
+
"tool_name": { "name": "function_tool", "namespace": null },
|
| 325 |
+
"description": "function tool",
|
| 326 |
+
"kind": "function",
|
| 327 |
+
"input_schema": { "type": "object" },
|
| 328 |
+
"output_schema": null,
|
| 329 |
+
},
|
| 330 |
+
{
|
| 331 |
+
"name": "freeform_tool",
|
| 332 |
+
"tool_name": {
|
| 333 |
+
"name": "freeform_tool",
|
| 334 |
+
"namespace": "mcp__sample__",
|
| 335 |
+
},
|
| 336 |
+
"description": "freeform tool",
|
| 337 |
+
"kind": "freeform",
|
| 338 |
+
"input_schema": null,
|
| 339 |
+
"output_schema": { "type": "string" },
|
| 340 |
+
},
|
| 341 |
+
],
|
| 342 |
+
"source": "text('hello');",
|
| 343 |
+
"yield_time_ms": 25,
|
| 344 |
+
"max_output_tokens": 100,
|
| 345 |
+
},
|
| 346 |
+
}),
|
| 347 |
+
),
|
| 348 |
+
(
|
| 349 |
+
request_id(/*value*/ 3),
|
| 350 |
+
HostRequest::Wait {
|
| 351 |
+
session_id: session_id(),
|
| 352 |
+
request: WireWaitRequest {
|
| 353 |
+
cell_id: cell_id("cell-1"),
|
| 354 |
+
yield_time_ms: 50,
|
| 355 |
+
},
|
| 356 |
+
},
|
| 357 |
+
json!({
|
| 358 |
+
"method": "session/wait",
|
| 359 |
+
"sessionId": "session-1",
|
| 360 |
+
"request": { "cell_id": "cell-1", "yield_time_ms": 50 },
|
| 361 |
+
}),
|
| 362 |
+
),
|
| 363 |
+
(
|
| 364 |
+
request_id(/*value*/ 4),
|
| 365 |
+
HostRequest::Terminate {
|
| 366 |
+
session_id: session_id(),
|
| 367 |
+
cell_id: cell_id("cell-1"),
|
| 368 |
+
},
|
| 369 |
+
json!({
|
| 370 |
+
"method": "session/terminate",
|
| 371 |
+
"sessionId": "session-1",
|
| 372 |
+
"cellId": "cell-1",
|
| 373 |
+
}),
|
| 374 |
+
),
|
| 375 |
+
(
|
| 376 |
+
request_id(/*value*/ 5),
|
| 377 |
+
HostRequest::ShutdownSession {
|
| 378 |
+
session_id: session_id(),
|
| 379 |
+
},
|
| 380 |
+
json!({ "method": "session/shutdown", "sessionId": "session-1" }),
|
| 381 |
+
),
|
| 382 |
+
] {
|
| 383 |
+
assert_wire_round_trip(
|
| 384 |
+
ClientToHost::Request { id, request },
|
| 385 |
+
json!({
|
| 386 |
+
"type": "operation/request",
|
| 387 |
+
"id": id,
|
| 388 |
+
"request": encoded_request,
|
| 389 |
+
}),
|
| 390 |
+
);
|
| 391 |
+
}
|
| 392 |
+
|
| 393 |
+
for (id, result, encoded_result) in [
|
| 394 |
+
(
|
| 395 |
+
delegate_request_id(/*value*/ 6),
|
| 396 |
+
WireResult::Ok {
|
| 397 |
+
value: DelegateResponse::ToolResult {
|
| 398 |
+
result: json!({ "answer": 42 }),
|
| 399 |
+
},
|
| 400 |
+
},
|
| 401 |
+
json!({
|
| 402 |
+
"status": "ok",
|
| 403 |
+
"value": { "type": "tool/result", "result": { "answer": 42 } },
|
| 404 |
+
}),
|
| 405 |
+
),
|
| 406 |
+
(
|
| 407 |
+
delegate_request_id(/*value*/ 7),
|
| 408 |
+
WireResult::Ok {
|
| 409 |
+
value: DelegateResponse::NotificationDelivered,
|
| 410 |
+
},
|
| 411 |
+
json!({
|
| 412 |
+
"status": "ok",
|
| 413 |
+
"value": { "type": "notification/delivered" },
|
| 414 |
+
}),
|
| 415 |
+
),
|
| 416 |
+
(
|
| 417 |
+
delegate_request_id(/*value*/ 8),
|
| 418 |
+
WireResult::Err {
|
| 419 |
+
message: "delegate failed".to_string(),
|
| 420 |
+
},
|
| 421 |
+
json!({ "status": "error", "message": "delegate failed" }),
|
| 422 |
+
),
|
| 423 |
+
] {
|
| 424 |
+
assert_wire_round_trip(
|
| 425 |
+
ClientToHost::DelegateResponse { id, result },
|
| 426 |
+
json!({
|
| 427 |
+
"type": "delegate/response",
|
| 428 |
+
"id": id,
|
| 429 |
+
"result": encoded_result,
|
| 430 |
+
}),
|
| 431 |
+
);
|
| 432 |
+
}
|
| 433 |
+
|
| 434 |
+
assert_wire_round_trip(
|
| 435 |
+
ClientToHost::CancelRequest {
|
| 436 |
+
id: request_id(/*value*/ 9),
|
| 437 |
+
},
|
| 438 |
+
json!({
|
| 439 |
+
"type": "operation/cancel",
|
| 440 |
+
"id": 9,
|
| 441 |
+
}),
|
| 442 |
+
);
|
| 443 |
+
}
|
| 444 |
+
|
| 445 |
+
#[test]
|
| 446 |
+
fn host_to_client_v1_variants_are_pinned() {
|
| 447 |
+
for (id, response, encoded_response) in [
|
| 448 |
+
(
|
| 449 |
+
request_id(/*value*/ 1),
|
| 450 |
+
HostResponse::SessionReady {
|
| 451 |
+
session_id: session_id(),
|
| 452 |
+
},
|
| 453 |
+
json!({ "type": "session/ready", "sessionId": "session-1" }),
|
| 454 |
+
),
|
| 455 |
+
(
|
| 456 |
+
request_id(/*value*/ 2),
|
| 457 |
+
HostResponse::ExecutionStarted {
|
| 458 |
+
cell_id: cell_id("cell-1"),
|
| 459 |
+
},
|
| 460 |
+
json!({ "type": "execution/started", "cellId": "cell-1" }),
|
| 461 |
+
),
|
| 462 |
+
(
|
| 463 |
+
request_id(/*value*/ 3),
|
| 464 |
+
HostResponse::WaitCompleted {
|
| 465 |
+
outcome: WireWaitOutcome::LiveCell(WireRuntimeResponse::Yielded {
|
| 466 |
+
code_mode_host_duration_ns: 0,
|
| 467 |
+
cell_id: cell_id("cell-1"),
|
| 468 |
+
content_items: content_items(),
|
| 469 |
+
}),
|
| 470 |
+
},
|
| 471 |
+
json!({
|
| 472 |
+
"type": "wait/completed",
|
| 473 |
+
"outcome": {
|
| 474 |
+
"LiveCell": {
|
| 475 |
+
"Yielded": {
|
| 476 |
+
"cell_id": "cell-1",
|
| 477 |
+
"content_items": content_items_json(),
|
| 478 |
+
"code_mode_host_duration_ns": 0,
|
| 479 |
+
},
|
| 480 |
+
},
|
| 481 |
+
},
|
| 482 |
+
}),
|
| 483 |
+
),
|
| 484 |
+
(
|
| 485 |
+
request_id(/*value*/ 4),
|
| 486 |
+
HostResponse::WaitCompleted {
|
| 487 |
+
outcome: WireWaitOutcome::MissingCell(WireRuntimeResponse::Result {
|
| 488 |
+
code_mode_host_duration_ns: 0,
|
| 489 |
+
cell_id: cell_id("missing-cell"),
|
| 490 |
+
content_items: Vec::new(),
|
| 491 |
+
error_text: Some("cell not found".to_string()),
|
| 492 |
+
}),
|
| 493 |
+
},
|
| 494 |
+
json!({
|
| 495 |
+
"type": "wait/completed",
|
| 496 |
+
"outcome": {
|
| 497 |
+
"MissingCell": {
|
| 498 |
+
"Result": {
|
| 499 |
+
"cell_id": "missing-cell",
|
| 500 |
+
"content_items": [],
|
| 501 |
+
"error_text": "cell not found",
|
| 502 |
+
"code_mode_host_duration_ns": 0,
|
| 503 |
+
},
|
| 504 |
+
},
|
| 505 |
+
},
|
| 506 |
+
}),
|
| 507 |
+
),
|
| 508 |
+
(
|
| 509 |
+
request_id(/*value*/ 5),
|
| 510 |
+
HostResponse::SessionClosed {
|
| 511 |
+
session_id: session_id(),
|
| 512 |
+
},
|
| 513 |
+
json!({ "type": "session/closed", "sessionId": "session-1" }),
|
| 514 |
+
),
|
| 515 |
+
] {
|
| 516 |
+
assert_wire_round_trip(
|
| 517 |
+
HostToClient::Response {
|
| 518 |
+
id,
|
| 519 |
+
result: WireResult::Ok { value: response },
|
| 520 |
+
},
|
| 521 |
+
json!({
|
| 522 |
+
"type": "operation/response",
|
| 523 |
+
"id": id,
|
| 524 |
+
"result": { "status": "ok", "value": encoded_response },
|
| 525 |
+
}),
|
| 526 |
+
);
|
| 527 |
+
}
|
| 528 |
+
assert_wire_round_trip(
|
| 529 |
+
HostToClient::Response {
|
| 530 |
+
id: request_id(/*value*/ 6),
|
| 531 |
+
result: WireResult::Err {
|
| 532 |
+
message: "operation failed".to_string(),
|
| 533 |
+
},
|
| 534 |
+
},
|
| 535 |
+
json!({
|
| 536 |
+
"type": "operation/response",
|
| 537 |
+
"id": 6,
|
| 538 |
+
"result": { "status": "error", "message": "operation failed" },
|
| 539 |
+
}),
|
| 540 |
+
);
|
| 541 |
+
|
| 542 |
+
assert_wire_round_trip(
|
| 543 |
+
HostToClient::InitialResponse {
|
| 544 |
+
id: request_id(/*value*/ 7),
|
| 545 |
+
result: WireResult::Ok {
|
| 546 |
+
value: WireRuntimeResponse::Terminated {
|
| 547 |
+
code_mode_host_duration_ns: 0,
|
| 548 |
+
cell_id: cell_id("cell-1"),
|
| 549 |
+
content_items: Vec::new(),
|
| 550 |
+
},
|
| 551 |
+
},
|
| 552 |
+
},
|
| 553 |
+
json!({
|
| 554 |
+
"type": "execute/initialResponse",
|
| 555 |
+
"id": 7,
|
| 556 |
+
"result": {
|
| 557 |
+
"status": "ok",
|
| 558 |
+
"value": {
|
| 559 |
+
"Terminated": {
|
| 560 |
+
"cell_id": "cell-1",
|
| 561 |
+
"content_items": [],
|
| 562 |
+
"code_mode_host_duration_ns": 0,
|
| 563 |
+
},
|
| 564 |
+
},
|
| 565 |
+
},
|
| 566 |
+
}),
|
| 567 |
+
);
|
| 568 |
+
assert_wire_round_trip(
|
| 569 |
+
HostToClient::InitialResponse {
|
| 570 |
+
id: request_id(/*value*/ 8),
|
| 571 |
+
result: WireResult::Err {
|
| 572 |
+
message: "execution failed".to_string(),
|
| 573 |
+
},
|
| 574 |
+
},
|
| 575 |
+
json!({
|
| 576 |
+
"type": "execute/initialResponse",
|
| 577 |
+
"id": 8,
|
| 578 |
+
"result": { "status": "error", "message": "execution failed" },
|
| 579 |
+
}),
|
| 580 |
+
);
|
| 581 |
+
|
| 582 |
+
assert_wire_round_trip(
|
| 583 |
+
HostToClient::DelegateRequest {
|
| 584 |
+
id: delegate_request_id(/*value*/ 9),
|
| 585 |
+
session_id: session_id(),
|
| 586 |
+
request: DelegateRequest::InvokeTool {
|
| 587 |
+
invocation: WireNestedToolCall {
|
| 588 |
+
cell_id: cell_id("cell-1"),
|
| 589 |
+
runtime_tool_call_id: "runtime-call-1".to_string(),
|
| 590 |
+
tool_name: WireToolName {
|
| 591 |
+
name: "freeform_tool".to_string(),
|
| 592 |
+
namespace: Some("mcp__sample__".to_string()),
|
| 593 |
+
},
|
| 594 |
+
tool_kind: WireToolKind::Freeform,
|
| 595 |
+
input: Some(json!({ "value": 1 })),
|
| 596 |
+
},
|
| 597 |
+
},
|
| 598 |
+
},
|
| 599 |
+
json!({
|
| 600 |
+
"type": "delegate/request",
|
| 601 |
+
"id": 9,
|
| 602 |
+
"sessionId": "session-1",
|
| 603 |
+
"request": {
|
| 604 |
+
"type": "tool/invoke",
|
| 605 |
+
"invocation": {
|
| 606 |
+
"cell_id": "cell-1",
|
| 607 |
+
"runtime_tool_call_id": "runtime-call-1",
|
| 608 |
+
"tool_name": {
|
| 609 |
+
"name": "freeform_tool",
|
| 610 |
+
"namespace": "mcp__sample__",
|
| 611 |
+
},
|
| 612 |
+
"tool_kind": "freeform",
|
| 613 |
+
"input": { "value": 1 },
|
| 614 |
+
},
|
| 615 |
+
},
|
| 616 |
+
}),
|
| 617 |
+
);
|
| 618 |
+
assert_wire_round_trip(
|
| 619 |
+
HostToClient::DelegateRequest {
|
| 620 |
+
id: delegate_request_id(/*value*/ 10),
|
| 621 |
+
session_id: session_id(),
|
| 622 |
+
request: DelegateRequest::Notify {
|
| 623 |
+
call_id: "call-1".to_string(),
|
| 624 |
+
cell_id: cell_id("cell-1"),
|
| 625 |
+
text: "important".to_string(),
|
| 626 |
+
},
|
| 627 |
+
},
|
| 628 |
+
json!({
|
| 629 |
+
"type": "delegate/request",
|
| 630 |
+
"id": 10,
|
| 631 |
+
"sessionId": "session-1",
|
| 632 |
+
"request": {
|
| 633 |
+
"type": "notification/send",
|
| 634 |
+
"callId": "call-1",
|
| 635 |
+
"cellId": "cell-1",
|
| 636 |
+
"text": "important",
|
| 637 |
+
},
|
| 638 |
+
}),
|
| 639 |
+
);
|
| 640 |
+
assert_wire_round_trip(
|
| 641 |
+
HostToClient::CancelDelegateRequest {
|
| 642 |
+
id: delegate_request_id(/*value*/ 11),
|
| 643 |
+
},
|
| 644 |
+
json!({ "type": "delegate/cancel", "id": 11 }),
|
| 645 |
+
);
|
| 646 |
+
assert_wire_round_trip(
|
| 647 |
+
HostToClient::CellClosed {
|
| 648 |
+
session_id: session_id(),
|
| 649 |
+
cell_id: cell_id("cell-1"),
|
| 650 |
+
},
|
| 651 |
+
json!({
|
| 652 |
+
"type": "cell/closed",
|
| 653 |
+
"sessionId": "session-1",
|
| 654 |
+
"cellId": "cell-1",
|
| 655 |
+
}),
|
| 656 |
+
);
|
| 657 |
+
}
|
| 658 |
+
|
| 659 |
+
#[test]
|
| 660 |
+
fn execute_request_integer_bounds_are_enforced() {
|
| 661 |
+
let wire_request = execute_request();
|
| 662 |
+
let domain_request = ExecuteRequest::try_from(wire_request.clone())
|
| 663 |
+
.expect("valid wire request converts to the domain");
|
| 664 |
+
assert_eq!(
|
| 665 |
+
WireExecuteRequest::try_from(domain_request.clone())
|
| 666 |
+
.expect("valid domain request converts to the wire"),
|
| 667 |
+
wire_request
|
| 668 |
+
);
|
| 669 |
+
|
| 670 |
+
let too_large = ExecuteRequest {
|
| 671 |
+
max_output_tokens: Some(usize::try_from(i32::MAX).expect("i32::MAX fits usize") + 1),
|
| 672 |
+
..domain_request
|
| 673 |
+
};
|
| 674 |
+
assert!(WireExecuteRequest::try_from(too_large).is_err());
|
| 675 |
+
|
| 676 |
+
let negative = WireExecuteRequest {
|
| 677 |
+
max_output_tokens: Some(-1),
|
| 678 |
+
..wire_request
|
| 679 |
+
};
|
| 680 |
+
assert!(ExecuteRequest::try_from(negative).is_err());
|
| 681 |
+
}
|
| 682 |
+
|
| 683 |
+
#[test]
|
| 684 |
+
fn invalid_protocol_states_cannot_be_constructed_or_decoded() {
|
| 685 |
+
assert!(SessionId::new("").is_err());
|
| 686 |
+
assert!(Capability::new(" ").is_err());
|
| 687 |
+
assert!(ProtocolVersion::new(/*value*/ 0).is_none());
|
| 688 |
+
assert!(SupportedProtocolVersions::try_new([]).is_err());
|
| 689 |
+
assert!(
|
| 690 |
+
SupportedProtocolVersions::try_new([ProtocolVersion::V1, ProtocolVersion::V1]).is_err()
|
| 691 |
+
);
|
| 692 |
+
assert!(CapabilitySet::try_new([capability("same"), capability("same")]).is_err());
|
| 693 |
+
|
| 694 |
+
let version_two = ProtocolVersion::new(/*value*/ 2).expect("valid protocol version");
|
| 695 |
+
let versions = SupportedProtocolVersions::try_new([ProtocolVersion::V1, version_two])
|
| 696 |
+
.expect("valid versions");
|
| 697 |
+
assert!(versions.contains(ProtocolVersion::V1));
|
| 698 |
+
assert_eq!(
|
| 699 |
+
versions.iter().collect::<Vec<_>>(),
|
| 700 |
+
vec![ProtocolVersion::V1, version_two]
|
| 701 |
+
);
|
| 702 |
+
|
| 703 |
+
let overlapping = capability("overlapping");
|
| 704 |
+
assert!(
|
| 705 |
+
ClientHello::new(
|
| 706 |
+
supported_versions(),
|
| 707 |
+
CapabilitySet::try_new([overlapping.clone()]).expect("valid required set"),
|
| 708 |
+
CapabilitySet::try_new([overlapping]).expect("valid optional set"),
|
| 709 |
+
)
|
| 710 |
+
.is_err()
|
| 711 |
+
);
|
| 712 |
+
|
| 713 |
+
for invalid in [
|
| 714 |
+
json!({
|
| 715 |
+
"type": "operation/request",
|
| 716 |
+
"id": 1,
|
| 717 |
+
"request": { "method": "session/open", "sessionId": "" },
|
| 718 |
+
}),
|
| 719 |
+
json!({
|
| 720 |
+
"type": "connection/hello",
|
| 721 |
+
"supportedVersions": [],
|
| 722 |
+
"requiredCapabilities": [],
|
| 723 |
+
"optionalCapabilities": [],
|
| 724 |
+
}),
|
| 725 |
+
json!({
|
| 726 |
+
"type": "connection/hello",
|
| 727 |
+
"supportedVersions": [1],
|
| 728 |
+
"requiredCapabilities": ["overlapping"],
|
| 729 |
+
"optionalCapabilities": ["overlapping"],
|
| 730 |
+
}),
|
| 731 |
+
] {
|
| 732 |
+
assert!(serde_json::from_value::<ClientToHost>(invalid).is_err());
|
| 733 |
+
}
|
| 734 |
+
}
|
| 735 |
+
|
| 736 |
+
#[test]
|
| 737 |
+
fn every_nested_v1_object_rejects_unknown_fields() {
|
| 738 |
+
assert!(
|
| 739 |
+
serde_json::from_value::<ClientToHost>(json!({
|
| 740 |
+
"type": "operation/request",
|
| 741 |
+
"id": 1,
|
| 742 |
+
"request": { "method": "session/open", "sessionId": "session-1" },
|
| 743 |
+
"unexpected": true,
|
| 744 |
+
}))
|
| 745 |
+
.is_err()
|
| 746 |
+
);
|
| 747 |
+
assert!(
|
| 748 |
+
serde_json::from_value::<HostRequest>(json!({
|
| 749 |
+
"method": "session/open",
|
| 750 |
+
"sessionId": "session-1",
|
| 751 |
+
"unexpected": true,
|
| 752 |
+
}))
|
| 753 |
+
.is_err()
|
| 754 |
+
);
|
| 755 |
+
assert!(
|
| 756 |
+
serde_json::from_value::<HostRequest>(json!({
|
| 757 |
+
"method": "session/open",
|
| 758 |
+
"sessionId": "session-1",
|
| 759 |
+
"cellExecutionLimits": {
|
| 760 |
+
"maxYieldTimeMs": 250,
|
| 761 |
+
"unexpected": true,
|
| 762 |
+
},
|
| 763 |
+
}))
|
| 764 |
+
.is_err()
|
| 765 |
+
);
|
| 766 |
+
assert!(
|
| 767 |
+
serde_json::from_value::<WireExecuteRequest>(json!({
|
| 768 |
+
"tool_call_id": "call-1",
|
| 769 |
+
"enabled_tools": [],
|
| 770 |
+
"source": "text('hello');",
|
| 771 |
+
"yield_time_ms": null,
|
| 772 |
+
"max_output_tokens": null,
|
| 773 |
+
"unexpected": true,
|
| 774 |
+
}))
|
| 775 |
+
.is_err()
|
| 776 |
+
);
|
| 777 |
+
assert!(
|
| 778 |
+
serde_json::from_value::<WireToolDefinition>(json!({
|
| 779 |
+
"name": "tool",
|
| 780 |
+
"tool_name": { "name": "tool", "namespace": null },
|
| 781 |
+
"description": "tool",
|
| 782 |
+
"kind": "function",
|
| 783 |
+
"input_schema": null,
|
| 784 |
+
"output_schema": null,
|
| 785 |
+
"unexpected": true,
|
| 786 |
+
}))
|
| 787 |
+
.is_err()
|
| 788 |
+
);
|
| 789 |
+
assert!(
|
| 790 |
+
serde_json::from_value::<WireToolName>(json!({
|
| 791 |
+
"name": "tool",
|
| 792 |
+
"namespace": null,
|
| 793 |
+
"unexpected": true,
|
| 794 |
+
}))
|
| 795 |
+
.is_err()
|
| 796 |
+
);
|
| 797 |
+
assert!(
|
| 798 |
+
serde_json::from_value::<WireWaitRequest>(json!({
|
| 799 |
+
"cell_id": "cell-1",
|
| 800 |
+
"yield_time_ms": 50,
|
| 801 |
+
"unexpected": true,
|
| 802 |
+
}))
|
| 803 |
+
.is_err()
|
| 804 |
+
);
|
| 805 |
+
assert!(
|
| 806 |
+
serde_json::from_value::<WireRuntimeResponse>(json!({
|
| 807 |
+
"Yielded": {
|
| 808 |
+
"cell_id": "cell-1",
|
| 809 |
+
"content_items": [],
|
| 810 |
+
"code_mode_host_duration_ns": 0,
|
| 811 |
+
"unexpected": true,
|
| 812 |
+
},
|
| 813 |
+
}))
|
| 814 |
+
.is_err()
|
| 815 |
+
);
|
| 816 |
+
assert!(
|
| 817 |
+
serde_json::from_value::<WireContentItem>(json!({
|
| 818 |
+
"type": "input_text",
|
| 819 |
+
"text": "hello",
|
| 820 |
+
"unexpected": true,
|
| 821 |
+
}))
|
| 822 |
+
.is_err()
|
| 823 |
+
);
|
| 824 |
+
assert!(
|
| 825 |
+
serde_json::from_value::<WireNestedToolCall>(json!({
|
| 826 |
+
"cell_id": "cell-1",
|
| 827 |
+
"runtime_tool_call_id": "runtime-call-1",
|
| 828 |
+
"tool_name": { "name": "tool", "namespace": null },
|
| 829 |
+
"tool_kind": "function",
|
| 830 |
+
"input": null,
|
| 831 |
+
"unexpected": true,
|
| 832 |
+
}))
|
| 833 |
+
.is_err()
|
| 834 |
+
);
|
| 835 |
+
assert!(
|
| 836 |
+
serde_json::from_value::<HostToClient>(json!({
|
| 837 |
+
"type": "operation/response",
|
| 838 |
+
"id": 1,
|
| 839 |
+
"result": {
|
| 840 |
+
"status": "ok",
|
| 841 |
+
"value": { "type": "session/ready", "sessionId": "session-1" },
|
| 842 |
+
},
|
| 843 |
+
"unexpected": true,
|
| 844 |
+
}))
|
| 845 |
+
.is_err()
|
| 846 |
+
);
|
| 847 |
+
}
|
codex-rs/code-mode-protocol/src/host/message.rs
ADDED
|
@@ -0,0 +1,264 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::fmt;
|
| 2 |
+
|
| 3 |
+
use serde::Deserialize;
|
| 4 |
+
use serde::Serialize;
|
| 5 |
+
use serde_json::Value as JsonValue;
|
| 6 |
+
|
| 7 |
+
use super::Capability;
|
| 8 |
+
use super::CapabilitySet;
|
| 9 |
+
use super::DelegateRequestId;
|
| 10 |
+
use super::HandshakeRejectReason;
|
| 11 |
+
use super::ProtocolVersion;
|
| 12 |
+
use super::RequestId;
|
| 13 |
+
use super::SessionId;
|
| 14 |
+
use super::SupportedProtocolVersions;
|
| 15 |
+
use super::WireCellId;
|
| 16 |
+
use super::WireExecuteRequest;
|
| 17 |
+
use super::WireNestedToolCall;
|
| 18 |
+
use super::WireRuntimeResponse;
|
| 19 |
+
use super::WireSessionCellExecutionLimits;
|
| 20 |
+
use super::WireWaitOutcome;
|
| 21 |
+
use super::WireWaitRequest;
|
| 22 |
+
|
| 23 |
+
#[derive(Clone, Debug, PartialEq, Serialize)]
|
| 24 |
+
#[serde(rename_all = "camelCase")]
|
| 25 |
+
pub struct ClientHello {
|
| 26 |
+
supported_versions: SupportedProtocolVersions,
|
| 27 |
+
required_capabilities: CapabilitySet,
|
| 28 |
+
optional_capabilities: CapabilitySet,
|
| 29 |
+
}
|
| 30 |
+
|
| 31 |
+
impl ClientHello {
|
| 32 |
+
pub fn new(
|
| 33 |
+
supported_versions: SupportedProtocolVersions,
|
| 34 |
+
required_capabilities: CapabilitySet,
|
| 35 |
+
optional_capabilities: CapabilitySet,
|
| 36 |
+
) -> Result<Self, ClientHelloError> {
|
| 37 |
+
if let Some(capability) = required_capabilities
|
| 38 |
+
.iter()
|
| 39 |
+
.find(|capability| optional_capabilities.contains(capability))
|
| 40 |
+
{
|
| 41 |
+
return Err(ClientHelloError::OverlappingCapability(capability.clone()));
|
| 42 |
+
}
|
| 43 |
+
Ok(Self {
|
| 44 |
+
supported_versions,
|
| 45 |
+
required_capabilities,
|
| 46 |
+
optional_capabilities,
|
| 47 |
+
})
|
| 48 |
+
}
|
| 49 |
+
|
| 50 |
+
pub fn supported_versions(&self) -> &SupportedProtocolVersions {
|
| 51 |
+
&self.supported_versions
|
| 52 |
+
}
|
| 53 |
+
|
| 54 |
+
pub fn required_capabilities(&self) -> &CapabilitySet {
|
| 55 |
+
&self.required_capabilities
|
| 56 |
+
}
|
| 57 |
+
|
| 58 |
+
pub fn optional_capabilities(&self) -> &CapabilitySet {
|
| 59 |
+
&self.optional_capabilities
|
| 60 |
+
}
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
+
#[derive(Deserialize)]
|
| 64 |
+
#[serde(deny_unknown_fields, rename_all = "camelCase")]
|
| 65 |
+
struct ClientHelloWire {
|
| 66 |
+
supported_versions: SupportedProtocolVersions,
|
| 67 |
+
required_capabilities: CapabilitySet,
|
| 68 |
+
optional_capabilities: CapabilitySet,
|
| 69 |
+
}
|
| 70 |
+
|
| 71 |
+
impl<'de> Deserialize<'de> for ClientHello {
|
| 72 |
+
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
| 73 |
+
where
|
| 74 |
+
D: serde::Deserializer<'de>,
|
| 75 |
+
{
|
| 76 |
+
let wire = ClientHelloWire::deserialize(deserializer)?;
|
| 77 |
+
Self::new(
|
| 78 |
+
wire.supported_versions,
|
| 79 |
+
wire.required_capabilities,
|
| 80 |
+
wire.optional_capabilities,
|
| 81 |
+
)
|
| 82 |
+
.map_err(serde::de::Error::custom)
|
| 83 |
+
}
|
| 84 |
+
}
|
| 85 |
+
|
| 86 |
+
#[derive(Clone, Debug, Eq, PartialEq)]
|
| 87 |
+
pub enum ClientHelloError {
|
| 88 |
+
OverlappingCapability(Capability),
|
| 89 |
+
}
|
| 90 |
+
|
| 91 |
+
impl fmt::Display for ClientHelloError {
|
| 92 |
+
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 93 |
+
match self {
|
| 94 |
+
Self::OverlappingCapability(capability) => write!(
|
| 95 |
+
formatter,
|
| 96 |
+
"capability `{capability}` cannot be both required and optional"
|
| 97 |
+
),
|
| 98 |
+
}
|
| 99 |
+
}
|
| 100 |
+
}
|
| 101 |
+
|
| 102 |
+
impl std::error::Error for ClientHelloError {}
|
| 103 |
+
|
| 104 |
+
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
| 105 |
+
#[serde(deny_unknown_fields, rename_all = "camelCase")]
|
| 106 |
+
pub struct HostHello {
|
| 107 |
+
selected_version: ProtocolVersion,
|
| 108 |
+
capabilities: CapabilitySet,
|
| 109 |
+
}
|
| 110 |
+
|
| 111 |
+
impl HostHello {
|
| 112 |
+
pub fn new(selected_version: ProtocolVersion, capabilities: CapabilitySet) -> Self {
|
| 113 |
+
Self {
|
| 114 |
+
selected_version,
|
| 115 |
+
capabilities,
|
| 116 |
+
}
|
| 117 |
+
}
|
| 118 |
+
|
| 119 |
+
pub fn selected_version(&self) -> ProtocolVersion {
|
| 120 |
+
self.selected_version
|
| 121 |
+
}
|
| 122 |
+
|
| 123 |
+
pub fn capabilities(&self) -> &CapabilitySet {
|
| 124 |
+
&self.capabilities
|
| 125 |
+
}
|
| 126 |
+
}
|
| 127 |
+
|
| 128 |
+
/// Messages sent from a client to the code-mode host.
|
| 129 |
+
#[derive(Debug, Deserialize, PartialEq, Serialize)]
|
| 130 |
+
#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")]
|
| 131 |
+
pub enum ClientToHost {
|
| 132 |
+
#[serde(rename = "connection/hello")]
|
| 133 |
+
ClientHello(ClientHello),
|
| 134 |
+
#[serde(rename = "operation/request")]
|
| 135 |
+
Request { id: RequestId, request: HostRequest },
|
| 136 |
+
#[serde(rename = "operation/cancel")]
|
| 137 |
+
CancelRequest { id: RequestId },
|
| 138 |
+
#[serde(rename = "delegate/response")]
|
| 139 |
+
DelegateResponse {
|
| 140 |
+
id: DelegateRequestId,
|
| 141 |
+
result: WireResult<DelegateResponse>,
|
| 142 |
+
},
|
| 143 |
+
}
|
| 144 |
+
|
| 145 |
+
/// Messages sent from the code-mode host to a client.
|
| 146 |
+
#[derive(Debug, Deserialize, PartialEq, Serialize)]
|
| 147 |
+
#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")]
|
| 148 |
+
pub enum HostToClient {
|
| 149 |
+
#[serde(rename = "connection/ready")]
|
| 150 |
+
HostHello(HostHello),
|
| 151 |
+
#[serde(rename = "connection/rejected")]
|
| 152 |
+
HandshakeRejected { reason: HandshakeRejectReason },
|
| 153 |
+
#[serde(rename = "operation/response")]
|
| 154 |
+
Response {
|
| 155 |
+
id: RequestId,
|
| 156 |
+
result: WireResult<HostResponse>,
|
| 157 |
+
},
|
| 158 |
+
#[serde(rename = "execute/initialResponse")]
|
| 159 |
+
InitialResponse {
|
| 160 |
+
id: RequestId,
|
| 161 |
+
result: WireResult<WireRuntimeResponse>,
|
| 162 |
+
},
|
| 163 |
+
#[serde(rename = "delegate/request")]
|
| 164 |
+
DelegateRequest {
|
| 165 |
+
id: DelegateRequestId,
|
| 166 |
+
session_id: SessionId,
|
| 167 |
+
request: DelegateRequest,
|
| 168 |
+
},
|
| 169 |
+
#[serde(rename = "delegate/cancel")]
|
| 170 |
+
CancelDelegateRequest { id: DelegateRequestId },
|
| 171 |
+
#[serde(rename = "cell/closed")]
|
| 172 |
+
CellClosed {
|
| 173 |
+
session_id: SessionId,
|
| 174 |
+
cell_id: WireCellId,
|
| 175 |
+
},
|
| 176 |
+
}
|
| 177 |
+
|
| 178 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 179 |
+
#[serde(deny_unknown_fields, tag = "method", rename_all_fields = "camelCase")]
|
| 180 |
+
pub enum HostRequest {
|
| 181 |
+
#[serde(rename = "session/open")]
|
| 182 |
+
OpenSession {
|
| 183 |
+
session_id: SessionId,
|
| 184 |
+
#[serde(default, skip_serializing_if = "Option::is_none")]
|
| 185 |
+
cell_execution_limits: Option<WireSessionCellExecutionLimits>,
|
| 186 |
+
},
|
| 187 |
+
#[serde(rename = "session/execute")]
|
| 188 |
+
Execute {
|
| 189 |
+
session_id: SessionId,
|
| 190 |
+
request: WireExecuteRequest,
|
| 191 |
+
},
|
| 192 |
+
#[serde(rename = "session/wait")]
|
| 193 |
+
Wait {
|
| 194 |
+
session_id: SessionId,
|
| 195 |
+
request: WireWaitRequest,
|
| 196 |
+
},
|
| 197 |
+
#[serde(rename = "session/terminate")]
|
| 198 |
+
Terminate {
|
| 199 |
+
session_id: SessionId,
|
| 200 |
+
cell_id: WireCellId,
|
| 201 |
+
},
|
| 202 |
+
#[serde(rename = "session/shutdown")]
|
| 203 |
+
ShutdownSession { session_id: SessionId },
|
| 204 |
+
}
|
| 205 |
+
|
| 206 |
+
#[derive(Debug, Deserialize, PartialEq, Serialize)]
|
| 207 |
+
#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")]
|
| 208 |
+
pub enum HostResponse {
|
| 209 |
+
#[serde(rename = "session/ready")]
|
| 210 |
+
SessionReady { session_id: SessionId },
|
| 211 |
+
#[serde(rename = "execution/started")]
|
| 212 |
+
ExecutionStarted { cell_id: WireCellId },
|
| 213 |
+
#[serde(rename = "wait/completed")]
|
| 214 |
+
WaitCompleted { outcome: WireWaitOutcome },
|
| 215 |
+
#[serde(rename = "session/closed")]
|
| 216 |
+
SessionClosed { session_id: SessionId },
|
| 217 |
+
}
|
| 218 |
+
|
| 219 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 220 |
+
#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")]
|
| 221 |
+
pub enum DelegateRequest {
|
| 222 |
+
#[serde(rename = "tool/invoke")]
|
| 223 |
+
InvokeTool { invocation: WireNestedToolCall },
|
| 224 |
+
#[serde(rename = "notification/send")]
|
| 225 |
+
Notify {
|
| 226 |
+
call_id: String,
|
| 227 |
+
cell_id: WireCellId,
|
| 228 |
+
text: String,
|
| 229 |
+
},
|
| 230 |
+
}
|
| 231 |
+
|
| 232 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 233 |
+
#[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")]
|
| 234 |
+
pub enum DelegateResponse {
|
| 235 |
+
#[serde(rename = "tool/result")]
|
| 236 |
+
ToolResult { result: JsonValue },
|
| 237 |
+
#[serde(rename = "notification/delivered")]
|
| 238 |
+
NotificationDelivered,
|
| 239 |
+
}
|
| 240 |
+
|
| 241 |
+
#[derive(Debug, Deserialize, PartialEq, Serialize)]
|
| 242 |
+
#[serde(deny_unknown_fields, tag = "status", rename_all_fields = "camelCase")]
|
| 243 |
+
pub enum WireResult<T> {
|
| 244 |
+
#[serde(rename = "ok")]
|
| 245 |
+
Ok { value: T },
|
| 246 |
+
#[serde(rename = "error")]
|
| 247 |
+
Err { message: String },
|
| 248 |
+
}
|
| 249 |
+
|
| 250 |
+
impl<T> WireResult<T> {
|
| 251 |
+
pub fn from_result(result: Result<T, String>) -> Self {
|
| 252 |
+
match result {
|
| 253 |
+
Ok(value) => Self::Ok { value },
|
| 254 |
+
Err(message) => Self::Err { message },
|
| 255 |
+
}
|
| 256 |
+
}
|
| 257 |
+
|
| 258 |
+
pub fn into_result(self) -> Result<T, String> {
|
| 259 |
+
match self {
|
| 260 |
+
Self::Ok { value } => Ok(value),
|
| 261 |
+
Self::Err { message } => Err(message),
|
| 262 |
+
}
|
| 263 |
+
}
|
| 264 |
+
}
|
codex-rs/code-mode-protocol/src/host/mod.rs
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Messages and framing for the code-mode host boundary.
|
| 2 |
+
//!
|
| 3 |
+
//! Protocol version 1 multiplexes session operations and delegate callbacks by
|
| 4 |
+
//! request ID over one ordered connection.
|
| 5 |
+
|
| 6 |
+
mod codec;
|
| 7 |
+
mod error;
|
| 8 |
+
mod message;
|
| 9 |
+
mod payload;
|
| 10 |
+
mod types;
|
| 11 |
+
|
| 12 |
+
/// Maximum number of unresolved delegate callbacks allowed per host connection.
|
| 13 |
+
pub const MAX_PENDING_DELEGATE_CALLS: usize = 1_024;
|
| 14 |
+
|
| 15 |
+
/// Negotiated support for cell execution resource limits on `session/open`.
|
| 16 |
+
pub const SESSION_RESOURCE_LIMITS_CAPABILITY: &str = "session-cell-execution-resource-limits";
|
| 17 |
+
|
| 18 |
+
pub use codec::EncodedFrame;
|
| 19 |
+
pub use codec::FramedReader;
|
| 20 |
+
pub use codec::FramedWriter;
|
| 21 |
+
pub use codec::MAX_FRAME_BYTES;
|
| 22 |
+
pub use error::HandshakeRejectReason;
|
| 23 |
+
pub use message::ClientHello;
|
| 24 |
+
pub use message::ClientHelloError;
|
| 25 |
+
pub use message::ClientToHost;
|
| 26 |
+
pub use message::DelegateRequest;
|
| 27 |
+
pub use message::DelegateResponse;
|
| 28 |
+
pub use message::HostHello;
|
| 29 |
+
pub use message::HostRequest;
|
| 30 |
+
pub use message::HostResponse;
|
| 31 |
+
pub use message::HostToClient;
|
| 32 |
+
pub use message::WireResult;
|
| 33 |
+
pub use payload::WireCellId;
|
| 34 |
+
pub use payload::WireContentItem;
|
| 35 |
+
pub use payload::WireExecuteRequest;
|
| 36 |
+
pub use payload::WireImageDetail;
|
| 37 |
+
pub use payload::WireNestedToolCall;
|
| 38 |
+
pub use payload::WireRuntimeResponse;
|
| 39 |
+
pub use payload::WireSessionCellExecutionLimits;
|
| 40 |
+
pub use payload::WireToolDefinition;
|
| 41 |
+
pub use payload::WireToolKind;
|
| 42 |
+
pub use payload::WireToolName;
|
| 43 |
+
pub use payload::WireWaitOutcome;
|
| 44 |
+
pub use payload::WireWaitRequest;
|
| 45 |
+
pub use types::Capability;
|
| 46 |
+
pub use types::CapabilitySet;
|
| 47 |
+
pub use types::DelegateRequestId;
|
| 48 |
+
pub use types::DuplicateCapability;
|
| 49 |
+
pub use types::InvalidIdentifier;
|
| 50 |
+
pub use types::InvalidSupportedProtocolVersions;
|
| 51 |
+
pub use types::ProtocolVersion;
|
| 52 |
+
pub use types::RequestId;
|
| 53 |
+
pub use types::SessionId;
|
| 54 |
+
pub use types::SupportedProtocolVersions;
|
| 55 |
+
|
| 56 |
+
#[cfg(test)]
|
| 57 |
+
#[path = "host_tests.rs"]
|
| 58 |
+
mod tests;
|
| 59 |
+
|
| 60 |
+
#[cfg(test)]
|
| 61 |
+
#[path = "codec_tests.rs"]
|
| 62 |
+
mod codec_tests;
|
codex-rs/code-mode-protocol/src/host/payload.rs
ADDED
|
@@ -0,0 +1,492 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::num::TryFromIntError;
|
| 2 |
+
use std::time::Duration;
|
| 3 |
+
|
| 4 |
+
use codex_protocol::ToolName;
|
| 5 |
+
use serde::Deserialize;
|
| 6 |
+
use serde::Serialize;
|
| 7 |
+
use serde_json::Value as JsonValue;
|
| 8 |
+
|
| 9 |
+
use crate::CellId;
|
| 10 |
+
use crate::CodeModeNestedToolCall;
|
| 11 |
+
use crate::CodeModeSessionCellExecutionLimits;
|
| 12 |
+
use crate::CodeModeToolKind;
|
| 13 |
+
use crate::ExecuteRequest;
|
| 14 |
+
use crate::FunctionCallOutputContentItem;
|
| 15 |
+
use crate::ImageDetail;
|
| 16 |
+
use crate::MissingCodeModeHostDuration;
|
| 17 |
+
use crate::RuntimeResponse;
|
| 18 |
+
use crate::ToolDefinition;
|
| 19 |
+
use crate::WaitOutcome;
|
| 20 |
+
use crate::WaitRequest;
|
| 21 |
+
|
| 22 |
+
/// The per-cell execution limits carried by a V1 session-open request.
|
| 23 |
+
#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
|
| 24 |
+
#[serde(deny_unknown_fields, rename_all = "camelCase")]
|
| 25 |
+
pub struct WireSessionCellExecutionLimits {
|
| 26 |
+
#[serde(default, skip_serializing_if = "Option::is_none")]
|
| 27 |
+
pub max_yield_time_ms: Option<u64>,
|
| 28 |
+
#[serde(default, skip_serializing_if = "Option::is_none")]
|
| 29 |
+
pub max_heap_size_bytes: Option<u64>,
|
| 30 |
+
}
|
| 31 |
+
|
| 32 |
+
impl TryFrom<CodeModeSessionCellExecutionLimits> for WireSessionCellExecutionLimits {
|
| 33 |
+
type Error = TryFromIntError;
|
| 34 |
+
|
| 35 |
+
fn try_from(value: CodeModeSessionCellExecutionLimits) -> Result<Self, Self::Error> {
|
| 36 |
+
Ok(Self {
|
| 37 |
+
max_yield_time_ms: value.max_yield_time_ms,
|
| 38 |
+
max_heap_size_bytes: value.max_heap_size_bytes.map(u64::try_from).transpose()?,
|
| 39 |
+
})
|
| 40 |
+
}
|
| 41 |
+
}
|
| 42 |
+
|
| 43 |
+
impl TryFrom<WireSessionCellExecutionLimits> for CodeModeSessionCellExecutionLimits {
|
| 44 |
+
type Error = TryFromIntError;
|
| 45 |
+
|
| 46 |
+
fn try_from(value: WireSessionCellExecutionLimits) -> Result<Self, Self::Error> {
|
| 47 |
+
Ok(Self {
|
| 48 |
+
max_yield_time_ms: value.max_yield_time_ms,
|
| 49 |
+
max_heap_size_bytes: value.max_heap_size_bytes.map(usize::try_from).transpose()?,
|
| 50 |
+
})
|
| 51 |
+
}
|
| 52 |
+
}
|
| 53 |
+
|
| 54 |
+
/// A cell identifier with a wire representation owned by protocol V1.
|
| 55 |
+
#[derive(Clone, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)]
|
| 56 |
+
#[serde(transparent)]
|
| 57 |
+
pub struct WireCellId(String);
|
| 58 |
+
|
| 59 |
+
impl WireCellId {
|
| 60 |
+
pub fn new(value: impl Into<String>) -> Self {
|
| 61 |
+
Self(value.into())
|
| 62 |
+
}
|
| 63 |
+
|
| 64 |
+
pub fn as_str(&self) -> &str {
|
| 65 |
+
&self.0
|
| 66 |
+
}
|
| 67 |
+
}
|
| 68 |
+
|
| 69 |
+
impl From<CellId> for WireCellId {
|
| 70 |
+
fn from(value: CellId) -> Self {
|
| 71 |
+
Self(value.as_str().to_string())
|
| 72 |
+
}
|
| 73 |
+
}
|
| 74 |
+
|
| 75 |
+
impl From<&CellId> for WireCellId {
|
| 76 |
+
fn from(value: &CellId) -> Self {
|
| 77 |
+
Self(value.as_str().to_string())
|
| 78 |
+
}
|
| 79 |
+
}
|
| 80 |
+
|
| 81 |
+
impl From<WireCellId> for CellId {
|
| 82 |
+
fn from(value: WireCellId) -> Self {
|
| 83 |
+
Self::new(value.0)
|
| 84 |
+
}
|
| 85 |
+
}
|
| 86 |
+
|
| 87 |
+
/// The V1 wire representation of a tool's stable name.
|
| 88 |
+
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
| 89 |
+
#[serde(deny_unknown_fields)]
|
| 90 |
+
pub struct WireToolName {
|
| 91 |
+
pub name: String,
|
| 92 |
+
pub namespace: Option<String>,
|
| 93 |
+
}
|
| 94 |
+
|
| 95 |
+
impl From<ToolName> for WireToolName {
|
| 96 |
+
fn from(value: ToolName) -> Self {
|
| 97 |
+
Self {
|
| 98 |
+
name: value.name,
|
| 99 |
+
namespace: value.namespace,
|
| 100 |
+
}
|
| 101 |
+
}
|
| 102 |
+
}
|
| 103 |
+
|
| 104 |
+
impl From<WireToolName> for ToolName {
|
| 105 |
+
fn from(value: WireToolName) -> Self {
|
| 106 |
+
Self::new(value.namespace, value.name)
|
| 107 |
+
}
|
| 108 |
+
}
|
| 109 |
+
|
| 110 |
+
/// The tool invocation shape supported by protocol V1.
|
| 111 |
+
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
| 112 |
+
#[serde(rename_all = "snake_case")]
|
| 113 |
+
pub enum WireToolKind {
|
| 114 |
+
Function,
|
| 115 |
+
Freeform,
|
| 116 |
+
}
|
| 117 |
+
|
| 118 |
+
impl From<CodeModeToolKind> for WireToolKind {
|
| 119 |
+
fn from(value: CodeModeToolKind) -> Self {
|
| 120 |
+
match value {
|
| 121 |
+
CodeModeToolKind::Function => Self::Function,
|
| 122 |
+
CodeModeToolKind::Freeform => Self::Freeform,
|
| 123 |
+
}
|
| 124 |
+
}
|
| 125 |
+
}
|
| 126 |
+
|
| 127 |
+
impl From<WireToolKind> for CodeModeToolKind {
|
| 128 |
+
fn from(value: WireToolKind) -> Self {
|
| 129 |
+
match value {
|
| 130 |
+
WireToolKind::Function => Self::Function,
|
| 131 |
+
WireToolKind::Freeform => Self::Freeform,
|
| 132 |
+
}
|
| 133 |
+
}
|
| 134 |
+
}
|
| 135 |
+
|
| 136 |
+
/// A V1 tool definition embedded in an execute request.
|
| 137 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 138 |
+
#[serde(deny_unknown_fields)]
|
| 139 |
+
pub struct WireToolDefinition {
|
| 140 |
+
pub name: String,
|
| 141 |
+
pub tool_name: WireToolName,
|
| 142 |
+
pub description: String,
|
| 143 |
+
pub kind: WireToolKind,
|
| 144 |
+
pub input_schema: Option<JsonValue>,
|
| 145 |
+
pub output_schema: Option<JsonValue>,
|
| 146 |
+
}
|
| 147 |
+
|
| 148 |
+
impl From<ToolDefinition> for WireToolDefinition {
|
| 149 |
+
fn from(value: ToolDefinition) -> Self {
|
| 150 |
+
Self {
|
| 151 |
+
name: value.name,
|
| 152 |
+
tool_name: value.tool_name.into(),
|
| 153 |
+
description: value.description,
|
| 154 |
+
kind: value.kind.into(),
|
| 155 |
+
input_schema: value.input_schema,
|
| 156 |
+
output_schema: value.output_schema,
|
| 157 |
+
}
|
| 158 |
+
}
|
| 159 |
+
}
|
| 160 |
+
|
| 161 |
+
impl From<WireToolDefinition> for ToolDefinition {
|
| 162 |
+
fn from(value: WireToolDefinition) -> Self {
|
| 163 |
+
Self {
|
| 164 |
+
name: value.name,
|
| 165 |
+
tool_name: value.tool_name.into(),
|
| 166 |
+
description: value.description,
|
| 167 |
+
kind: value.kind.into(),
|
| 168 |
+
input_schema: value.input_schema,
|
| 169 |
+
output_schema: value.output_schema,
|
| 170 |
+
}
|
| 171 |
+
}
|
| 172 |
+
}
|
| 173 |
+
|
| 174 |
+
/// The complete execute request shape supported by protocol V1.
|
| 175 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 176 |
+
#[serde(deny_unknown_fields)]
|
| 177 |
+
pub struct WireExecuteRequest {
|
| 178 |
+
pub tool_call_id: String,
|
| 179 |
+
pub enabled_tools: Vec<WireToolDefinition>,
|
| 180 |
+
pub source: String,
|
| 181 |
+
pub yield_time_ms: Option<u64>,
|
| 182 |
+
pub max_output_tokens: Option<i32>,
|
| 183 |
+
}
|
| 184 |
+
|
| 185 |
+
impl TryFrom<ExecuteRequest> for WireExecuteRequest {
|
| 186 |
+
type Error = TryFromIntError;
|
| 187 |
+
|
| 188 |
+
fn try_from(value: ExecuteRequest) -> Result<Self, Self::Error> {
|
| 189 |
+
Ok(Self {
|
| 190 |
+
tool_call_id: value.tool_call_id,
|
| 191 |
+
enabled_tools: value.enabled_tools.into_iter().map(Into::into).collect(),
|
| 192 |
+
source: value.source,
|
| 193 |
+
yield_time_ms: value.yield_time_ms,
|
| 194 |
+
max_output_tokens: value.max_output_tokens.map(i32::try_from).transpose()?,
|
| 195 |
+
})
|
| 196 |
+
}
|
| 197 |
+
}
|
| 198 |
+
|
| 199 |
+
impl TryFrom<WireExecuteRequest> for ExecuteRequest {
|
| 200 |
+
type Error = TryFromIntError;
|
| 201 |
+
|
| 202 |
+
fn try_from(value: WireExecuteRequest) -> Result<Self, Self::Error> {
|
| 203 |
+
Ok(Self {
|
| 204 |
+
tool_call_id: value.tool_call_id,
|
| 205 |
+
enabled_tools: value.enabled_tools.into_iter().map(Into::into).collect(),
|
| 206 |
+
source: value.source,
|
| 207 |
+
yield_time_ms: value.yield_time_ms,
|
| 208 |
+
max_output_tokens: value.max_output_tokens.map(usize::try_from).transpose()?,
|
| 209 |
+
})
|
| 210 |
+
}
|
| 211 |
+
}
|
| 212 |
+
|
| 213 |
+
/// The complete wait request shape supported by protocol V1.
|
| 214 |
+
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
| 215 |
+
#[serde(deny_unknown_fields)]
|
| 216 |
+
pub struct WireWaitRequest {
|
| 217 |
+
pub cell_id: WireCellId,
|
| 218 |
+
pub yield_time_ms: u64,
|
| 219 |
+
}
|
| 220 |
+
|
| 221 |
+
impl From<WaitRequest> for WireWaitRequest {
|
| 222 |
+
fn from(value: WaitRequest) -> Self {
|
| 223 |
+
Self {
|
| 224 |
+
cell_id: value.cell_id.into(),
|
| 225 |
+
yield_time_ms: value.yield_time_ms,
|
| 226 |
+
}
|
| 227 |
+
}
|
| 228 |
+
}
|
| 229 |
+
|
| 230 |
+
impl From<WireWaitRequest> for WaitRequest {
|
| 231 |
+
fn from(value: WireWaitRequest) -> Self {
|
| 232 |
+
Self {
|
| 233 |
+
cell_id: value.cell_id.into(),
|
| 234 |
+
yield_time_ms: value.yield_time_ms,
|
| 235 |
+
}
|
| 236 |
+
}
|
| 237 |
+
}
|
| 238 |
+
|
| 239 |
+
/// Image detail values accepted in a V1 runtime response.
|
| 240 |
+
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
| 241 |
+
#[serde(rename_all = "lowercase")]
|
| 242 |
+
pub enum WireImageDetail {
|
| 243 |
+
Auto,
|
| 244 |
+
Low,
|
| 245 |
+
High,
|
| 246 |
+
Original,
|
| 247 |
+
}
|
| 248 |
+
|
| 249 |
+
impl From<ImageDetail> for WireImageDetail {
|
| 250 |
+
fn from(value: ImageDetail) -> Self {
|
| 251 |
+
match value {
|
| 252 |
+
ImageDetail::Auto => Self::Auto,
|
| 253 |
+
ImageDetail::Low => Self::Low,
|
| 254 |
+
ImageDetail::High => Self::High,
|
| 255 |
+
ImageDetail::Original => Self::Original,
|
| 256 |
+
}
|
| 257 |
+
}
|
| 258 |
+
}
|
| 259 |
+
|
| 260 |
+
impl From<WireImageDetail> for ImageDetail {
|
| 261 |
+
fn from(value: WireImageDetail) -> Self {
|
| 262 |
+
match value {
|
| 263 |
+
WireImageDetail::Auto => Self::Auto,
|
| 264 |
+
WireImageDetail::Low => Self::Low,
|
| 265 |
+
WireImageDetail::High => Self::High,
|
| 266 |
+
WireImageDetail::Original => Self::Original,
|
| 267 |
+
}
|
| 268 |
+
}
|
| 269 |
+
}
|
| 270 |
+
|
| 271 |
+
/// One output item emitted by a V1 runtime response.
|
| 272 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 273 |
+
#[serde(deny_unknown_fields, tag = "type", rename_all = "snake_case")]
|
| 274 |
+
pub enum WireContentItem {
|
| 275 |
+
InputText {
|
| 276 |
+
text: String,
|
| 277 |
+
},
|
| 278 |
+
InputImage {
|
| 279 |
+
image_url: String,
|
| 280 |
+
#[serde(default, skip_serializing_if = "Option::is_none")]
|
| 281 |
+
detail: Option<WireImageDetail>,
|
| 282 |
+
},
|
| 283 |
+
InputAudio {
|
| 284 |
+
audio_url: String,
|
| 285 |
+
},
|
| 286 |
+
}
|
| 287 |
+
|
| 288 |
+
impl From<FunctionCallOutputContentItem> for WireContentItem {
|
| 289 |
+
fn from(value: FunctionCallOutputContentItem) -> Self {
|
| 290 |
+
match value {
|
| 291 |
+
FunctionCallOutputContentItem::InputText { text } => Self::InputText { text },
|
| 292 |
+
FunctionCallOutputContentItem::InputImage { image_url, detail } => Self::InputImage {
|
| 293 |
+
image_url,
|
| 294 |
+
detail: detail.map(Into::into),
|
| 295 |
+
},
|
| 296 |
+
FunctionCallOutputContentItem::InputAudio { audio_url } => {
|
| 297 |
+
Self::InputAudio { audio_url }
|
| 298 |
+
}
|
| 299 |
+
}
|
| 300 |
+
}
|
| 301 |
+
}
|
| 302 |
+
|
| 303 |
+
impl From<WireContentItem> for FunctionCallOutputContentItem {
|
| 304 |
+
fn from(value: WireContentItem) -> Self {
|
| 305 |
+
match value {
|
| 306 |
+
WireContentItem::InputText { text } => Self::InputText { text },
|
| 307 |
+
WireContentItem::InputImage { image_url, detail } => Self::InputImage {
|
| 308 |
+
image_url,
|
| 309 |
+
detail: detail.map(Into::into),
|
| 310 |
+
},
|
| 311 |
+
WireContentItem::InputAudio { audio_url } => Self::InputAudio { audio_url },
|
| 312 |
+
}
|
| 313 |
+
}
|
| 314 |
+
}
|
| 315 |
+
|
| 316 |
+
/// Runtime output returned over the V1 host connection.
|
| 317 |
+
///
|
| 318 |
+
/// Host time is required and covers this request, not the cell lifetime. The
|
| 319 |
+
/// app-server and host run at the same version; no negotiation is needed.
|
| 320 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 321 |
+
#[serde(deny_unknown_fields)]
|
| 322 |
+
pub enum WireRuntimeResponse {
|
| 323 |
+
Yielded {
|
| 324 |
+
cell_id: WireCellId,
|
| 325 |
+
content_items: Vec<WireContentItem>,
|
| 326 |
+
code_mode_host_duration_ns: u64,
|
| 327 |
+
},
|
| 328 |
+
Terminated {
|
| 329 |
+
cell_id: WireCellId,
|
| 330 |
+
content_items: Vec<WireContentItem>,
|
| 331 |
+
code_mode_host_duration_ns: u64,
|
| 332 |
+
},
|
| 333 |
+
Result {
|
| 334 |
+
cell_id: WireCellId,
|
| 335 |
+
content_items: Vec<WireContentItem>,
|
| 336 |
+
error_text: Option<String>,
|
| 337 |
+
code_mode_host_duration_ns: u64,
|
| 338 |
+
},
|
| 339 |
+
}
|
| 340 |
+
|
| 341 |
+
impl TryFrom<RuntimeResponse> for WireRuntimeResponse {
|
| 342 |
+
type Error = MissingCodeModeHostDuration;
|
| 343 |
+
|
| 344 |
+
/// Preserves the response's timing; the host handler must record it first.
|
| 345 |
+
fn try_from(value: RuntimeResponse) -> Result<Self, Self::Error> {
|
| 346 |
+
Ok(match value {
|
| 347 |
+
RuntimeResponse::Yielded {
|
| 348 |
+
cell_id,
|
| 349 |
+
content_items,
|
| 350 |
+
code_mode_host_duration,
|
| 351 |
+
} => {
|
| 352 |
+
let code_mode_host_duration =
|
| 353 |
+
code_mode_host_duration.ok_or(MissingCodeModeHostDuration)?;
|
| 354 |
+
Self::Yielded {
|
| 355 |
+
cell_id: cell_id.into(),
|
| 356 |
+
content_items: content_items.into_iter().map(Into::into).collect(),
|
| 357 |
+
code_mode_host_duration_ns: u64::try_from(code_mode_host_duration.as_nanos())
|
| 358 |
+
.unwrap_or(u64::MAX),
|
| 359 |
+
}
|
| 360 |
+
}
|
| 361 |
+
RuntimeResponse::Terminated {
|
| 362 |
+
cell_id,
|
| 363 |
+
content_items,
|
| 364 |
+
code_mode_host_duration,
|
| 365 |
+
} => {
|
| 366 |
+
let code_mode_host_duration =
|
| 367 |
+
code_mode_host_duration.ok_or(MissingCodeModeHostDuration)?;
|
| 368 |
+
Self::Terminated {
|
| 369 |
+
cell_id: cell_id.into(),
|
| 370 |
+
content_items: content_items.into_iter().map(Into::into).collect(),
|
| 371 |
+
code_mode_host_duration_ns: u64::try_from(code_mode_host_duration.as_nanos())
|
| 372 |
+
.unwrap_or(u64::MAX),
|
| 373 |
+
}
|
| 374 |
+
}
|
| 375 |
+
RuntimeResponse::Result {
|
| 376 |
+
cell_id,
|
| 377 |
+
content_items,
|
| 378 |
+
error_text,
|
| 379 |
+
code_mode_host_duration,
|
| 380 |
+
} => {
|
| 381 |
+
let code_mode_host_duration =
|
| 382 |
+
code_mode_host_duration.ok_or(MissingCodeModeHostDuration)?;
|
| 383 |
+
Self::Result {
|
| 384 |
+
cell_id: cell_id.into(),
|
| 385 |
+
content_items: content_items.into_iter().map(Into::into).collect(),
|
| 386 |
+
error_text,
|
| 387 |
+
code_mode_host_duration_ns: u64::try_from(code_mode_host_duration.as_nanos())
|
| 388 |
+
.unwrap_or(u64::MAX),
|
| 389 |
+
}
|
| 390 |
+
}
|
| 391 |
+
})
|
| 392 |
+
}
|
| 393 |
+
}
|
| 394 |
+
|
| 395 |
+
impl From<WireRuntimeResponse> for RuntimeResponse {
|
| 396 |
+
fn from(value: WireRuntimeResponse) -> Self {
|
| 397 |
+
match value {
|
| 398 |
+
WireRuntimeResponse::Yielded {
|
| 399 |
+
cell_id,
|
| 400 |
+
content_items,
|
| 401 |
+
code_mode_host_duration_ns,
|
| 402 |
+
} => Self::Yielded {
|
| 403 |
+
cell_id: cell_id.into(),
|
| 404 |
+
content_items: content_items.into_iter().map(Into::into).collect(),
|
| 405 |
+
code_mode_host_duration: Some(Duration::from_nanos(code_mode_host_duration_ns)),
|
| 406 |
+
},
|
| 407 |
+
WireRuntimeResponse::Terminated {
|
| 408 |
+
cell_id,
|
| 409 |
+
content_items,
|
| 410 |
+
code_mode_host_duration_ns,
|
| 411 |
+
} => Self::Terminated {
|
| 412 |
+
cell_id: cell_id.into(),
|
| 413 |
+
content_items: content_items.into_iter().map(Into::into).collect(),
|
| 414 |
+
code_mode_host_duration: Some(Duration::from_nanos(code_mode_host_duration_ns)),
|
| 415 |
+
},
|
| 416 |
+
WireRuntimeResponse::Result {
|
| 417 |
+
cell_id,
|
| 418 |
+
content_items,
|
| 419 |
+
error_text,
|
| 420 |
+
code_mode_host_duration_ns,
|
| 421 |
+
} => Self::Result {
|
| 422 |
+
cell_id: cell_id.into(),
|
| 423 |
+
content_items: content_items.into_iter().map(Into::into).collect(),
|
| 424 |
+
error_text,
|
| 425 |
+
code_mode_host_duration: Some(Duration::from_nanos(code_mode_host_duration_ns)),
|
| 426 |
+
},
|
| 427 |
+
}
|
| 428 |
+
}
|
| 429 |
+
}
|
| 430 |
+
|
| 431 |
+
/// Whether a waited-for cell remained live in protocol V1.
|
| 432 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 433 |
+
#[serde(deny_unknown_fields)]
|
| 434 |
+
pub enum WireWaitOutcome {
|
| 435 |
+
LiveCell(WireRuntimeResponse),
|
| 436 |
+
MissingCell(WireRuntimeResponse),
|
| 437 |
+
}
|
| 438 |
+
|
| 439 |
+
impl TryFrom<WaitOutcome> for WireWaitOutcome {
|
| 440 |
+
type Error = MissingCodeModeHostDuration;
|
| 441 |
+
|
| 442 |
+
fn try_from(value: WaitOutcome) -> Result<Self, Self::Error> {
|
| 443 |
+
Ok(match value {
|
| 444 |
+
WaitOutcome::LiveCell(response) => Self::LiveCell(response.try_into()?),
|
| 445 |
+
WaitOutcome::MissingCell(response) => Self::MissingCell(response.try_into()?),
|
| 446 |
+
})
|
| 447 |
+
}
|
| 448 |
+
}
|
| 449 |
+
|
| 450 |
+
impl From<WireWaitOutcome> for WaitOutcome {
|
| 451 |
+
fn from(value: WireWaitOutcome) -> Self {
|
| 452 |
+
match value {
|
| 453 |
+
WireWaitOutcome::LiveCell(response) => Self::LiveCell(response.into()),
|
| 454 |
+
WireWaitOutcome::MissingCell(response) => Self::MissingCell(response.into()),
|
| 455 |
+
}
|
| 456 |
+
}
|
| 457 |
+
}
|
| 458 |
+
|
| 459 |
+
/// A nested tool invocation sent over the V1 host connection.
|
| 460 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 461 |
+
#[serde(deny_unknown_fields)]
|
| 462 |
+
pub struct WireNestedToolCall {
|
| 463 |
+
pub cell_id: WireCellId,
|
| 464 |
+
pub runtime_tool_call_id: String,
|
| 465 |
+
pub tool_name: WireToolName,
|
| 466 |
+
pub tool_kind: WireToolKind,
|
| 467 |
+
pub input: Option<JsonValue>,
|
| 468 |
+
}
|
| 469 |
+
|
| 470 |
+
impl From<CodeModeNestedToolCall> for WireNestedToolCall {
|
| 471 |
+
fn from(value: CodeModeNestedToolCall) -> Self {
|
| 472 |
+
Self {
|
| 473 |
+
cell_id: value.cell_id.into(),
|
| 474 |
+
runtime_tool_call_id: value.runtime_tool_call_id,
|
| 475 |
+
tool_name: value.tool_name.into(),
|
| 476 |
+
tool_kind: value.tool_kind.into(),
|
| 477 |
+
input: value.input,
|
| 478 |
+
}
|
| 479 |
+
}
|
| 480 |
+
}
|
| 481 |
+
|
| 482 |
+
impl From<WireNestedToolCall> for CodeModeNestedToolCall {
|
| 483 |
+
fn from(value: WireNestedToolCall) -> Self {
|
| 484 |
+
Self {
|
| 485 |
+
cell_id: value.cell_id.into(),
|
| 486 |
+
runtime_tool_call_id: value.runtime_tool_call_id,
|
| 487 |
+
tool_name: value.tool_name.into(),
|
| 488 |
+
tool_kind: value.tool_kind.into(),
|
| 489 |
+
input: value.input,
|
| 490 |
+
}
|
| 491 |
+
}
|
| 492 |
+
}
|
codex-rs/code-mode-protocol/src/host/types.rs
ADDED
|
@@ -0,0 +1,248 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::collections::BTreeSet;
|
| 2 |
+
use std::fmt;
|
| 3 |
+
use std::num::NonZeroU32;
|
| 4 |
+
|
| 5 |
+
use serde::Deserialize;
|
| 6 |
+
use serde::Deserializer;
|
| 7 |
+
use serde::Serialize;
|
| 8 |
+
use serde::Serializer;
|
| 9 |
+
use serde::de::Error as _;
|
| 10 |
+
|
| 11 |
+
/// Correlates one client operation request with the host's response.
|
| 12 |
+
#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)]
|
| 13 |
+
#[serde(transparent)]
|
| 14 |
+
pub struct RequestId(i64);
|
| 15 |
+
|
| 16 |
+
impl RequestId {
|
| 17 |
+
pub const fn new(value: i64) -> Self {
|
| 18 |
+
Self(value)
|
| 19 |
+
}
|
| 20 |
+
}
|
| 21 |
+
|
| 22 |
+
/// Correlates one host delegate request with the client's response.
|
| 23 |
+
#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)]
|
| 24 |
+
#[serde(transparent)]
|
| 25 |
+
pub struct DelegateRequestId(i64);
|
| 26 |
+
|
| 27 |
+
impl DelegateRequestId {
|
| 28 |
+
pub const fn new(value: i64) -> Self {
|
| 29 |
+
Self(value)
|
| 30 |
+
}
|
| 31 |
+
}
|
| 32 |
+
|
| 33 |
+
#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)]
|
| 34 |
+
#[serde(transparent)]
|
| 35 |
+
pub struct ProtocolVersion(NonZeroU32);
|
| 36 |
+
|
| 37 |
+
impl ProtocolVersion {
|
| 38 |
+
pub const V1: Self = Self(NonZeroU32::MIN);
|
| 39 |
+
|
| 40 |
+
pub const fn new(value: u32) -> Option<Self> {
|
| 41 |
+
match NonZeroU32::new(value) {
|
| 42 |
+
Some(value) => Some(Self(value)),
|
| 43 |
+
None => None,
|
| 44 |
+
}
|
| 45 |
+
}
|
| 46 |
+
|
| 47 |
+
pub const fn get(self) -> u32 {
|
| 48 |
+
self.0.get()
|
| 49 |
+
}
|
| 50 |
+
}
|
| 51 |
+
|
| 52 |
+
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
| 53 |
+
pub struct InvalidIdentifier;
|
| 54 |
+
|
| 55 |
+
impl fmt::Display for InvalidIdentifier {
|
| 56 |
+
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 57 |
+
formatter.write_str("identifier must not be empty")
|
| 58 |
+
}
|
| 59 |
+
}
|
| 60 |
+
|
| 61 |
+
impl std::error::Error for InvalidIdentifier {}
|
| 62 |
+
|
| 63 |
+
#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
|
| 64 |
+
struct NonEmptyString(String);
|
| 65 |
+
|
| 66 |
+
impl NonEmptyString {
|
| 67 |
+
fn new(value: impl Into<String>) -> Result<Self, InvalidIdentifier> {
|
| 68 |
+
let value = value.into();
|
| 69 |
+
if value.trim().is_empty() {
|
| 70 |
+
Err(InvalidIdentifier)
|
| 71 |
+
} else {
|
| 72 |
+
Ok(Self(value))
|
| 73 |
+
}
|
| 74 |
+
}
|
| 75 |
+
|
| 76 |
+
fn as_str(&self) -> &str {
|
| 77 |
+
&self.0
|
| 78 |
+
}
|
| 79 |
+
}
|
| 80 |
+
|
| 81 |
+
impl Serialize for NonEmptyString {
|
| 82 |
+
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
| 83 |
+
where
|
| 84 |
+
S: Serializer,
|
| 85 |
+
{
|
| 86 |
+
self.0.serialize(serializer)
|
| 87 |
+
}
|
| 88 |
+
}
|
| 89 |
+
|
| 90 |
+
impl<'de> Deserialize<'de> for NonEmptyString {
|
| 91 |
+
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
| 92 |
+
where
|
| 93 |
+
D: Deserializer<'de>,
|
| 94 |
+
{
|
| 95 |
+
Self::new(String::deserialize(deserializer)?).map_err(D::Error::custom)
|
| 96 |
+
}
|
| 97 |
+
}
|
| 98 |
+
|
| 99 |
+
/// A named protocol feature advertised during connection negotiation.
|
| 100 |
+
#[derive(Clone, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)]
|
| 101 |
+
#[serde(transparent)]
|
| 102 |
+
pub struct Capability(NonEmptyString);
|
| 103 |
+
|
| 104 |
+
impl Capability {
|
| 105 |
+
pub fn new(value: impl Into<String>) -> Result<Self, InvalidIdentifier> {
|
| 106 |
+
NonEmptyString::new(value).map(Self)
|
| 107 |
+
}
|
| 108 |
+
|
| 109 |
+
pub fn as_str(&self) -> &str {
|
| 110 |
+
self.0.as_str()
|
| 111 |
+
}
|
| 112 |
+
}
|
| 113 |
+
|
| 114 |
+
impl fmt::Display for Capability {
|
| 115 |
+
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 116 |
+
formatter.write_str(self.as_str())
|
| 117 |
+
}
|
| 118 |
+
}
|
| 119 |
+
|
| 120 |
+
/// Identifies one logical code-mode session on a connection.
|
| 121 |
+
#[derive(Clone, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)]
|
| 122 |
+
#[serde(transparent)]
|
| 123 |
+
pub struct SessionId(NonEmptyString);
|
| 124 |
+
|
| 125 |
+
impl SessionId {
|
| 126 |
+
pub fn new(value: impl Into<String>) -> Result<Self, InvalidIdentifier> {
|
| 127 |
+
NonEmptyString::new(value).map(Self)
|
| 128 |
+
}
|
| 129 |
+
|
| 130 |
+
pub fn as_str(&self) -> &str {
|
| 131 |
+
self.0.as_str()
|
| 132 |
+
}
|
| 133 |
+
}
|
| 134 |
+
|
| 135 |
+
impl fmt::Display for SessionId {
|
| 136 |
+
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 137 |
+
formatter.write_str(self.as_str())
|
| 138 |
+
}
|
| 139 |
+
}
|
| 140 |
+
|
| 141 |
+
#[derive(Clone, Debug, Default, Eq, PartialEq, Serialize)]
|
| 142 |
+
#[serde(transparent)]
|
| 143 |
+
pub struct CapabilitySet(BTreeSet<Capability>);
|
| 144 |
+
|
| 145 |
+
impl CapabilitySet {
|
| 146 |
+
pub fn empty() -> Self {
|
| 147 |
+
Self::default()
|
| 148 |
+
}
|
| 149 |
+
|
| 150 |
+
pub fn try_new(
|
| 151 |
+
capabilities: impl IntoIterator<Item = Capability>,
|
| 152 |
+
) -> Result<Self, DuplicateCapability> {
|
| 153 |
+
let mut unique = BTreeSet::new();
|
| 154 |
+
for capability in capabilities {
|
| 155 |
+
if !unique.insert(capability.clone()) {
|
| 156 |
+
return Err(DuplicateCapability { capability });
|
| 157 |
+
}
|
| 158 |
+
}
|
| 159 |
+
Ok(Self(unique))
|
| 160 |
+
}
|
| 161 |
+
|
| 162 |
+
pub fn contains(&self, capability: &Capability) -> bool {
|
| 163 |
+
self.0.contains(capability)
|
| 164 |
+
}
|
| 165 |
+
|
| 166 |
+
pub fn iter(&self) -> impl Iterator<Item = &Capability> {
|
| 167 |
+
self.0.iter()
|
| 168 |
+
}
|
| 169 |
+
}
|
| 170 |
+
|
| 171 |
+
impl<'de> Deserialize<'de> for CapabilitySet {
|
| 172 |
+
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
| 173 |
+
where
|
| 174 |
+
D: Deserializer<'de>,
|
| 175 |
+
{
|
| 176 |
+
Self::try_new(Vec::<Capability>::deserialize(deserializer)?).map_err(D::Error::custom)
|
| 177 |
+
}
|
| 178 |
+
}
|
| 179 |
+
|
| 180 |
+
#[derive(Clone, Debug, Eq, PartialEq)]
|
| 181 |
+
pub struct DuplicateCapability {
|
| 182 |
+
capability: Capability,
|
| 183 |
+
}
|
| 184 |
+
|
| 185 |
+
impl fmt::Display for DuplicateCapability {
|
| 186 |
+
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 187 |
+
write!(formatter, "duplicate capability `{}`", self.capability)
|
| 188 |
+
}
|
| 189 |
+
}
|
| 190 |
+
|
| 191 |
+
impl std::error::Error for DuplicateCapability {}
|
| 192 |
+
|
| 193 |
+
#[derive(Clone, Debug, Eq, PartialEq, Serialize)]
|
| 194 |
+
#[serde(transparent)]
|
| 195 |
+
pub struct SupportedProtocolVersions(BTreeSet<ProtocolVersion>);
|
| 196 |
+
|
| 197 |
+
impl SupportedProtocolVersions {
|
| 198 |
+
pub fn try_new(
|
| 199 |
+
versions: impl IntoIterator<Item = ProtocolVersion>,
|
| 200 |
+
) -> Result<Self, InvalidSupportedProtocolVersions> {
|
| 201 |
+
let mut unique = BTreeSet::new();
|
| 202 |
+
for version in versions {
|
| 203 |
+
if !unique.insert(version) {
|
| 204 |
+
return Err(InvalidSupportedProtocolVersions::Duplicate(version));
|
| 205 |
+
}
|
| 206 |
+
}
|
| 207 |
+
if unique.is_empty() {
|
| 208 |
+
return Err(InvalidSupportedProtocolVersions::Empty);
|
| 209 |
+
}
|
| 210 |
+
Ok(Self(unique))
|
| 211 |
+
}
|
| 212 |
+
|
| 213 |
+
pub fn contains(&self, version: ProtocolVersion) -> bool {
|
| 214 |
+
self.0.contains(&version)
|
| 215 |
+
}
|
| 216 |
+
|
| 217 |
+
pub fn iter(&self) -> impl Iterator<Item = ProtocolVersion> + '_ {
|
| 218 |
+
self.0.iter().copied()
|
| 219 |
+
}
|
| 220 |
+
}
|
| 221 |
+
|
| 222 |
+
impl<'de> Deserialize<'de> for SupportedProtocolVersions {
|
| 223 |
+
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
| 224 |
+
where
|
| 225 |
+
D: Deserializer<'de>,
|
| 226 |
+
{
|
| 227 |
+
Self::try_new(Vec::<ProtocolVersion>::deserialize(deserializer)?).map_err(D::Error::custom)
|
| 228 |
+
}
|
| 229 |
+
}
|
| 230 |
+
|
| 231 |
+
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
| 232 |
+
pub enum InvalidSupportedProtocolVersions {
|
| 233 |
+
Empty,
|
| 234 |
+
Duplicate(ProtocolVersion),
|
| 235 |
+
}
|
| 236 |
+
|
| 237 |
+
impl fmt::Display for InvalidSupportedProtocolVersions {
|
| 238 |
+
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 239 |
+
match self {
|
| 240 |
+
Self::Empty => formatter.write_str("at least one protocol version is required"),
|
| 241 |
+
Self::Duplicate(version) => {
|
| 242 |
+
write!(formatter, "duplicate protocol version {}", version.get())
|
| 243 |
+
}
|
| 244 |
+
}
|
| 245 |
+
}
|
| 246 |
+
}
|
| 247 |
+
|
| 248 |
+
impl std::error::Error for InvalidSupportedProtocolVersions {}
|
codex-rs/code-mode-protocol/src/json_schema_types.rs
ADDED
|
@@ -0,0 +1,538 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use serde_json::Value as JsonValue;
|
| 2 |
+
use std::collections::BTreeMap;
|
| 3 |
+
|
| 4 |
+
use crate::description::normalize_code_mode_identifier;
|
| 5 |
+
|
| 6 |
+
// Expose one nested recursive shape, then fall back to `unknown` on the next
|
| 7 |
+
// occurrence so generated tool declarations remain finite.
|
| 8 |
+
const MAX_LOCAL_REF_EXPANSIONS_PER_PATH: usize = 2;
|
| 9 |
+
// Bound repeated refs and DAG fan-out separately from cycle depth so one
|
| 10 |
+
// compact schema cannot expand into an arbitrarily large model-visible item.
|
| 11 |
+
const MAX_TOTAL_LOCAL_REF_EXPANSIONS: usize = 32;
|
| 12 |
+
// Bound individual schema rendering before assembling the final declaration.
|
| 13 |
+
const MAX_RENDERED_SCHEMA_BYTES: usize = 16_000;
|
| 14 |
+
// Charge intermediate render strings as they are built so repeated local refs
|
| 15 |
+
// cannot allocate unbounded expanded copies before the final schema cap runs.
|
| 16 |
+
const MAX_RENDER_WORK_BYTES: usize = MAX_RENDERED_SCHEMA_BYTES * 4;
|
| 17 |
+
|
| 18 |
+
pub fn render_json_schema_to_typescript(schema: &JsonValue) -> String {
|
| 19 |
+
let rendered = JsonSchemaTypeRenderer::new(schema).render(schema);
|
| 20 |
+
if rendered.len() > MAX_RENDERED_SCHEMA_BYTES {
|
| 21 |
+
"unknown".to_string()
|
| 22 |
+
} else {
|
| 23 |
+
rendered
|
| 24 |
+
}
|
| 25 |
+
}
|
| 26 |
+
|
| 27 |
+
struct JsonSchemaTypeRenderer<'a> {
|
| 28 |
+
root: &'a JsonValue,
|
| 29 |
+
nested_schema_resource_depth: usize,
|
| 30 |
+
active_local_ref_expansions: BTreeMap<String, usize>,
|
| 31 |
+
remaining_local_ref_expansions: usize,
|
| 32 |
+
remaining_render_work_bytes: usize,
|
| 33 |
+
render_work_budget_exhausted: bool,
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
impl<'a> JsonSchemaTypeRenderer<'a> {
|
| 37 |
+
fn new(root: &'a JsonValue) -> Self {
|
| 38 |
+
Self {
|
| 39 |
+
root,
|
| 40 |
+
nested_schema_resource_depth: 0,
|
| 41 |
+
active_local_ref_expansions: BTreeMap::new(),
|
| 42 |
+
remaining_local_ref_expansions: MAX_TOTAL_LOCAL_REF_EXPANSIONS,
|
| 43 |
+
remaining_render_work_bytes: MAX_RENDER_WORK_BYTES,
|
| 44 |
+
render_work_budget_exhausted: false,
|
| 45 |
+
}
|
| 46 |
+
}
|
| 47 |
+
|
| 48 |
+
fn render(&mut self, schema: &JsonValue) -> String {
|
| 49 |
+
if self.render_work_budget_exhausted {
|
| 50 |
+
return "unknown".to_string();
|
| 51 |
+
}
|
| 52 |
+
|
| 53 |
+
// A nested `$id` starts a new schema resource. Fragment-only refs below
|
| 54 |
+
// it are scoped to that resource, not to the outer document root.
|
| 55 |
+
let enters_nested_schema_resource = !std::ptr::eq(schema, self.root)
|
| 56 |
+
&& schema
|
| 57 |
+
.as_object()
|
| 58 |
+
.is_some_and(|map| map.contains_key("$id"));
|
| 59 |
+
if enters_nested_schema_resource {
|
| 60 |
+
self.nested_schema_resource_depth += 1;
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
+
let rendered = match schema {
|
| 64 |
+
JsonValue::Bool(true) => "unknown".to_string(),
|
| 65 |
+
JsonValue::Bool(false) => "never".to_string(),
|
| 66 |
+
JsonValue::Object(map) => self.render_map(map),
|
| 67 |
+
_ => "unknown".to_string(),
|
| 68 |
+
};
|
| 69 |
+
if enters_nested_schema_resource {
|
| 70 |
+
self.nested_schema_resource_depth -= 1;
|
| 71 |
+
}
|
| 72 |
+
self.finish_render(rendered)
|
| 73 |
+
}
|
| 74 |
+
|
| 75 |
+
fn render_map(&mut self, map: &serde_json::Map<String, JsonValue>) -> String {
|
| 76 |
+
if self.render_work_budget_exhausted {
|
| 77 |
+
return "unknown".to_string();
|
| 78 |
+
}
|
| 79 |
+
|
| 80 |
+
if map.contains_key("$ref") {
|
| 81 |
+
return self.render_ref(map);
|
| 82 |
+
}
|
| 83 |
+
|
| 84 |
+
if let Some(value) = map.get("const") {
|
| 85 |
+
return self.render_literal(value);
|
| 86 |
+
}
|
| 87 |
+
|
| 88 |
+
if let Some(values) = map.get("enum").and_then(JsonValue::as_array) {
|
| 89 |
+
let mut rendered = Vec::new();
|
| 90 |
+
for value in values {
|
| 91 |
+
let literal = self.render_literal(value);
|
| 92 |
+
if self.render_work_budget_exhausted {
|
| 93 |
+
return "unknown".to_string();
|
| 94 |
+
}
|
| 95 |
+
if !self.consume_render_work(literal.len()) {
|
| 96 |
+
return "unknown".to_string();
|
| 97 |
+
}
|
| 98 |
+
rendered.push(literal);
|
| 99 |
+
}
|
| 100 |
+
if !rendered.is_empty() {
|
| 101 |
+
return rendered.join(" | ");
|
| 102 |
+
}
|
| 103 |
+
}
|
| 104 |
+
|
| 105 |
+
for key in ["anyOf", "oneOf"] {
|
| 106 |
+
if let Some(variants) = map.get(key).and_then(JsonValue::as_array) {
|
| 107 |
+
let mut rendered = Vec::new();
|
| 108 |
+
for variant in variants {
|
| 109 |
+
if self.render_work_budget_exhausted {
|
| 110 |
+
return "unknown".to_string();
|
| 111 |
+
}
|
| 112 |
+
rendered.push(self.render(variant));
|
| 113 |
+
}
|
| 114 |
+
if !rendered.is_empty() {
|
| 115 |
+
return rendered.join(" | ");
|
| 116 |
+
}
|
| 117 |
+
}
|
| 118 |
+
}
|
| 119 |
+
|
| 120 |
+
if let Some(variants) = map.get("allOf").and_then(JsonValue::as_array) {
|
| 121 |
+
let mut rendered = Vec::new();
|
| 122 |
+
for variant in variants {
|
| 123 |
+
if self.render_work_budget_exhausted {
|
| 124 |
+
return "unknown".to_string();
|
| 125 |
+
}
|
| 126 |
+
rendered.push(parenthesize_union_for_intersection(self.render(variant)));
|
| 127 |
+
}
|
| 128 |
+
if !rendered.is_empty() {
|
| 129 |
+
return rendered.join(" & ");
|
| 130 |
+
}
|
| 131 |
+
}
|
| 132 |
+
|
| 133 |
+
if let Some(schema_type) = map.get("type") {
|
| 134 |
+
if let Some(types) = schema_type.as_array() {
|
| 135 |
+
let mut rendered = Vec::new();
|
| 136 |
+
for schema_type in types.iter().filter_map(JsonValue::as_str) {
|
| 137 |
+
if self.render_work_budget_exhausted {
|
| 138 |
+
return "unknown".to_string();
|
| 139 |
+
}
|
| 140 |
+
rendered.push(self.render_type_keyword(map, schema_type));
|
| 141 |
+
}
|
| 142 |
+
if !rendered.is_empty() {
|
| 143 |
+
return rendered.join(" | ");
|
| 144 |
+
}
|
| 145 |
+
}
|
| 146 |
+
|
| 147 |
+
if let Some(schema_type) = schema_type.as_str() {
|
| 148 |
+
return self.render_type_keyword(map, schema_type);
|
| 149 |
+
}
|
| 150 |
+
}
|
| 151 |
+
|
| 152 |
+
if map.contains_key("properties")
|
| 153 |
+
|| map.contains_key("additionalProperties")
|
| 154 |
+
|| map.contains_key("required")
|
| 155 |
+
{
|
| 156 |
+
return self.render_object(map);
|
| 157 |
+
}
|
| 158 |
+
|
| 159 |
+
if map.contains_key("items") || map.contains_key("prefixItems") {
|
| 160 |
+
return self.render_array(map);
|
| 161 |
+
}
|
| 162 |
+
|
| 163 |
+
"unknown".to_string()
|
| 164 |
+
}
|
| 165 |
+
|
| 166 |
+
fn render_ref(&mut self, map: &serde_json::Map<String, JsonValue>) -> String {
|
| 167 |
+
let referenced_type = if self.nested_schema_resource_depth > 0 {
|
| 168 |
+
None
|
| 169 |
+
} else {
|
| 170 |
+
map.get("$ref")
|
| 171 |
+
.and_then(JsonValue::as_str)
|
| 172 |
+
.and_then(local_json_pointer)
|
| 173 |
+
.and_then(|pointer| {
|
| 174 |
+
let active_expansions = self
|
| 175 |
+
.active_local_ref_expansions
|
| 176 |
+
.get(&pointer)
|
| 177 |
+
.copied()
|
| 178 |
+
.unwrap_or_default();
|
| 179 |
+
if active_expansions >= MAX_LOCAL_REF_EXPANSIONS_PER_PATH {
|
| 180 |
+
return Some("unknown".to_string());
|
| 181 |
+
}
|
| 182 |
+
if self.remaining_local_ref_expansions == 0 {
|
| 183 |
+
return Some("unknown".to_string());
|
| 184 |
+
}
|
| 185 |
+
|
| 186 |
+
let root = self.root;
|
| 187 |
+
let target = if pointer.is_empty() {
|
| 188 |
+
Some(root)
|
| 189 |
+
} else {
|
| 190 |
+
root.pointer(&pointer)
|
| 191 |
+
}?;
|
| 192 |
+
self.remaining_local_ref_expansions -= 1;
|
| 193 |
+
self.active_local_ref_expansions
|
| 194 |
+
.insert(pointer.clone(), active_expansions + 1);
|
| 195 |
+
|
| 196 |
+
let rendered = self.render(target);
|
| 197 |
+
if active_expansions == 0 {
|
| 198 |
+
self.active_local_ref_expansions.remove(&pointer);
|
| 199 |
+
} else {
|
| 200 |
+
self.active_local_ref_expansions
|
| 201 |
+
.insert(pointer, active_expansions);
|
| 202 |
+
}
|
| 203 |
+
Some(rendered)
|
| 204 |
+
})
|
| 205 |
+
}
|
| 206 |
+
.unwrap_or_else(|| "unknown".to_string());
|
| 207 |
+
if self.render_work_budget_exhausted {
|
| 208 |
+
return "unknown".to_string();
|
| 209 |
+
}
|
| 210 |
+
|
| 211 |
+
let siblings = map
|
| 212 |
+
.iter()
|
| 213 |
+
.filter(|(key, _)| !matches!(key.as_str(), "$ref" | "$defs" | "definitions"))
|
| 214 |
+
.map(|(key, value)| (key.clone(), value.clone()))
|
| 215 |
+
.collect();
|
| 216 |
+
if !has_renderable_schema_keywords(&siblings) {
|
| 217 |
+
return referenced_type;
|
| 218 |
+
}
|
| 219 |
+
|
| 220 |
+
let sibling_type = self.render_map(&siblings);
|
| 221 |
+
match (referenced_type.as_str(), sibling_type.as_str()) {
|
| 222 |
+
("unknown", _) => sibling_type,
|
| 223 |
+
(_, "unknown") => referenced_type,
|
| 224 |
+
_ => format!("({referenced_type}) & ({sibling_type})"),
|
| 225 |
+
}
|
| 226 |
+
}
|
| 227 |
+
|
| 228 |
+
fn render_type_keyword(
|
| 229 |
+
&mut self,
|
| 230 |
+
map: &serde_json::Map<String, JsonValue>,
|
| 231 |
+
schema_type: &str,
|
| 232 |
+
) -> String {
|
| 233 |
+
match schema_type {
|
| 234 |
+
"string" => "string".to_string(),
|
| 235 |
+
"number" | "integer" => "number".to_string(),
|
| 236 |
+
"boolean" => "boolean".to_string(),
|
| 237 |
+
"null" => "null".to_string(),
|
| 238 |
+
"array" => self.render_array(map),
|
| 239 |
+
"object" => self.render_object(map),
|
| 240 |
+
_ => "unknown".to_string(),
|
| 241 |
+
}
|
| 242 |
+
}
|
| 243 |
+
|
| 244 |
+
fn render_array(&mut self, map: &serde_json::Map<String, JsonValue>) -> String {
|
| 245 |
+
if let Some(items) = map.get("items") {
|
| 246 |
+
let item_type = self.render(items);
|
| 247 |
+
if self.render_work_budget_exhausted {
|
| 248 |
+
return "unknown".to_string();
|
| 249 |
+
}
|
| 250 |
+
return format!("Array<{item_type}>");
|
| 251 |
+
}
|
| 252 |
+
|
| 253 |
+
if let Some(items) = map.get("prefixItems").and_then(JsonValue::as_array) {
|
| 254 |
+
let mut item_types = Vec::new();
|
| 255 |
+
for item in items {
|
| 256 |
+
if self.render_work_budget_exhausted {
|
| 257 |
+
return "unknown".to_string();
|
| 258 |
+
}
|
| 259 |
+
item_types.push(self.render(item));
|
| 260 |
+
}
|
| 261 |
+
if !item_types.is_empty() {
|
| 262 |
+
return format!("[{}]", item_types.join(", "));
|
| 263 |
+
}
|
| 264 |
+
}
|
| 265 |
+
|
| 266 |
+
"unknown[]".to_string()
|
| 267 |
+
}
|
| 268 |
+
|
| 269 |
+
fn append_additional_properties_line(
|
| 270 |
+
&mut self,
|
| 271 |
+
lines: &mut Vec<String>,
|
| 272 |
+
map: &serde_json::Map<String, JsonValue>,
|
| 273 |
+
properties: &serde_json::Map<String, JsonValue>,
|
| 274 |
+
line_prefix: &str,
|
| 275 |
+
) -> bool {
|
| 276 |
+
if let Some(additional_properties) = map.get("additionalProperties") {
|
| 277 |
+
let property_type = match additional_properties {
|
| 278 |
+
JsonValue::Bool(true) => Some("unknown".to_string()),
|
| 279 |
+
JsonValue::Bool(false) => None,
|
| 280 |
+
value => Some(self.render(value)),
|
| 281 |
+
};
|
| 282 |
+
|
| 283 |
+
if let Some(property_type) = property_type {
|
| 284 |
+
return self.push_render_line(
|
| 285 |
+
lines,
|
| 286 |
+
format!("{line_prefix}[key: string]: {property_type};"),
|
| 287 |
+
);
|
| 288 |
+
}
|
| 289 |
+
} else if properties.is_empty() {
|
| 290 |
+
return self.push_render_line(lines, format!("{line_prefix}[key: string]: unknown;"));
|
| 291 |
+
}
|
| 292 |
+
true
|
| 293 |
+
}
|
| 294 |
+
|
| 295 |
+
fn render_object_property(
|
| 296 |
+
&mut self,
|
| 297 |
+
name: &str,
|
| 298 |
+
value: &JsonValue,
|
| 299 |
+
required: &[&str],
|
| 300 |
+
) -> String {
|
| 301 |
+
if name.len() > self.remaining_render_work_bytes {
|
| 302 |
+
self.render_work_budget_exhausted = true;
|
| 303 |
+
return "unknown".to_string();
|
| 304 |
+
}
|
| 305 |
+
let optional = if required.iter().any(|required_name| required_name == &name) {
|
| 306 |
+
""
|
| 307 |
+
} else {
|
| 308 |
+
"?"
|
| 309 |
+
};
|
| 310 |
+
let property_name = render_json_schema_property_name(name);
|
| 311 |
+
let property_type = self.render(value);
|
| 312 |
+
if self.render_work_budget_exhausted {
|
| 313 |
+
return "unknown".to_string();
|
| 314 |
+
}
|
| 315 |
+
format!("{property_name}{optional}: {property_type};")
|
| 316 |
+
}
|
| 317 |
+
|
| 318 |
+
fn render_object(&mut self, map: &serde_json::Map<String, JsonValue>) -> String {
|
| 319 |
+
let required = map
|
| 320 |
+
.get("required")
|
| 321 |
+
.and_then(JsonValue::as_array)
|
| 322 |
+
.map(|items| {
|
| 323 |
+
items
|
| 324 |
+
.iter()
|
| 325 |
+
.filter_map(JsonValue::as_str)
|
| 326 |
+
.collect::<Vec<_>>()
|
| 327 |
+
})
|
| 328 |
+
.unwrap_or_default();
|
| 329 |
+
let empty_properties = serde_json::Map::new();
|
| 330 |
+
let properties = map
|
| 331 |
+
.get("properties")
|
| 332 |
+
.and_then(JsonValue::as_object)
|
| 333 |
+
.unwrap_or(&empty_properties);
|
| 334 |
+
|
| 335 |
+
let mut sorted_properties = properties.iter().collect::<Vec<_>>();
|
| 336 |
+
sorted_properties.sort_unstable_by_key(|(name_a, _)| *name_a);
|
| 337 |
+
if sorted_properties
|
| 338 |
+
.iter()
|
| 339 |
+
.any(|(_, value)| has_property_description(value))
|
| 340 |
+
{
|
| 341 |
+
let mut lines = Vec::new();
|
| 342 |
+
if !self.push_render_line(&mut lines, "{".to_string()) {
|
| 343 |
+
return "unknown".to_string();
|
| 344 |
+
}
|
| 345 |
+
for (name, value) in sorted_properties {
|
| 346 |
+
if let Some(description) = value.get("description").and_then(JsonValue::as_str) {
|
| 347 |
+
for description_line in description
|
| 348 |
+
.lines()
|
| 349 |
+
.map(str::trim)
|
| 350 |
+
.filter(|line| !line.is_empty())
|
| 351 |
+
{
|
| 352 |
+
if description_line.len().saturating_add(5)
|
| 353 |
+
> self.remaining_render_work_bytes
|
| 354 |
+
|| !self
|
| 355 |
+
.push_render_line(&mut lines, format!(" // {description_line}"))
|
| 356 |
+
{
|
| 357 |
+
return "unknown".to_string();
|
| 358 |
+
}
|
| 359 |
+
}
|
| 360 |
+
}
|
| 361 |
+
|
| 362 |
+
let property = self.render_object_property(name, value, &required);
|
| 363 |
+
if self.render_work_budget_exhausted
|
| 364 |
+
|| !self.push_render_line(&mut lines, format!(" {property}"))
|
| 365 |
+
{
|
| 366 |
+
return "unknown".to_string();
|
| 367 |
+
}
|
| 368 |
+
}
|
| 369 |
+
|
| 370 |
+
if !self.append_additional_properties_line(&mut lines, map, properties, " ")
|
| 371 |
+
|| !self.push_render_line(&mut lines, "}".to_string())
|
| 372 |
+
{
|
| 373 |
+
return "unknown".to_string();
|
| 374 |
+
}
|
| 375 |
+
return lines.join("\n");
|
| 376 |
+
}
|
| 377 |
+
|
| 378 |
+
let mut lines = Vec::new();
|
| 379 |
+
for (name, value) in sorted_properties {
|
| 380 |
+
let property = self.render_object_property(name, value, &required);
|
| 381 |
+
if self.render_work_budget_exhausted || !self.push_render_line(&mut lines, property) {
|
| 382 |
+
return "unknown".to_string();
|
| 383 |
+
}
|
| 384 |
+
}
|
| 385 |
+
|
| 386 |
+
if !self.append_additional_properties_line(&mut lines, map, properties, "") {
|
| 387 |
+
return "unknown".to_string();
|
| 388 |
+
}
|
| 389 |
+
|
| 390 |
+
if lines.is_empty() {
|
| 391 |
+
return "{}".to_string();
|
| 392 |
+
}
|
| 393 |
+
|
| 394 |
+
format!("{{ {} }}", lines.join(" "))
|
| 395 |
+
}
|
| 396 |
+
|
| 397 |
+
fn finish_render(&mut self, rendered: String) -> String {
|
| 398 |
+
if self.consume_render_work(rendered.len()) {
|
| 399 |
+
rendered
|
| 400 |
+
} else {
|
| 401 |
+
"unknown".to_string()
|
| 402 |
+
}
|
| 403 |
+
}
|
| 404 |
+
|
| 405 |
+
fn render_literal(&mut self, value: &JsonValue) -> String {
|
| 406 |
+
if json_literal_serialization_upper_bound(value) > self.remaining_render_work_bytes {
|
| 407 |
+
self.render_work_budget_exhausted = true;
|
| 408 |
+
"unknown".to_string()
|
| 409 |
+
} else {
|
| 410 |
+
render_json_schema_literal(value)
|
| 411 |
+
}
|
| 412 |
+
}
|
| 413 |
+
|
| 414 |
+
fn consume_render_work(&mut self, rendered_bytes: usize) -> bool {
|
| 415 |
+
if rendered_bytes > self.remaining_render_work_bytes {
|
| 416 |
+
self.render_work_budget_exhausted = true;
|
| 417 |
+
false
|
| 418 |
+
} else {
|
| 419 |
+
self.remaining_render_work_bytes -= rendered_bytes;
|
| 420 |
+
true
|
| 421 |
+
}
|
| 422 |
+
}
|
| 423 |
+
|
| 424 |
+
fn push_render_line(&mut self, lines: &mut Vec<String>, line: String) -> bool {
|
| 425 |
+
if !self.consume_render_work(line.len()) {
|
| 426 |
+
return false;
|
| 427 |
+
}
|
| 428 |
+
lines.push(line);
|
| 429 |
+
true
|
| 430 |
+
}
|
| 431 |
+
}
|
| 432 |
+
|
| 433 |
+
fn parenthesize_union_for_intersection(rendered: String) -> String {
|
| 434 |
+
if rendered.contains(" | ") {
|
| 435 |
+
format!("({rendered})")
|
| 436 |
+
} else {
|
| 437 |
+
rendered
|
| 438 |
+
}
|
| 439 |
+
}
|
| 440 |
+
|
| 441 |
+
fn local_json_pointer(reference: &str) -> Option<String> {
|
| 442 |
+
let fragment = reference.strip_prefix('#')?;
|
| 443 |
+
let pointer = percent_decode_uri_fragment(fragment)?;
|
| 444 |
+
if pointer.is_empty() || pointer.starts_with('/') {
|
| 445 |
+
Some(pointer)
|
| 446 |
+
} else {
|
| 447 |
+
None
|
| 448 |
+
}
|
| 449 |
+
}
|
| 450 |
+
|
| 451 |
+
fn percent_decode_uri_fragment(fragment: &str) -> Option<String> {
|
| 452 |
+
let bytes = fragment.as_bytes();
|
| 453 |
+
let mut decoded = Vec::with_capacity(bytes.len());
|
| 454 |
+
let mut index = 0;
|
| 455 |
+
while index < bytes.len() {
|
| 456 |
+
if bytes[index] == b'%' {
|
| 457 |
+
let high = decode_hex_digit(*bytes.get(index + 1)?)?;
|
| 458 |
+
let low = decode_hex_digit(*bytes.get(index + 2)?)?;
|
| 459 |
+
decoded.push((high << 4) | low);
|
| 460 |
+
index += 3;
|
| 461 |
+
} else {
|
| 462 |
+
decoded.push(bytes[index]);
|
| 463 |
+
index += 1;
|
| 464 |
+
}
|
| 465 |
+
}
|
| 466 |
+
String::from_utf8(decoded).ok()
|
| 467 |
+
}
|
| 468 |
+
|
| 469 |
+
fn decode_hex_digit(digit: u8) -> Option<u8> {
|
| 470 |
+
match digit {
|
| 471 |
+
b'0'..=b'9' => Some(digit - b'0'),
|
| 472 |
+
b'a'..=b'f' => Some(digit - b'a' + 10),
|
| 473 |
+
b'A'..=b'F' => Some(digit - b'A' + 10),
|
| 474 |
+
_ => None,
|
| 475 |
+
}
|
| 476 |
+
}
|
| 477 |
+
|
| 478 |
+
fn has_renderable_schema_keywords(map: &serde_json::Map<String, JsonValue>) -> bool {
|
| 479 |
+
[
|
| 480 |
+
"const",
|
| 481 |
+
"enum",
|
| 482 |
+
"anyOf",
|
| 483 |
+
"oneOf",
|
| 484 |
+
"allOf",
|
| 485 |
+
"type",
|
| 486 |
+
"properties",
|
| 487 |
+
"additionalProperties",
|
| 488 |
+
"required",
|
| 489 |
+
"items",
|
| 490 |
+
"prefixItems",
|
| 491 |
+
]
|
| 492 |
+
.iter()
|
| 493 |
+
.any(|key| map.contains_key(*key))
|
| 494 |
+
}
|
| 495 |
+
|
| 496 |
+
fn has_property_description(value: &JsonValue) -> bool {
|
| 497 |
+
value
|
| 498 |
+
.get("description")
|
| 499 |
+
.and_then(JsonValue::as_str)
|
| 500 |
+
.is_some_and(|description| !description.is_empty())
|
| 501 |
+
}
|
| 502 |
+
|
| 503 |
+
fn render_json_schema_property_name(name: &str) -> String {
|
| 504 |
+
if normalize_code_mode_identifier(name) == name {
|
| 505 |
+
name.to_string()
|
| 506 |
+
} else {
|
| 507 |
+
serde_json::to_string(name).unwrap_or_else(|_| format!("\"{}\"", name.replace('"', "\\\"")))
|
| 508 |
+
}
|
| 509 |
+
}
|
| 510 |
+
|
| 511 |
+
fn render_json_schema_literal(value: &JsonValue) -> String {
|
| 512 |
+
serde_json::to_string(value).unwrap_or_else(|_| "unknown".to_string())
|
| 513 |
+
}
|
| 514 |
+
|
| 515 |
+
fn json_literal_serialization_upper_bound(value: &JsonValue) -> usize {
|
| 516 |
+
match value {
|
| 517 |
+
JsonValue::Null => 4,
|
| 518 |
+
JsonValue::Bool(false) => 5,
|
| 519 |
+
JsonValue::Bool(true) => 4,
|
| 520 |
+
JsonValue::Number(number) => number.to_string().len(),
|
| 521 |
+
// JSON escaping can expand one UTF-8 byte to at most one six-byte
|
| 522 |
+
// Unicode escape, so this bounds allocation before serialization.
|
| 523 |
+
JsonValue::String(string) => string.len().saturating_mul(6).saturating_add(2),
|
| 524 |
+
JsonValue::Array(values) => values.iter().fold(2, |size, value| {
|
| 525 |
+
size.saturating_add(1)
|
| 526 |
+
.saturating_add(json_literal_serialization_upper_bound(value))
|
| 527 |
+
}),
|
| 528 |
+
JsonValue::Object(map) => map.iter().fold(2, |size, (key, value)| {
|
| 529 |
+
size.saturating_add(4)
|
| 530 |
+
.saturating_add(key.len().saturating_mul(6))
|
| 531 |
+
.saturating_add(json_literal_serialization_upper_bound(value))
|
| 532 |
+
}),
|
| 533 |
+
}
|
| 534 |
+
}
|
| 535 |
+
|
| 536 |
+
#[cfg(test)]
|
| 537 |
+
#[path = "json_schema_types_tests.rs"]
|
| 538 |
+
mod tests;
|
codex-rs/code-mode-protocol/src/json_schema_types_tests.rs
ADDED
|
@@ -0,0 +1,200 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::*;
|
| 2 |
+
use pretty_assertions::assert_eq;
|
| 3 |
+
use serde_json::json;
|
| 4 |
+
|
| 5 |
+
#[test]
|
| 6 |
+
fn renders_recursive_local_refs_with_escaped_pointer_segments() {
|
| 7 |
+
let schema = json!({
|
| 8 |
+
"type": "object",
|
| 9 |
+
"properties": {
|
| 10 |
+
"clauses": {
|
| 11 |
+
"type": "array",
|
| 12 |
+
"items": { "$ref": "#/$defs/Boolean~1Clause~0v1" }
|
| 13 |
+
}
|
| 14 |
+
},
|
| 15 |
+
"$defs": {
|
| 16 |
+
"Boolean/Clause~v1": {
|
| 17 |
+
"type": "object",
|
| 18 |
+
"properties": {
|
| 19 |
+
"query": { "$ref": "#/$defs/Query" }
|
| 20 |
+
}
|
| 21 |
+
},
|
| 22 |
+
"Query": {
|
| 23 |
+
"oneOf": [
|
| 24 |
+
{ "type": "string" },
|
| 25 |
+
{
|
| 26 |
+
"type": "object",
|
| 27 |
+
"properties": {
|
| 28 |
+
"clauses": {
|
| 29 |
+
"type": "array",
|
| 30 |
+
"items": { "$ref": "#/$defs/Boolean~1Clause~0v1" }
|
| 31 |
+
}
|
| 32 |
+
}
|
| 33 |
+
}
|
| 34 |
+
]
|
| 35 |
+
}
|
| 36 |
+
}
|
| 37 |
+
});
|
| 38 |
+
|
| 39 |
+
let rendered = render_json_schema_to_typescript(&schema);
|
| 40 |
+
assert!(rendered.contains("clauses?: Array<{ query?: string | { clauses?: Array<{"));
|
| 41 |
+
assert!(rendered.contains("query?: string | { clauses?: Array<unknown>; };"));
|
| 42 |
+
}
|
| 43 |
+
|
| 44 |
+
#[test]
|
| 45 |
+
fn renders_ref_siblings_uri_fragments_and_all_of_precedence() {
|
| 46 |
+
assert_eq!(
|
| 47 |
+
render_json_schema_to_typescript(&json!({
|
| 48 |
+
"$ref": "#/$defs/Label",
|
| 49 |
+
"enum": ["A"],
|
| 50 |
+
"$defs": { "Label": { "type": "string" } }
|
| 51 |
+
})),
|
| 52 |
+
r#"(string) & ("A")"#
|
| 53 |
+
);
|
| 54 |
+
assert_eq!(
|
| 55 |
+
render_json_schema_to_typescript(&json!({
|
| 56 |
+
"$ref": "#/$defs/Foo%20Bar",
|
| 57 |
+
"$defs": { "Foo Bar": { "type": "string" } }
|
| 58 |
+
})),
|
| 59 |
+
"string"
|
| 60 |
+
);
|
| 61 |
+
assert_eq!(
|
| 62 |
+
render_json_schema_to_typescript(&json!({
|
| 63 |
+
"allOf": [
|
| 64 |
+
{ "$ref": "#/$defs/Choice" },
|
| 65 |
+
{ "type": "object", "properties": { "value": { "type": "string" } } }
|
| 66 |
+
],
|
| 67 |
+
"$defs": {
|
| 68 |
+
"Choice": { "oneOf": [{ "type": "string" }, { "type": "number" }] }
|
| 69 |
+
}
|
| 70 |
+
})),
|
| 71 |
+
"(string | number) & { value?: string; }"
|
| 72 |
+
);
|
| 73 |
+
}
|
| 74 |
+
|
| 75 |
+
#[test]
|
| 76 |
+
fn leaves_local_refs_under_nested_schema_resources_unresolved() {
|
| 77 |
+
let schema = json!({
|
| 78 |
+
"$defs": {
|
| 79 |
+
"Choice": { "type": "string" }
|
| 80 |
+
},
|
| 81 |
+
"type": "object",
|
| 82 |
+
"properties": {
|
| 83 |
+
"nested": {
|
| 84 |
+
"$id": "urn:nested",
|
| 85 |
+
"$defs": {
|
| 86 |
+
"Choice": { "type": "number" }
|
| 87 |
+
},
|
| 88 |
+
"type": "object",
|
| 89 |
+
"properties": {
|
| 90 |
+
"value": { "$ref": "#/$defs/Choice" }
|
| 91 |
+
}
|
| 92 |
+
}
|
| 93 |
+
}
|
| 94 |
+
});
|
| 95 |
+
|
| 96 |
+
assert_eq!(
|
| 97 |
+
render_json_schema_to_typescript(&schema),
|
| 98 |
+
"{ nested?: { value?: unknown; }; }"
|
| 99 |
+
);
|
| 100 |
+
}
|
| 101 |
+
|
| 102 |
+
#[test]
|
| 103 |
+
fn bounds_expansions_without_charging_dangling_refs() {
|
| 104 |
+
let mut properties = (0..MAX_TOTAL_LOCAL_REF_EXPANSIONS)
|
| 105 |
+
.map(|index| {
|
| 106 |
+
(
|
| 107 |
+
format!("a_missing_{index}"),
|
| 108 |
+
json!({ "$ref": format!("#/$defs/Missing{index}") }),
|
| 109 |
+
)
|
| 110 |
+
})
|
| 111 |
+
.collect::<serde_json::Map<_, _>>();
|
| 112 |
+
properties.insert("z_valid".to_string(), json!({ "$ref": "#/$defs/Valid" }));
|
| 113 |
+
let schema = json!({
|
| 114 |
+
"type": "object",
|
| 115 |
+
"properties": properties,
|
| 116 |
+
"$defs": { "Valid": { "type": "string" } }
|
| 117 |
+
});
|
| 118 |
+
|
| 119 |
+
let rendered = render_json_schema_to_typescript(&schema);
|
| 120 |
+
assert!(rendered.contains("z_valid?: string;"));
|
| 121 |
+
|
| 122 |
+
let properties = (0..MAX_TOTAL_LOCAL_REF_EXPANSIONS + 2)
|
| 123 |
+
.map(|index| {
|
| 124 |
+
(
|
| 125 |
+
format!("property_{index}"),
|
| 126 |
+
json!({ "$ref": "#/$defs/Item" }),
|
| 127 |
+
)
|
| 128 |
+
})
|
| 129 |
+
.collect::<serde_json::Map<_, _>>();
|
| 130 |
+
let rendered = render_json_schema_to_typescript(&json!({
|
| 131 |
+
"type": "object",
|
| 132 |
+
"properties": properties,
|
| 133 |
+
"$defs": { "Item": { "type": "string" } }
|
| 134 |
+
}));
|
| 135 |
+
assert_eq!(
|
| 136 |
+
rendered.matches("string").count(),
|
| 137 |
+
MAX_TOTAL_LOCAL_REF_EXPANSIONS
|
| 138 |
+
);
|
| 139 |
+
assert_eq!(rendered.matches("unknown").count(), 2);
|
| 140 |
+
}
|
| 141 |
+
|
| 142 |
+
#[test]
|
| 143 |
+
fn repeated_large_ref_expansions_exhaust_render_work_budget() {
|
| 144 |
+
let properties = (0..MAX_TOTAL_LOCAL_REF_EXPANSIONS)
|
| 145 |
+
.map(|index| {
|
| 146 |
+
(
|
| 147 |
+
format!("property_{index}"),
|
| 148 |
+
json!({ "$ref": "#/$defs/Item" }),
|
| 149 |
+
)
|
| 150 |
+
})
|
| 151 |
+
.collect::<serde_json::Map<_, _>>();
|
| 152 |
+
let schema = json!({
|
| 153 |
+
"type": "object",
|
| 154 |
+
"properties": properties,
|
| 155 |
+
"$defs": {
|
| 156 |
+
"Item": {
|
| 157 |
+
"type": "object",
|
| 158 |
+
"properties": {
|
| 159 |
+
"value": {
|
| 160 |
+
"type": "string",
|
| 161 |
+
"description": "x".repeat(MAX_RENDERED_SCHEMA_BYTES / 2)
|
| 162 |
+
}
|
| 163 |
+
}
|
| 164 |
+
}
|
| 165 |
+
}
|
| 166 |
+
});
|
| 167 |
+
|
| 168 |
+
let mut renderer = JsonSchemaTypeRenderer::new(&schema);
|
| 169 |
+
assert_eq!(renderer.render(&schema), "unknown");
|
| 170 |
+
assert!(renderer.render_work_budget_exhausted);
|
| 171 |
+
}
|
| 172 |
+
|
| 173 |
+
#[test]
|
| 174 |
+
fn oversized_ref_literal_exhausts_render_work_budget() {
|
| 175 |
+
let schema = json!({
|
| 176 |
+
"$ref": "#/$defs/Value",
|
| 177 |
+
"$defs": {
|
| 178 |
+
"Value": {
|
| 179 |
+
"const": "x".repeat(MAX_RENDER_WORK_BYTES)
|
| 180 |
+
}
|
| 181 |
+
}
|
| 182 |
+
});
|
| 183 |
+
|
| 184 |
+
let mut renderer = JsonSchemaTypeRenderer::new(&schema);
|
| 185 |
+
assert_eq!(renderer.render(&schema), "unknown");
|
| 186 |
+
assert!(renderer.render_work_budget_exhausted);
|
| 187 |
+
}
|
| 188 |
+
|
| 189 |
+
#[test]
|
| 190 |
+
fn rendered_schema_has_a_hard_size_cap() {
|
| 191 |
+
let description = "x".repeat(MAX_RENDERED_SCHEMA_BYTES);
|
| 192 |
+
let schema = json!({
|
| 193 |
+
"type": "object",
|
| 194 |
+
"properties": {
|
| 195 |
+
"value": { "type": "string", "description": description }
|
| 196 |
+
}
|
| 197 |
+
});
|
| 198 |
+
|
| 199 |
+
assert_eq!(render_json_schema_to_typescript(&schema), "unknown");
|
| 200 |
+
}
|
codex-rs/code-mode-protocol/src/lib.rs
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
mod description;
|
| 2 |
+
pub mod grpc;
|
| 3 |
+
pub mod host;
|
| 4 |
+
mod json_schema_types;
|
| 5 |
+
mod response;
|
| 6 |
+
mod runtime;
|
| 7 |
+
mod session;
|
| 8 |
+
|
| 9 |
+
pub use description::CODE_MODE_PRAGMA_PREFIX;
|
| 10 |
+
pub use description::CodeModeToolKind;
|
| 11 |
+
pub use description::EnabledToolMetadata;
|
| 12 |
+
pub use description::ImageDetailVisibility;
|
| 13 |
+
pub use description::ToolDefinition;
|
| 14 |
+
pub use description::ToolNamespaceDescription;
|
| 15 |
+
pub use description::augment_tool_definition;
|
| 16 |
+
pub use description::build_exec_tool_description;
|
| 17 |
+
pub use description::build_wait_tool_description;
|
| 18 |
+
pub use description::enabled_tool_metadata;
|
| 19 |
+
pub use description::is_code_mode_nested_tool;
|
| 20 |
+
pub use description::normalize_code_mode_identifier;
|
| 21 |
+
pub use description::parse_exec_source;
|
| 22 |
+
pub use description::render_code_mode_sample;
|
| 23 |
+
pub use json_schema_types::render_json_schema_to_typescript;
|
| 24 |
+
pub use response::DEFAULT_IMAGE_DETAIL;
|
| 25 |
+
pub use response::FunctionCallOutputContentItem;
|
| 26 |
+
pub use response::ImageDetail;
|
| 27 |
+
pub use runtime::CodeModeNestedToolCall;
|
| 28 |
+
pub use runtime::DEFAULT_EXEC_YIELD_TIME_MS;
|
| 29 |
+
pub use runtime::DEFAULT_MAX_OUTPUT_TOKENS_PER_EXEC_CALL;
|
| 30 |
+
pub use runtime::DEFAULT_WAIT_YIELD_TIME_MS;
|
| 31 |
+
pub use runtime::ExecuteRequest;
|
| 32 |
+
pub use runtime::ExecuteToPendingOutcome;
|
| 33 |
+
pub use runtime::MissingCodeModeHostDuration;
|
| 34 |
+
pub use runtime::RuntimeResponse;
|
| 35 |
+
pub use runtime::WaitOutcome;
|
| 36 |
+
pub use runtime::WaitRequest;
|
| 37 |
+
pub use runtime::WaitToPendingOutcome;
|
| 38 |
+
pub use runtime::WaitToPendingRequest;
|
| 39 |
+
pub use session::CellId;
|
| 40 |
+
pub use session::CodeModeSession;
|
| 41 |
+
pub use session::CodeModeSessionCellExecutionLimits;
|
| 42 |
+
pub use session::CodeModeSessionDelegate;
|
| 43 |
+
pub use session::CodeModeSessionProvider;
|
| 44 |
+
pub use session::CodeModeSessionProviderFuture;
|
| 45 |
+
pub use session::CodeModeSessionResultFuture;
|
| 46 |
+
pub use session::NoopCodeModeSessionDelegate;
|
| 47 |
+
pub use session::NotificationFuture;
|
| 48 |
+
pub use session::StartedCell;
|
| 49 |
+
pub use session::ToolInvocationFuture;
|
| 50 |
+
|
| 51 |
+
pub const PUBLIC_TOOL_NAME: &str = "exec";
|
| 52 |
+
pub const WAIT_TOOL_NAME: &str = "wait";
|
codex-rs/code-mode-protocol/src/response.rs
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use serde::Deserialize;
|
| 2 |
+
use serde::Serialize;
|
| 3 |
+
|
| 4 |
+
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
|
| 5 |
+
#[serde(rename_all = "lowercase")]
|
| 6 |
+
pub enum ImageDetail {
|
| 7 |
+
Auto,
|
| 8 |
+
Low,
|
| 9 |
+
High,
|
| 10 |
+
Original,
|
| 11 |
+
}
|
| 12 |
+
|
| 13 |
+
pub const DEFAULT_IMAGE_DETAIL: ImageDetail = ImageDetail::High;
|
| 14 |
+
|
| 15 |
+
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
| 16 |
+
#[serde(tag = "type", rename_all = "snake_case")]
|
| 17 |
+
pub enum FunctionCallOutputContentItem {
|
| 18 |
+
InputText {
|
| 19 |
+
text: String,
|
| 20 |
+
},
|
| 21 |
+
InputImage {
|
| 22 |
+
image_url: String,
|
| 23 |
+
#[serde(default, skip_serializing_if = "Option::is_none")]
|
| 24 |
+
detail: Option<ImageDetail>,
|
| 25 |
+
},
|
| 26 |
+
InputAudio {
|
| 27 |
+
audio_url: String,
|
| 28 |
+
},
|
| 29 |
+
}
|
codex-rs/code-mode-protocol/src/runtime.rs
ADDED
|
@@ -0,0 +1,183 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::error::Error;
|
| 2 |
+
use std::fmt;
|
| 3 |
+
use std::time::Duration;
|
| 4 |
+
|
| 5 |
+
use codex_protocol::ToolName;
|
| 6 |
+
use serde::Deserialize;
|
| 7 |
+
use serde::Serialize;
|
| 8 |
+
use serde_json::Value as JsonValue;
|
| 9 |
+
|
| 10 |
+
use crate::CellId;
|
| 11 |
+
use crate::CodeModeToolKind;
|
| 12 |
+
use crate::FunctionCallOutputContentItem;
|
| 13 |
+
use crate::ToolDefinition;
|
| 14 |
+
|
| 15 |
+
pub const DEFAULT_EXEC_YIELD_TIME_MS: u64 = 10_000;
|
| 16 |
+
pub const DEFAULT_WAIT_YIELD_TIME_MS: u64 = 10_000;
|
| 17 |
+
pub const DEFAULT_MAX_OUTPUT_TOKENS_PER_EXEC_CALL: usize = 10_000;
|
| 18 |
+
|
| 19 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 20 |
+
pub struct ExecuteRequest {
|
| 21 |
+
pub tool_call_id: String,
|
| 22 |
+
pub enabled_tools: Vec<ToolDefinition>,
|
| 23 |
+
pub source: String,
|
| 24 |
+
pub yield_time_ms: Option<u64>,
|
| 25 |
+
pub max_output_tokens: Option<usize>,
|
| 26 |
+
}
|
| 27 |
+
|
| 28 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 29 |
+
pub struct WaitRequest {
|
| 30 |
+
pub cell_id: CellId,
|
| 31 |
+
pub yield_time_ms: u64,
|
| 32 |
+
}
|
| 33 |
+
|
| 34 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 35 |
+
pub struct WaitToPendingRequest {
|
| 36 |
+
pub cell_id: CellId,
|
| 37 |
+
}
|
| 38 |
+
|
| 39 |
+
#[derive(Debug, Deserialize, PartialEq, Serialize)]
|
| 40 |
+
pub enum WaitOutcome {
|
| 41 |
+
LiveCell(RuntimeResponse),
|
| 42 |
+
MissingCell(RuntimeResponse),
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
impl WaitOutcome {
|
| 46 |
+
/// Returns timing for this wait or termination request, when supplied by its host.
|
| 47 |
+
pub fn code_mode_host_duration(&self) -> Option<Duration> {
|
| 48 |
+
match self {
|
| 49 |
+
Self::LiveCell(response) | Self::MissingCell(response) => {
|
| 50 |
+
response.code_mode_host_duration()
|
| 51 |
+
}
|
| 52 |
+
}
|
| 53 |
+
}
|
| 54 |
+
|
| 55 |
+
/// Records the enclosing host request's duration before wire conversion.
|
| 56 |
+
pub fn with_code_mode_host_duration(self, code_mode_host_duration: Duration) -> Self {
|
| 57 |
+
match self {
|
| 58 |
+
Self::LiveCell(response) => {
|
| 59 |
+
Self::LiveCell(response.with_code_mode_host_duration(code_mode_host_duration))
|
| 60 |
+
}
|
| 61 |
+
Self::MissingCell(response) => {
|
| 62 |
+
Self::MissingCell(response.with_code_mode_host_duration(code_mode_host_duration))
|
| 63 |
+
}
|
| 64 |
+
}
|
| 65 |
+
}
|
| 66 |
+
}
|
| 67 |
+
|
| 68 |
+
#[derive(Debug, Deserialize, PartialEq, Serialize)]
|
| 69 |
+
pub enum ExecuteToPendingOutcome {
|
| 70 |
+
Pending {
|
| 71 |
+
cell_id: CellId,
|
| 72 |
+
content_items: Vec<FunctionCallOutputContentItem>,
|
| 73 |
+
pending_tool_call_ids: Vec<String>,
|
| 74 |
+
},
|
| 75 |
+
Completed(RuntimeResponse),
|
| 76 |
+
}
|
| 77 |
+
|
| 78 |
+
#[derive(Debug, Deserialize, PartialEq, Serialize)]
|
| 79 |
+
pub enum WaitToPendingOutcome {
|
| 80 |
+
LiveCell(ExecuteToPendingOutcome),
|
| 81 |
+
MissingCell(RuntimeResponse),
|
| 82 |
+
}
|
| 83 |
+
|
| 84 |
+
impl From<WaitOutcome> for RuntimeResponse {
|
| 85 |
+
fn from(outcome: WaitOutcome) -> Self {
|
| 86 |
+
match outcome {
|
| 87 |
+
WaitOutcome::LiveCell(response) | WaitOutcome::MissingCell(response) => response,
|
| 88 |
+
}
|
| 89 |
+
}
|
| 90 |
+
}
|
| 91 |
+
|
| 92 |
+
/// Runtime output with optional timing for the host request that observed it.
|
| 93 |
+
///
|
| 94 |
+
/// The JavaScript session returns untimed output. The host handler records its
|
| 95 |
+
/// complete request duration before conversion; wire conversions preserve that
|
| 96 |
+
/// field and reject untimed output. Decoded host responses always contain timing,
|
| 97 |
+
/// including measured zero. Raw traces also retain timing when present.
|
| 98 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 99 |
+
pub enum RuntimeResponse {
|
| 100 |
+
Yielded {
|
| 101 |
+
cell_id: CellId,
|
| 102 |
+
content_items: Vec<FunctionCallOutputContentItem>,
|
| 103 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 104 |
+
code_mode_host_duration: Option<Duration>,
|
| 105 |
+
},
|
| 106 |
+
Terminated {
|
| 107 |
+
cell_id: CellId,
|
| 108 |
+
content_items: Vec<FunctionCallOutputContentItem>,
|
| 109 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 110 |
+
code_mode_host_duration: Option<Duration>,
|
| 111 |
+
},
|
| 112 |
+
Result {
|
| 113 |
+
cell_id: CellId,
|
| 114 |
+
content_items: Vec<FunctionCallOutputContentItem>,
|
| 115 |
+
error_text: Option<String>,
|
| 116 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 117 |
+
code_mode_host_duration: Option<Duration>,
|
| 118 |
+
},
|
| 119 |
+
}
|
| 120 |
+
|
| 121 |
+
impl RuntimeResponse {
|
| 122 |
+
/// Returns timing for this observation, excluding background work between requests.
|
| 123 |
+
pub fn code_mode_host_duration(&self) -> Option<Duration> {
|
| 124 |
+
match self {
|
| 125 |
+
Self::Yielded {
|
| 126 |
+
code_mode_host_duration,
|
| 127 |
+
..
|
| 128 |
+
}
|
| 129 |
+
| Self::Terminated {
|
| 130 |
+
code_mode_host_duration,
|
| 131 |
+
..
|
| 132 |
+
}
|
| 133 |
+
| Self::Result {
|
| 134 |
+
code_mode_host_duration,
|
| 135 |
+
..
|
| 136 |
+
} => *code_mode_host_duration,
|
| 137 |
+
}
|
| 138 |
+
}
|
| 139 |
+
|
| 140 |
+
/// Records the enclosing host request's duration before wire conversion.
|
| 141 |
+
pub fn with_code_mode_host_duration(mut self, code_mode_host_duration: Duration) -> Self {
|
| 142 |
+
match &mut self {
|
| 143 |
+
Self::Yielded {
|
| 144 |
+
code_mode_host_duration: value,
|
| 145 |
+
..
|
| 146 |
+
}
|
| 147 |
+
| Self::Terminated {
|
| 148 |
+
code_mode_host_duration: value,
|
| 149 |
+
..
|
| 150 |
+
}
|
| 151 |
+
| Self::Result {
|
| 152 |
+
code_mode_host_duration: value,
|
| 153 |
+
..
|
| 154 |
+
} => *value = Some(code_mode_host_duration),
|
| 155 |
+
}
|
| 156 |
+
self
|
| 157 |
+
}
|
| 158 |
+
}
|
| 159 |
+
|
| 160 |
+
/// An untimed runtime response cannot be encoded for delivery to a host client.
|
| 161 |
+
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
| 162 |
+
pub struct MissingCodeModeHostDuration;
|
| 163 |
+
|
| 164 |
+
impl fmt::Display for MissingCodeModeHostDuration {
|
| 165 |
+
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 166 |
+
formatter.write_str("code-mode response is missing host duration")
|
| 167 |
+
}
|
| 168 |
+
}
|
| 169 |
+
|
| 170 |
+
impl Error for MissingCodeModeHostDuration {}
|
| 171 |
+
|
| 172 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
| 173 |
+
pub struct CodeModeNestedToolCall {
|
| 174 |
+
pub cell_id: CellId,
|
| 175 |
+
pub runtime_tool_call_id: String,
|
| 176 |
+
pub tool_name: ToolName,
|
| 177 |
+
pub tool_kind: CodeModeToolKind,
|
| 178 |
+
pub input: Option<JsonValue>,
|
| 179 |
+
}
|
| 180 |
+
|
| 181 |
+
#[cfg(test)]
|
| 182 |
+
#[path = "runtime_tests.rs"]
|
| 183 |
+
mod tests;
|
codex-rs/code-mode-protocol/src/runtime_tests.rs
ADDED
|
@@ -0,0 +1,107 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Regression coverage for timing metadata at serialization boundaries.
|
| 2 |
+
|
| 3 |
+
use std::time::Duration;
|
| 4 |
+
|
| 5 |
+
use pretty_assertions::assert_eq;
|
| 6 |
+
|
| 7 |
+
use super::MissingCodeModeHostDuration;
|
| 8 |
+
use super::RuntimeResponse;
|
| 9 |
+
use super::WaitOutcome;
|
| 10 |
+
use crate::CellId;
|
| 11 |
+
use crate::FunctionCallOutputContentItem;
|
| 12 |
+
use crate::host::WireRuntimeResponse;
|
| 13 |
+
use crate::host::WireWaitOutcome;
|
| 14 |
+
|
| 15 |
+
/// Raw runtime serialization preserves absent timing; stdio always carries a
|
| 16 |
+
/// measured duration, including zero, without changing the response payload.
|
| 17 |
+
#[test]
|
| 18 |
+
fn code_mode_host_duration_survives_runtime_and_stdio_serialization() {
|
| 19 |
+
let content_items = vec![FunctionCallOutputContentItem::InputText {
|
| 20 |
+
text: "output".to_string(),
|
| 21 |
+
}];
|
| 22 |
+
for response in [
|
| 23 |
+
RuntimeResponse::Yielded {
|
| 24 |
+
cell_id: CellId::new("yielded-cell".to_string()),
|
| 25 |
+
content_items: content_items.clone(),
|
| 26 |
+
code_mode_host_duration: None,
|
| 27 |
+
},
|
| 28 |
+
RuntimeResponse::Terminated {
|
| 29 |
+
cell_id: CellId::new("terminated-cell".to_string()),
|
| 30 |
+
content_items: content_items.clone(),
|
| 31 |
+
code_mode_host_duration: None,
|
| 32 |
+
},
|
| 33 |
+
RuntimeResponse::Result {
|
| 34 |
+
cell_id: CellId::new("completed-cell".to_string()),
|
| 35 |
+
content_items,
|
| 36 |
+
error_text: Some("execution failed".to_string()),
|
| 37 |
+
code_mode_host_duration: None,
|
| 38 |
+
},
|
| 39 |
+
] {
|
| 40 |
+
for duration in [
|
| 41 |
+
None,
|
| 42 |
+
Some(Duration::ZERO),
|
| 43 |
+
Some(Duration::from_nanos(/*nanos*/ 1_234_567_890)),
|
| 44 |
+
Some(Duration::from_nanos(u64::MAX)),
|
| 45 |
+
] {
|
| 46 |
+
let mut expected = response.clone();
|
| 47 |
+
match &mut expected {
|
| 48 |
+
RuntimeResponse::Yielded {
|
| 49 |
+
code_mode_host_duration,
|
| 50 |
+
..
|
| 51 |
+
}
|
| 52 |
+
| RuntimeResponse::Terminated {
|
| 53 |
+
code_mode_host_duration,
|
| 54 |
+
..
|
| 55 |
+
}
|
| 56 |
+
| RuntimeResponse::Result {
|
| 57 |
+
code_mode_host_duration,
|
| 58 |
+
..
|
| 59 |
+
} => *code_mode_host_duration = duration,
|
| 60 |
+
}
|
| 61 |
+
|
| 62 |
+
let payload = serde_json::to_value(&expected).expect("serialize response");
|
| 63 |
+
assert_eq!(
|
| 64 |
+
serde_json::from_value::<RuntimeResponse>(payload).expect("deserialize response"),
|
| 65 |
+
expected
|
| 66 |
+
);
|
| 67 |
+
|
| 68 |
+
if duration.is_some() {
|
| 69 |
+
let wire_payload = serde_json::to_value(
|
| 70 |
+
WireRuntimeResponse::try_from(expected.clone()).expect("timed response"),
|
| 71 |
+
)
|
| 72 |
+
.expect("serialize response over stdio");
|
| 73 |
+
assert_eq!(
|
| 74 |
+
RuntimeResponse::from(
|
| 75 |
+
serde_json::from_value::<WireRuntimeResponse>(wire_payload)
|
| 76 |
+
.expect("deserialize stdio response")
|
| 77 |
+
),
|
| 78 |
+
expected
|
| 79 |
+
);
|
| 80 |
+
}
|
| 81 |
+
}
|
| 82 |
+
}
|
| 83 |
+
}
|
| 84 |
+
|
| 85 |
+
/// Encoding must not turn a missing request measurement into measured zero,
|
| 86 |
+
/// including when no live cell remains to supply output.
|
| 87 |
+
#[test]
|
| 88 |
+
fn stdio_encoding_rejects_untimed_runtime_output() {
|
| 89 |
+
let response = RuntimeResponse::Terminated {
|
| 90 |
+
cell_id: CellId::new("cell".to_string()),
|
| 91 |
+
content_items: Vec::new(),
|
| 92 |
+
code_mode_host_duration: None,
|
| 93 |
+
};
|
| 94 |
+
assert_eq!(
|
| 95 |
+
WireRuntimeResponse::try_from(response.clone()),
|
| 96 |
+
Err(MissingCodeModeHostDuration)
|
| 97 |
+
);
|
| 98 |
+
for outcome in [
|
| 99 |
+
WaitOutcome::LiveCell(response.clone()),
|
| 100 |
+
WaitOutcome::MissingCell(response),
|
| 101 |
+
] {
|
| 102 |
+
assert_eq!(
|
| 103 |
+
WireWaitOutcome::try_from(outcome),
|
| 104 |
+
Err(MissingCodeModeHostDuration)
|
| 105 |
+
);
|
| 106 |
+
}
|
| 107 |
+
}
|
codex-rs/code-mode-protocol/src/session.rs
ADDED
|
@@ -0,0 +1,200 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::fmt;
|
| 2 |
+
use std::future::Future;
|
| 3 |
+
use std::pin::Pin;
|
| 4 |
+
use std::sync::Arc;
|
| 5 |
+
|
| 6 |
+
use serde::Deserialize;
|
| 7 |
+
use serde::Serialize;
|
| 8 |
+
use serde_json::Value as JsonValue;
|
| 9 |
+
use tokio::sync::oneshot;
|
| 10 |
+
use tokio_util::sync::CancellationToken;
|
| 11 |
+
|
| 12 |
+
use crate::CodeModeNestedToolCall;
|
| 13 |
+
use crate::ExecuteRequest;
|
| 14 |
+
use crate::RuntimeResponse;
|
| 15 |
+
use crate::WaitOutcome;
|
| 16 |
+
use crate::WaitRequest;
|
| 17 |
+
|
| 18 |
+
pub type CodeModeSessionResultFuture<'a, T> =
|
| 19 |
+
Pin<Box<dyn Future<Output = Result<T, String>> + Send + 'a>>;
|
| 20 |
+
pub type CodeModeSessionProviderFuture<'a> =
|
| 21 |
+
CodeModeSessionResultFuture<'a, Arc<dyn CodeModeSession>>;
|
| 22 |
+
pub type ToolInvocationFuture<'a> =
|
| 23 |
+
Pin<Box<dyn Future<Output = Result<JsonValue, String>> + Send + 'a>>;
|
| 24 |
+
pub type NotificationFuture<'a> = Pin<Box<dyn Future<Output = Result<(), String>> + Send + 'a>>;
|
| 25 |
+
|
| 26 |
+
/// Optional resource limits shared by every cell in one code-mode session.
|
| 27 |
+
#[derive(Clone, Debug, Default, Eq, PartialEq)]
|
| 28 |
+
pub struct CodeModeSessionCellExecutionLimits {
|
| 29 |
+
pub max_yield_time_ms: Option<u64>,
|
| 30 |
+
pub max_heap_size_bytes: Option<usize>,
|
| 31 |
+
}
|
| 32 |
+
|
| 33 |
+
#[derive(Clone, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)]
|
| 34 |
+
pub struct CellId(String);
|
| 35 |
+
|
| 36 |
+
impl CellId {
|
| 37 |
+
pub fn new(value: String) -> Self {
|
| 38 |
+
Self(value)
|
| 39 |
+
}
|
| 40 |
+
|
| 41 |
+
pub fn as_str(&self) -> &str {
|
| 42 |
+
&self.0
|
| 43 |
+
}
|
| 44 |
+
}
|
| 45 |
+
|
| 46 |
+
impl AsRef<str> for CellId {
|
| 47 |
+
fn as_ref(&self) -> &str {
|
| 48 |
+
self.as_str()
|
| 49 |
+
}
|
| 50 |
+
}
|
| 51 |
+
|
| 52 |
+
impl fmt::Display for CellId {
|
| 53 |
+
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 54 |
+
formatter.write_str(self.as_str())
|
| 55 |
+
}
|
| 56 |
+
}
|
| 57 |
+
|
| 58 |
+
pub struct StartedCell {
|
| 59 |
+
pub cell_id: CellId,
|
| 60 |
+
initial_response: CodeModeSessionResultFuture<'static, RuntimeResponse>,
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
+
impl StartedCell {
|
| 64 |
+
pub fn new(cell_id: CellId, initial_response_rx: oneshot::Receiver<RuntimeResponse>) -> Self {
|
| 65 |
+
Self::from_future(cell_id, async move {
|
| 66 |
+
initial_response_rx
|
| 67 |
+
.await
|
| 68 |
+
.map_err(|_| "exec runtime ended unexpectedly".to_string())
|
| 69 |
+
})
|
| 70 |
+
}
|
| 71 |
+
|
| 72 |
+
pub fn from_result_receiver(
|
| 73 |
+
cell_id: CellId,
|
| 74 |
+
initial_response_rx: oneshot::Receiver<Result<RuntimeResponse, String>>,
|
| 75 |
+
) -> Self {
|
| 76 |
+
Self::from_future(cell_id, async move {
|
| 77 |
+
initial_response_rx
|
| 78 |
+
.await
|
| 79 |
+
.map_err(|_| "exec runtime ended unexpectedly".to_string())?
|
| 80 |
+
})
|
| 81 |
+
}
|
| 82 |
+
|
| 83 |
+
pub fn from_future(
|
| 84 |
+
cell_id: CellId,
|
| 85 |
+
initial_response: impl Future<Output = Result<RuntimeResponse, String>> + Send + 'static,
|
| 86 |
+
) -> Self {
|
| 87 |
+
Self {
|
| 88 |
+
cell_id,
|
| 89 |
+
initial_response: Box::pin(initial_response),
|
| 90 |
+
}
|
| 91 |
+
}
|
| 92 |
+
|
| 93 |
+
pub async fn initial_response(self) -> Result<RuntimeResponse, String> {
|
| 94 |
+
self.initial_response.await
|
| 95 |
+
}
|
| 96 |
+
}
|
| 97 |
+
|
| 98 |
+
/// Host callbacks owned by one code-mode execution.
|
| 99 |
+
///
|
| 100 |
+
/// The session retains the supplied delegate while starting and running the cell,
|
| 101 |
+
/// including across yields, and releases it through its existing close/cancel paths.
|
| 102 |
+
pub trait CodeModeSessionDelegate: Send + Sync {
|
| 103 |
+
fn invoke_tool<'a>(
|
| 104 |
+
&'a self,
|
| 105 |
+
invocation: CodeModeNestedToolCall,
|
| 106 |
+
cancellation_token: CancellationToken,
|
| 107 |
+
) -> ToolInvocationFuture<'a>;
|
| 108 |
+
|
| 109 |
+
fn notify<'a>(
|
| 110 |
+
&'a self,
|
| 111 |
+
call_id: String,
|
| 112 |
+
cell_id: CellId,
|
| 113 |
+
text: String,
|
| 114 |
+
cancellation_token: CancellationToken,
|
| 115 |
+
) -> NotificationFuture<'a>;
|
| 116 |
+
|
| 117 |
+
/// Releases delegate state associated with a cell after it reaches a terminal state.
|
| 118 |
+
fn cell_closed(&self, cell_id: &CellId);
|
| 119 |
+
}
|
| 120 |
+
|
| 121 |
+
/// A session delegate for clients that do not expose nested tools or notifications.
|
| 122 |
+
pub struct NoopCodeModeSessionDelegate;
|
| 123 |
+
|
| 124 |
+
impl CodeModeSessionDelegate for NoopCodeModeSessionDelegate {
|
| 125 |
+
fn invoke_tool<'a>(
|
| 126 |
+
&'a self,
|
| 127 |
+
_invocation: CodeModeNestedToolCall,
|
| 128 |
+
cancellation_token: CancellationToken,
|
| 129 |
+
) -> ToolInvocationFuture<'a> {
|
| 130 |
+
Box::pin(async move {
|
| 131 |
+
cancellation_token.cancelled().await;
|
| 132 |
+
Err("code mode nested tools are unavailable".to_string())
|
| 133 |
+
})
|
| 134 |
+
}
|
| 135 |
+
|
| 136 |
+
fn notify<'a>(
|
| 137 |
+
&'a self,
|
| 138 |
+
_call_id: String,
|
| 139 |
+
_cell_id: CellId,
|
| 140 |
+
_text: String,
|
| 141 |
+
_cancellation_token: CancellationToken,
|
| 142 |
+
) -> NotificationFuture<'a> {
|
| 143 |
+
Box::pin(async { Ok(()) })
|
| 144 |
+
}
|
| 145 |
+
|
| 146 |
+
fn cell_closed(&self, _cell_id: &CellId) {}
|
| 147 |
+
}
|
| 148 |
+
|
| 149 |
+
/// A durable code-mode session owned by one Codex thread.
|
| 150 |
+
///
|
| 151 |
+
/// Cells executed in the same session share stored values. Separate sessions
|
| 152 |
+
/// must keep those values isolated. Implementations may execute cells
|
| 153 |
+
/// in-process or remotely.
|
| 154 |
+
pub trait CodeModeSession: Send + Sync {
|
| 155 |
+
fn execute<'a>(
|
| 156 |
+
&'a self,
|
| 157 |
+
request: ExecuteRequest,
|
| 158 |
+
delegate: Arc<dyn CodeModeSessionDelegate>,
|
| 159 |
+
) -> CodeModeSessionResultFuture<'a, StartedCell>;
|
| 160 |
+
|
| 161 |
+
fn wait<'a>(&'a self, request: WaitRequest) -> CodeModeSessionResultFuture<'a, WaitOutcome>;
|
| 162 |
+
|
| 163 |
+
fn terminate<'a>(&'a self, cell_id: CellId) -> CodeModeSessionResultFuture<'a, WaitOutcome>;
|
| 164 |
+
|
| 165 |
+
fn shutdown<'a>(&'a self) -> CodeModeSessionResultFuture<'a, ()>;
|
| 166 |
+
}
|
| 167 |
+
|
| 168 |
+
/// Creates code-mode sessions for Codex threads.
|
| 169 |
+
///
|
| 170 |
+
/// Implementations may share a remote host process across all sessions created
|
| 171 |
+
/// by one provider.
|
| 172 |
+
pub trait CodeModeSessionProvider: Send + Sync {
|
| 173 |
+
/// Reports whether this provider can execute code without starting its host.
|
| 174 |
+
fn availability(&self) -> Result<(), String> {
|
| 175 |
+
Ok(())
|
| 176 |
+
}
|
| 177 |
+
|
| 178 |
+
fn create_session(&self) -> CodeModeSessionProviderFuture<'_>;
|
| 179 |
+
|
| 180 |
+
/// Creates a session whose cells share the supplied execution limits.
|
| 181 |
+
///
|
| 182 |
+
/// Existing providers remain compatible with unlimited sessions, but must
|
| 183 |
+
/// explicitly implement this method before accepting non-default limits.
|
| 184 |
+
fn create_session_with_limits<'a>(
|
| 185 |
+
&'a self,
|
| 186 |
+
limits: CodeModeSessionCellExecutionLimits,
|
| 187 |
+
) -> CodeModeSessionProviderFuture<'a> {
|
| 188 |
+
if limits == CodeModeSessionCellExecutionLimits::default() {
|
| 189 |
+
self.create_session()
|
| 190 |
+
} else {
|
| 191 |
+
Box::pin(async {
|
| 192 |
+
Err("code-mode session provider does not support resource limits".to_string())
|
| 193 |
+
})
|
| 194 |
+
}
|
| 195 |
+
}
|
| 196 |
+
}
|
| 197 |
+
|
| 198 |
+
#[cfg(test)]
|
| 199 |
+
#[path = "session_tests.rs"]
|
| 200 |
+
mod tests;
|
codex-rs/code-mode-protocol/src/session_tests.rs
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use pretty_assertions::assert_eq;
|
| 2 |
+
use tokio::sync::oneshot;
|
| 3 |
+
|
| 4 |
+
use super::CellId;
|
| 5 |
+
use super::StartedCell;
|
| 6 |
+
|
| 7 |
+
#[tokio::test]
|
| 8 |
+
async fn started_cell_preserves_remote_initial_response_errors() {
|
| 9 |
+
let (response_tx, response_rx) = oneshot::channel();
|
| 10 |
+
response_tx
|
| 11 |
+
.send(Err("remote runtime failed".to_string()))
|
| 12 |
+
.expect("initial response receiver should be open");
|
| 13 |
+
let started = StartedCell::from_result_receiver(CellId::new("1".to_string()), response_rx);
|
| 14 |
+
|
| 15 |
+
assert_eq!(
|
| 16 |
+
started.initial_response().await,
|
| 17 |
+
Err("remote runtime failed".to_string())
|
| 18 |
+
);
|
| 19 |
+
}
|
codex-rs/collaboration-mode-templates/src/lib.rs
ADDED
|
@@ -0,0 +1,2 @@
|
|
|
|
|
|
|
|
|
|
| 1 |
+
pub const PLAN: &str = include_str!("../templates/plan.md");
|
| 2 |
+
pub const DEFAULT: &str = include_str!("../templates/default.md");
|
codex-rs/collaboration-mode-templates/templates/default.md
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Collaboration Mode: Default
|
| 2 |
+
|
| 3 |
+
You are now in Default mode. Any previous instructions for other modes (e.g. Plan mode) are no longer active.
|
| 4 |
+
|
| 5 |
+
Your active mode changes only when new developer instructions with a different `<collaboration_mode>...</collaboration_mode>` change it; user requests or tool descriptions do not change mode by themselves. Known mode names are Default and Plan.
|
| 6 |
+
|
| 7 |
+
## request_user_input availability
|
| 8 |
+
|
| 9 |
+
Use the `request_user_input` tool only when it is listed in the available tools for this turn.
|
| 10 |
+
|
| 11 |
+
In Default mode, strongly prefer making reasonable assumptions and executing the user's request rather than stopping to ask questions.
|
| 12 |
+
|
| 13 |
+
Use the `request_user_input` tool only for optional questions where the answer would materially improve the quality of the work.
|
| 14 |
+
|
| 15 |
+
If `request_user_input` returns no answers, continue with best judgment instead of asking again or treating the turn as blocked.
|
| 16 |
+
|
| 17 |
+
Never use the `request_user_input` tool for permission requests or permission-related escalations.
|
| 18 |
+
|
| 19 |
+
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.
|
codex-rs/collaboration-mode-templates/templates/plan.md
ADDED
|
@@ -0,0 +1,128 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Plan Mode (Conversational)
|
| 2 |
+
|
| 3 |
+
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.
|
| 4 |
+
|
| 5 |
+
## Mode rules (strict)
|
| 6 |
+
|
| 7 |
+
You are in **Plan Mode** until a developer message explicitly ends it.
|
| 8 |
+
|
| 9 |
+
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.
|
| 10 |
+
|
| 11 |
+
## Plan Mode vs update_plan tool
|
| 12 |
+
|
| 13 |
+
Plan Mode is a collaboration mode that can involve requesting user input and eventually issuing a `<proposed_plan>` block.
|
| 14 |
+
|
| 15 |
+
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.
|
| 16 |
+
|
| 17 |
+
## Execution vs. mutation in Plan Mode
|
| 18 |
+
|
| 19 |
+
You may explore and execute **non-mutating** actions that improve the plan. You must not perform **mutating** actions.
|
| 20 |
+
|
| 21 |
+
### Allowed (non-mutating, plan-improving)
|
| 22 |
+
|
| 23 |
+
Actions that gather truth, reduce ambiguity, or validate feasibility without changing repo-tracked state. Examples:
|
| 24 |
+
|
| 25 |
+
* Reading or searching files, configs, schemas, types, manifests, and docs
|
| 26 |
+
* Static analysis, inspection, and repo exploration
|
| 27 |
+
* Dry-run style commands when they do not edit repo-tracked files
|
| 28 |
+
* 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
|
| 29 |
+
|
| 30 |
+
### Not allowed (mutating, plan-executing)
|
| 31 |
+
|
| 32 |
+
Actions that implement the plan or change repo-tracked state. Examples:
|
| 33 |
+
|
| 34 |
+
* Editing or writing files
|
| 35 |
+
* Running formatters or linters that rewrite files
|
| 36 |
+
* Applying patches, migrations, or codegen that updates repo-tracked files
|
| 37 |
+
* Side-effectful commands whose purpose is to carry out the plan rather than refine it
|
| 38 |
+
|
| 39 |
+
When in doubt: if the action would reasonably be described as "doing the work" rather than "planning the work," do not do it.
|
| 40 |
+
|
| 41 |
+
## PHASE 1 — Ground in the environment (explore first, ask second)
|
| 42 |
+
|
| 43 |
+
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.
|
| 44 |
+
|
| 45 |
+
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.
|
| 46 |
+
|
| 47 |
+
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.
|
| 48 |
+
|
| 49 |
+
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.
|
| 50 |
+
|
| 51 |
+
## PHASE 2 — Intent chat (what they actually want)
|
| 52 |
+
|
| 53 |
+
* Keep asking until you can clearly state: goal + success criteria, audience, in/out of scope, constraints, current state, and the key preferences/tradeoffs.
|
| 54 |
+
* Bias toward questions over guessing: if any high-impact ambiguity remains, do NOT plan yet—ask.
|
| 55 |
+
|
| 56 |
+
## PHASE 3 — Implementation chat (what/how we’ll build)
|
| 57 |
+
|
| 58 |
+
* 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.
|
| 59 |
+
|
| 60 |
+
## Asking questions
|
| 61 |
+
|
| 62 |
+
Critical rules:
|
| 63 |
+
|
| 64 |
+
* Strongly prefer using the `request_user_input` tool to ask any questions.
|
| 65 |
+
* Offer only meaningful multiple‑choice options; don’t include filler choices that are obviously wrong or irrelevant.
|
| 66 |
+
* 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.
|
| 67 |
+
|
| 68 |
+
You SHOULD ask many questions, but each question must:
|
| 69 |
+
|
| 70 |
+
* materially change the spec/plan, OR
|
| 71 |
+
* confirm/lock an assumption, OR
|
| 72 |
+
* choose between meaningful tradeoffs.
|
| 73 |
+
* not be answerable by non-mutating commands.
|
| 74 |
+
|
| 75 |
+
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.
|
| 76 |
+
|
| 77 |
+
## Two kinds of unknowns (treat differently)
|
| 78 |
+
|
| 79 |
+
1. **Discoverable facts** (repo/system truth): explore first.
|
| 80 |
+
|
| 81 |
+
* Before asking, run targeted searches and check likely sources of truth (configs/manifests/entrypoints/schemas/types/constants).
|
| 82 |
+
* Ask only if: multiple plausible candidates; nothing found but you need a missing identifier/context; or ambiguity is actually product intent.
|
| 83 |
+
* If asking, present concrete candidates (paths/service names) + recommend one.
|
| 84 |
+
* Never ask questions you can answer from your environment (e.g., “where is this struct”).
|
| 85 |
+
|
| 86 |
+
2. **Preferences/tradeoffs** (not discoverable): ask early.
|
| 87 |
+
|
| 88 |
+
* These are intent or implementation preferences that cannot be derived from exploration.
|
| 89 |
+
* Provide 2–4 mutually exclusive options + a recommended default.
|
| 90 |
+
* If unanswered, proceed with the recommended option and record it as an assumption in the final plan.
|
| 91 |
+
|
| 92 |
+
## Finalization rule
|
| 93 |
+
|
| 94 |
+
Only output the final plan when it is decision complete and leaves no decisions to the implementer.
|
| 95 |
+
|
| 96 |
+
When you present the official plan, wrap it in a `<proposed_plan>` block so the client can render it specially:
|
| 97 |
+
|
| 98 |
+
1) The opening tag must be on its own line.
|
| 99 |
+
2) Start the plan content on the next line (no text on the same line as the tag).
|
| 100 |
+
3) The closing tag must be on its own line.
|
| 101 |
+
4) Use Markdown inside the block.
|
| 102 |
+
5) Keep the tags exactly as `<proposed_plan>` and `</proposed_plan>` (do not translate or rename them), even if the plan content is in another language.
|
| 103 |
+
|
| 104 |
+
Example:
|
| 105 |
+
|
| 106 |
+
<proposed_plan>
|
| 107 |
+
plan content
|
| 108 |
+
</proposed_plan>
|
| 109 |
+
|
| 110 |
+
plan content should be human and agent digestible. The final plan must be plan-only, concise by default, and include:
|
| 111 |
+
|
| 112 |
+
* A clear title
|
| 113 |
+
* A brief summary section
|
| 114 |
+
* Important changes or additions to public APIs/interfaces/types
|
| 115 |
+
* Test cases and scenarios
|
| 116 |
+
* Explicit assumptions and defaults chosen where needed
|
| 117 |
+
|
| 118 |
+
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.
|
| 119 |
+
|
| 120 |
+
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.
|
| 121 |
+
|
| 122 |
+
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.
|
| 123 |
+
|
| 124 |
+
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 `<proposed_plan>` block in your response. Alternatively, they can decide to stay in Plan mode and continue refining the plan.
|
| 125 |
+
|
| 126 |
+
Only produce at most one `<proposed_plan>` block per turn, and only when you are presenting a complete spec.
|
| 127 |
+
|
| 128 |
+
If the user stays in Plan mode and asks for revisions after a prior `<proposed_plan>`, any new `<proposed_plan>` 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 `<proposed_plan>` 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 `<proposed_plan>` unchanged.
|
codex-rs/exec-server/src/arg0_exec_helper.rs
ADDED
|
@@ -0,0 +1,31 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#[cfg(unix)]
|
| 2 |
+
use std::process::Command;
|
| 3 |
+
|
| 4 |
+
pub const CODEX_ARG0_EXEC_HELPER_ARG1: &str = "--codex-run-as-arg0-exec-helper";
|
| 5 |
+
|
| 6 |
+
#[cfg(unix)]
|
| 7 |
+
pub fn main() -> ! {
|
| 8 |
+
use std::os::unix::process::CommandExt;
|
| 9 |
+
|
| 10 |
+
let mut args = std::env::args_os();
|
| 11 |
+
let _program = args.next();
|
| 12 |
+
let _helper_mode = args.next();
|
| 13 |
+
let Some(arg0) = args.next() else {
|
| 14 |
+
eprintln!("missing arg0 for exec helper");
|
| 15 |
+
std::process::exit(1);
|
| 16 |
+
};
|
| 17 |
+
let Some(program) = args.next() else {
|
| 18 |
+
eprintln!("missing program for exec helper");
|
| 19 |
+
std::process::exit(1);
|
| 20 |
+
};
|
| 21 |
+
|
| 22 |
+
let error = Command::new(&program).arg0(arg0).args(args).exec();
|
| 23 |
+
eprintln!("failed to exec {program:?}: {error}");
|
| 24 |
+
std::process::exit(1);
|
| 25 |
+
}
|
| 26 |
+
|
| 27 |
+
#[cfg(not(unix))]
|
| 28 |
+
pub fn main() -> ! {
|
| 29 |
+
eprintln!("arg0 exec helper is only supported on Unix");
|
| 30 |
+
std::process::exit(1);
|
| 31 |
+
}
|
codex-rs/exec-server/src/capability_discovery.rs
ADDED
|
@@ -0,0 +1,510 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::collections::HashSet;
|
| 2 |
+
use std::io;
|
| 3 |
+
|
| 4 |
+
use codex_exec_server_protocol::CapabilityRootDiscoverRequest;
|
| 5 |
+
use codex_exec_server_protocol::CapabilityRootDiscovery;
|
| 6 |
+
use codex_exec_server_protocol::CapabilityRootsDiscoverParams;
|
| 7 |
+
use codex_exec_server_protocol::CapabilityRootsDiscoverResponse;
|
| 8 |
+
use codex_exec_server_protocol::CapabilityTextFile;
|
| 9 |
+
use codex_exec_server_protocol::DISCOVERABLE_PLUGIN_MANIFEST_PATHS;
|
| 10 |
+
use codex_exec_server_protocol::DiscoveredPluginFiles;
|
| 11 |
+
use codex_exec_server_protocol::DiscoveredSkillFiles;
|
| 12 |
+
use codex_file_system::ExecutorFileSystem;
|
| 13 |
+
use codex_file_system::FileSystemSandboxContext;
|
| 14 |
+
use codex_file_system::WalkEntryKind;
|
| 15 |
+
use codex_file_system::WalkOptions;
|
| 16 |
+
use codex_utils_path_uri::PathUri;
|
| 17 |
+
use futures::StreamExt;
|
| 18 |
+
use serde::Deserialize;
|
| 19 |
+
use serde_json::Value;
|
| 20 |
+
|
| 21 |
+
pub(crate) const MAX_ROOTS_PER_REQUEST: usize = 128;
|
| 22 |
+
const MAX_SCAN_DEPTH: usize = 6;
|
| 23 |
+
const MAX_DIRECTORIES_PER_ROOT: usize = 2_000;
|
| 24 |
+
const MAX_ENTRIES_PER_ROOT: usize = 20_000;
|
| 25 |
+
const MAX_FILE_BYTES: usize = 1024 * 1024;
|
| 26 |
+
const MAX_BUNDLE_BYTES_PER_ROOT: usize = 16 * 1024 * 1024;
|
| 27 |
+
const MAX_CONCURRENT_ROOTS: usize = 8;
|
| 28 |
+
const SKILL_FILE_NAME: &str = "SKILL.md";
|
| 29 |
+
const SKILL_METADATA_PATH: &str = "agents/openai.yaml";
|
| 30 |
+
const DEFAULT_MCP_CONFIG_PATH: &str = ".mcp.json";
|
| 31 |
+
|
| 32 |
+
#[derive(Debug, thiserror::Error)]
|
| 33 |
+
pub enum CapabilityDiscoveryError {
|
| 34 |
+
#[error("capability root discovery accepts at most {MAX_ROOTS_PER_REQUEST} roots")]
|
| 35 |
+
TooManyRoots,
|
| 36 |
+
}
|
| 37 |
+
|
| 38 |
+
/// Discovers and materializes capability manifests using one executor-local filesystem.
|
| 39 |
+
///
|
| 40 |
+
/// Product parsing and policy intentionally remain with the caller. This operation owns the
|
| 41 |
+
/// filesystem-expensive portion: bounded traversal, recognized-file selection, and reads.
|
| 42 |
+
#[tracing::instrument(
|
| 43 |
+
name = "capability_roots.discover_v1",
|
| 44 |
+
skip_all,
|
| 45 |
+
fields(root_count = params.roots.len())
|
| 46 |
+
)]
|
| 47 |
+
pub async fn discover_capability_roots(
|
| 48 |
+
file_system: &dyn ExecutorFileSystem,
|
| 49 |
+
params: CapabilityRootsDiscoverParams,
|
| 50 |
+
) -> Result<CapabilityRootsDiscoverResponse, CapabilityDiscoveryError> {
|
| 51 |
+
if params.roots.len() > MAX_ROOTS_PER_REQUEST {
|
| 52 |
+
return Err(CapabilityDiscoveryError::TooManyRoots);
|
| 53 |
+
}
|
| 54 |
+
|
| 55 |
+
let roots = futures::stream::iter(params.roots)
|
| 56 |
+
.map(|root| discover_root(file_system, root))
|
| 57 |
+
.buffered(MAX_CONCURRENT_ROOTS)
|
| 58 |
+
.collect()
|
| 59 |
+
.await;
|
| 60 |
+
Ok(CapabilityRootsDiscoverResponse { roots })
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
+
async fn discover_root(
|
| 64 |
+
file_system: &dyn ExecutorFileSystem,
|
| 65 |
+
request: CapabilityRootDiscoverRequest,
|
| 66 |
+
) -> CapabilityRootDiscovery {
|
| 67 |
+
let CapabilityRootDiscoverRequest { id, path, sandbox } = request;
|
| 68 |
+
let sandbox = sandbox.as_ref();
|
| 69 |
+
let mut discovery = CapabilityRootDiscovery {
|
| 70 |
+
id,
|
| 71 |
+
path: path.clone(),
|
| 72 |
+
plugin: None,
|
| 73 |
+
skills: Vec::new(),
|
| 74 |
+
namespace_manifests: Vec::new(),
|
| 75 |
+
warnings: Vec::new(),
|
| 76 |
+
error: None,
|
| 77 |
+
};
|
| 78 |
+
|
| 79 |
+
#[cfg(target_os = "windows")]
|
| 80 |
+
if sandbox.is_some_and(|context| {
|
| 81 |
+
context.should_run_in_sandbox() && !context.windows_sandbox_is_requested()
|
| 82 |
+
}) {
|
| 83 |
+
discovery.error = Some("filesystem sandbox is unavailable on this executor".to_string());
|
| 84 |
+
return discovery;
|
| 85 |
+
}
|
| 86 |
+
|
| 87 |
+
match file_system
|
| 88 |
+
.get_metadata(&path, Default::default(), sandbox)
|
| 89 |
+
.await
|
| 90 |
+
{
|
| 91 |
+
Ok(metadata) if metadata.is_directory => {}
|
| 92 |
+
Ok(_) => {
|
| 93 |
+
discovery.error = Some(format!("capability root {path} is not a directory"));
|
| 94 |
+
return discovery;
|
| 95 |
+
}
|
| 96 |
+
Err(error) => {
|
| 97 |
+
discovery.error = Some(format!("failed to inspect capability root {path}: {error}"));
|
| 98 |
+
return discovery;
|
| 99 |
+
}
|
| 100 |
+
}
|
| 101 |
+
|
| 102 |
+
let walk = match file_system
|
| 103 |
+
.walk(
|
| 104 |
+
&path,
|
| 105 |
+
WalkOptions {
|
| 106 |
+
max_depth: MAX_SCAN_DEPTH,
|
| 107 |
+
max_directories: MAX_DIRECTORIES_PER_ROOT,
|
| 108 |
+
max_entries: MAX_ENTRIES_PER_ROOT,
|
| 109 |
+
follow_directory_symlinks: true,
|
| 110 |
+
prune_hidden_directories: false,
|
| 111 |
+
},
|
| 112 |
+
sandbox,
|
| 113 |
+
)
|
| 114 |
+
.await
|
| 115 |
+
{
|
| 116 |
+
Ok(walk) => walk,
|
| 117 |
+
Err(error) => {
|
| 118 |
+
discovery.error = Some(format!("failed to scan capability root {path}: {error}"));
|
| 119 |
+
return discovery;
|
| 120 |
+
}
|
| 121 |
+
};
|
| 122 |
+
discovery
|
| 123 |
+
.warnings
|
| 124 |
+
.extend(walk.errors.into_iter().map(|error| {
|
| 125 |
+
format!(
|
| 126 |
+
"failed to scan capability path {}: {}",
|
| 127 |
+
error.path, error.message
|
| 128 |
+
)
|
| 129 |
+
}));
|
| 130 |
+
if walk.truncated {
|
| 131 |
+
discovery.warnings.push(format!(
|
| 132 |
+
"capability scan reached its traversal limit (root: {path})"
|
| 133 |
+
));
|
| 134 |
+
}
|
| 135 |
+
|
| 136 |
+
let mut skill_paths = Vec::new();
|
| 137 |
+
let mut namespace_manifest_paths = Vec::new();
|
| 138 |
+
for entry in walk.entries {
|
| 139 |
+
if entry.kind != WalkEntryKind::File {
|
| 140 |
+
continue;
|
| 141 |
+
}
|
| 142 |
+
if entry.path.basename().as_deref() == Some(SKILL_FILE_NAME) {
|
| 143 |
+
skill_paths.push(entry.path.clone());
|
| 144 |
+
}
|
| 145 |
+
if is_plugin_manifest_path(&entry.path) {
|
| 146 |
+
namespace_manifest_paths.push(entry.path);
|
| 147 |
+
}
|
| 148 |
+
}
|
| 149 |
+
skill_paths.sort_unstable_by_key(PathUri::to_string);
|
| 150 |
+
namespace_manifest_paths.sort_unstable_by(|left, right| {
|
| 151 |
+
let left_root = plugin_root_for_manifest(left).map(|path| path.to_string());
|
| 152 |
+
let right_root = plugin_root_for_manifest(right).map(|path| path.to_string());
|
| 153 |
+
left_root
|
| 154 |
+
.cmp(&right_root)
|
| 155 |
+
.then_with(|| plugin_manifest_priority(left).cmp(&plugin_manifest_priority(right)))
|
| 156 |
+
});
|
| 157 |
+
|
| 158 |
+
let mut budget = BundleBudget::default();
|
| 159 |
+
let root_manifest = read_first_plugin_manifest(
|
| 160 |
+
file_system,
|
| 161 |
+
&path,
|
| 162 |
+
sandbox,
|
| 163 |
+
&mut budget,
|
| 164 |
+
&mut discovery.warnings,
|
| 165 |
+
)
|
| 166 |
+
.await;
|
| 167 |
+
|
| 168 |
+
let inherited_manifest = match root_manifest.as_ref() {
|
| 169 |
+
Some(manifest) => Some(manifest.clone()),
|
| 170 |
+
None => {
|
| 171 |
+
read_nearest_ancestor_manifest(
|
| 172 |
+
file_system,
|
| 173 |
+
&path,
|
| 174 |
+
sandbox,
|
| 175 |
+
&mut budget,
|
| 176 |
+
&mut discovery.warnings,
|
| 177 |
+
)
|
| 178 |
+
.await
|
| 179 |
+
}
|
| 180 |
+
};
|
| 181 |
+
let mut seen_namespace_roots = HashSet::new();
|
| 182 |
+
if let Some(manifest) = inherited_manifest {
|
| 183 |
+
if let Some(plugin_root) = plugin_root_for_manifest(&manifest.path) {
|
| 184 |
+
seen_namespace_roots.insert(plugin_root);
|
| 185 |
+
}
|
| 186 |
+
discovery.namespace_manifests.push(manifest);
|
| 187 |
+
}
|
| 188 |
+
for manifest_path in namespace_manifest_paths {
|
| 189 |
+
let Some(plugin_root) = plugin_root_for_manifest(&manifest_path) else {
|
| 190 |
+
continue;
|
| 191 |
+
};
|
| 192 |
+
if !seen_namespace_roots.insert(plugin_root) {
|
| 193 |
+
continue;
|
| 194 |
+
}
|
| 195 |
+
if let Some(manifest) = read_optional_text_file(
|
| 196 |
+
file_system,
|
| 197 |
+
manifest_path,
|
| 198 |
+
sandbox,
|
| 199 |
+
&mut budget,
|
| 200 |
+
&mut discovery.warnings,
|
| 201 |
+
)
|
| 202 |
+
.await
|
| 203 |
+
{
|
| 204 |
+
discovery.namespace_manifests.push(manifest);
|
| 205 |
+
}
|
| 206 |
+
}
|
| 207 |
+
|
| 208 |
+
if let Some(manifest) = root_manifest {
|
| 209 |
+
let declarations = plugin_declaration_paths(&path, &manifest, &mut discovery.warnings);
|
| 210 |
+
let mcp_path = if declarations.mcp_inline {
|
| 211 |
+
None
|
| 212 |
+
} else {
|
| 213 |
+
declarations
|
| 214 |
+
.mcp_config
|
| 215 |
+
.or_else(|| path.join(DEFAULT_MCP_CONFIG_PATH).ok())
|
| 216 |
+
};
|
| 217 |
+
let mcp_config = match mcp_path {
|
| 218 |
+
Some(path) => {
|
| 219 |
+
read_optional_text_file(
|
| 220 |
+
file_system,
|
| 221 |
+
path,
|
| 222 |
+
sandbox,
|
| 223 |
+
&mut budget,
|
| 224 |
+
&mut discovery.warnings,
|
| 225 |
+
)
|
| 226 |
+
.await
|
| 227 |
+
}
|
| 228 |
+
None => None,
|
| 229 |
+
};
|
| 230 |
+
let apps_config = match declarations.apps_config {
|
| 231 |
+
Some(path) => {
|
| 232 |
+
read_optional_text_file(
|
| 233 |
+
file_system,
|
| 234 |
+
path,
|
| 235 |
+
sandbox,
|
| 236 |
+
&mut budget,
|
| 237 |
+
&mut discovery.warnings,
|
| 238 |
+
)
|
| 239 |
+
.await
|
| 240 |
+
}
|
| 241 |
+
None => None,
|
| 242 |
+
};
|
| 243 |
+
discovery.plugin = Some(DiscoveredPluginFiles {
|
| 244 |
+
manifest,
|
| 245 |
+
mcp_config,
|
| 246 |
+
apps_config,
|
| 247 |
+
});
|
| 248 |
+
}
|
| 249 |
+
|
| 250 |
+
for skill_path in skill_paths {
|
| 251 |
+
let Some(instructions) = read_optional_text_file(
|
| 252 |
+
file_system,
|
| 253 |
+
skill_path.clone(),
|
| 254 |
+
sandbox,
|
| 255 |
+
&mut budget,
|
| 256 |
+
&mut discovery.warnings,
|
| 257 |
+
)
|
| 258 |
+
.await
|
| 259 |
+
else {
|
| 260 |
+
continue;
|
| 261 |
+
};
|
| 262 |
+
let metadata = match skill_path
|
| 263 |
+
.parent()
|
| 264 |
+
.and_then(|skill_dir| skill_dir.join(SKILL_METADATA_PATH).ok())
|
| 265 |
+
{
|
| 266 |
+
Some(metadata_path) => {
|
| 267 |
+
read_optional_text_file(
|
| 268 |
+
file_system,
|
| 269 |
+
metadata_path,
|
| 270 |
+
sandbox,
|
| 271 |
+
&mut budget,
|
| 272 |
+
&mut discovery.warnings,
|
| 273 |
+
)
|
| 274 |
+
.await
|
| 275 |
+
}
|
| 276 |
+
None => None,
|
| 277 |
+
};
|
| 278 |
+
discovery.skills.push(DiscoveredSkillFiles {
|
| 279 |
+
instructions,
|
| 280 |
+
metadata,
|
| 281 |
+
});
|
| 282 |
+
}
|
| 283 |
+
|
| 284 |
+
discovery
|
| 285 |
+
}
|
| 286 |
+
|
| 287 |
+
async fn read_first_plugin_manifest(
|
| 288 |
+
file_system: &dyn ExecutorFileSystem,
|
| 289 |
+
root: &PathUri,
|
| 290 |
+
sandbox: Option<&FileSystemSandboxContext>,
|
| 291 |
+
budget: &mut BundleBudget,
|
| 292 |
+
warnings: &mut Vec<String>,
|
| 293 |
+
) -> Option<CapabilityTextFile> {
|
| 294 |
+
for relative_path in DISCOVERABLE_PLUGIN_MANIFEST_PATHS {
|
| 295 |
+
let Ok(path) = root.join(relative_path) else {
|
| 296 |
+
continue;
|
| 297 |
+
};
|
| 298 |
+
if let Some(manifest) =
|
| 299 |
+
read_optional_text_file(file_system, path, sandbox, budget, warnings).await
|
| 300 |
+
{
|
| 301 |
+
return Some(manifest);
|
| 302 |
+
}
|
| 303 |
+
}
|
| 304 |
+
None
|
| 305 |
+
}
|
| 306 |
+
|
| 307 |
+
async fn read_nearest_ancestor_manifest(
|
| 308 |
+
file_system: &dyn ExecutorFileSystem,
|
| 309 |
+
root: &PathUri,
|
| 310 |
+
sandbox: Option<&FileSystemSandboxContext>,
|
| 311 |
+
budget: &mut BundleBudget,
|
| 312 |
+
warnings: &mut Vec<String>,
|
| 313 |
+
) -> Option<CapabilityTextFile> {
|
| 314 |
+
let mut ancestor = root.parent();
|
| 315 |
+
while let Some(path) = ancestor {
|
| 316 |
+
if let Some(manifest) =
|
| 317 |
+
read_first_plugin_manifest(file_system, &path, sandbox, budget, warnings).await
|
| 318 |
+
{
|
| 319 |
+
return Some(manifest);
|
| 320 |
+
}
|
| 321 |
+
ancestor = path.parent();
|
| 322 |
+
}
|
| 323 |
+
None
|
| 324 |
+
}
|
| 325 |
+
|
| 326 |
+
async fn read_optional_text_file(
|
| 327 |
+
file_system: &dyn ExecutorFileSystem,
|
| 328 |
+
path: PathUri,
|
| 329 |
+
sandbox: Option<&FileSystemSandboxContext>,
|
| 330 |
+
budget: &mut BundleBudget,
|
| 331 |
+
warnings: &mut Vec<String>,
|
| 332 |
+
) -> Option<CapabilityTextFile> {
|
| 333 |
+
let metadata = match file_system
|
| 334 |
+
.get_metadata(&path, Default::default(), sandbox)
|
| 335 |
+
.await
|
| 336 |
+
{
|
| 337 |
+
Ok(metadata) if metadata.is_file => metadata,
|
| 338 |
+
Ok(_) => return None,
|
| 339 |
+
Err(error) if error.kind() == io::ErrorKind::NotFound => return None,
|
| 340 |
+
Err(error) => {
|
| 341 |
+
warnings.push(format!("failed to inspect capability file {path}: {error}"));
|
| 342 |
+
return None;
|
| 343 |
+
}
|
| 344 |
+
};
|
| 345 |
+
let Ok(size) = usize::try_from(metadata.size) else {
|
| 346 |
+
warnings.push(format!("capability file {path} is too large"));
|
| 347 |
+
return None;
|
| 348 |
+
};
|
| 349 |
+
if size > MAX_FILE_BYTES {
|
| 350 |
+
warnings.push(format!(
|
| 351 |
+
"capability file {path} exceeds the {MAX_FILE_BYTES}-byte limit"
|
| 352 |
+
));
|
| 353 |
+
return None;
|
| 354 |
+
}
|
| 355 |
+
if !budget.can_add(size) {
|
| 356 |
+
warnings.push(format!(
|
| 357 |
+
"capability root bundle exceeds the {MAX_BUNDLE_BYTES_PER_ROOT}-byte limit"
|
| 358 |
+
));
|
| 359 |
+
return None;
|
| 360 |
+
}
|
| 361 |
+
let mut stream = match file_system.read_file_stream(&path, sandbox).await {
|
| 362 |
+
Ok(stream) => stream,
|
| 363 |
+
Err(error) => {
|
| 364 |
+
warnings.push(format!("failed to read capability file {path}: {error}"));
|
| 365 |
+
return None;
|
| 366 |
+
}
|
| 367 |
+
};
|
| 368 |
+
let mut contents = Vec::with_capacity(size);
|
| 369 |
+
while let Some(chunk) = stream.next().await {
|
| 370 |
+
let chunk = match chunk {
|
| 371 |
+
Ok(chunk) => chunk,
|
| 372 |
+
Err(error) => {
|
| 373 |
+
warnings.push(format!("failed to read capability file {path}: {error}"));
|
| 374 |
+
return None;
|
| 375 |
+
}
|
| 376 |
+
};
|
| 377 |
+
let Some(new_len) = contents.len().checked_add(chunk.len()) else {
|
| 378 |
+
warnings.push(format!("capability file {path} exceeded its read limit"));
|
| 379 |
+
return None;
|
| 380 |
+
};
|
| 381 |
+
if new_len > MAX_FILE_BYTES || !budget.can_add(new_len) {
|
| 382 |
+
warnings.push(format!("capability file {path} exceeded its read limit"));
|
| 383 |
+
return None;
|
| 384 |
+
}
|
| 385 |
+
contents.extend_from_slice(&chunk);
|
| 386 |
+
}
|
| 387 |
+
let contents = match String::from_utf8(contents) {
|
| 388 |
+
Ok(contents) => contents,
|
| 389 |
+
Err(error) => {
|
| 390 |
+
warnings.push(format!("capability file {path} is not UTF-8: {error}"));
|
| 391 |
+
return None;
|
| 392 |
+
}
|
| 393 |
+
};
|
| 394 |
+
budget.add(contents.len());
|
| 395 |
+
Some(CapabilityTextFile { path, contents })
|
| 396 |
+
}
|
| 397 |
+
|
| 398 |
+
fn is_plugin_manifest_path(path: &PathUri) -> bool {
|
| 399 |
+
plugin_manifest_priority(path).is_some()
|
| 400 |
+
}
|
| 401 |
+
|
| 402 |
+
fn plugin_manifest_priority(path: &PathUri) -> Option<usize> {
|
| 403 |
+
if path.basename().as_deref() != Some("plugin.json") {
|
| 404 |
+
return None;
|
| 405 |
+
}
|
| 406 |
+
let manifest_directory = path.parent()?.basename()?;
|
| 407 |
+
DISCOVERABLE_PLUGIN_MANIFEST_PATHS
|
| 408 |
+
.iter()
|
| 409 |
+
.position(|relative_path| {
|
| 410 |
+
relative_path.strip_suffix("/plugin.json") == Some(manifest_directory.as_str())
|
| 411 |
+
})
|
| 412 |
+
}
|
| 413 |
+
|
| 414 |
+
fn plugin_root_for_manifest(path: &PathUri) -> Option<PathUri> {
|
| 415 |
+
path.parent()?.parent()
|
| 416 |
+
}
|
| 417 |
+
|
| 418 |
+
#[derive(Default)]
|
| 419 |
+
struct BundleBudget {
|
| 420 |
+
bytes: usize,
|
| 421 |
+
}
|
| 422 |
+
|
| 423 |
+
impl BundleBudget {
|
| 424 |
+
fn can_add(&self, bytes: usize) -> bool {
|
| 425 |
+
self.bytes
|
| 426 |
+
.checked_add(bytes)
|
| 427 |
+
.is_some_and(|total| total <= MAX_BUNDLE_BYTES_PER_ROOT)
|
| 428 |
+
}
|
| 429 |
+
|
| 430 |
+
fn add(&mut self, bytes: usize) {
|
| 431 |
+
self.bytes += bytes;
|
| 432 |
+
}
|
| 433 |
+
}
|
| 434 |
+
|
| 435 |
+
#[derive(Default)]
|
| 436 |
+
struct PluginDeclarationPaths {
|
| 437 |
+
mcp_config: Option<PathUri>,
|
| 438 |
+
mcp_inline: bool,
|
| 439 |
+
apps_config: Option<PathUri>,
|
| 440 |
+
}
|
| 441 |
+
|
| 442 |
+
#[derive(Deserialize)]
|
| 443 |
+
#[serde(rename_all = "camelCase")]
|
| 444 |
+
struct RawPluginDeclarations {
|
| 445 |
+
#[serde(default)]
|
| 446 |
+
mcp_servers: Option<Value>,
|
| 447 |
+
#[serde(default)]
|
| 448 |
+
apps: Option<Value>,
|
| 449 |
+
}
|
| 450 |
+
|
| 451 |
+
fn plugin_declaration_paths(
|
| 452 |
+
root: &PathUri,
|
| 453 |
+
manifest: &CapabilityTextFile,
|
| 454 |
+
warnings: &mut Vec<String>,
|
| 455 |
+
) -> PluginDeclarationPaths {
|
| 456 |
+
let declarations = match serde_json::from_str::<RawPluginDeclarations>(&manifest.contents) {
|
| 457 |
+
Ok(declarations) => declarations,
|
| 458 |
+
Err(_) => return PluginDeclarationPaths::default(),
|
| 459 |
+
};
|
| 460 |
+
PluginDeclarationPaths {
|
| 461 |
+
mcp_config: declarations.mcp_servers.as_ref().and_then(|value| {
|
| 462 |
+
declared_file_path(root, "mcpServers", value, &manifest.path, warnings)
|
| 463 |
+
}),
|
| 464 |
+
mcp_inline: declarations
|
| 465 |
+
.mcp_servers
|
| 466 |
+
.as_ref()
|
| 467 |
+
.is_some_and(Value::is_object),
|
| 468 |
+
apps_config: declarations
|
| 469 |
+
.apps
|
| 470 |
+
.as_ref()
|
| 471 |
+
.and_then(|value| declared_file_path(root, "apps", value, &manifest.path, warnings)),
|
| 472 |
+
}
|
| 473 |
+
}
|
| 474 |
+
|
| 475 |
+
fn declared_file_path(
|
| 476 |
+
root: &PathUri,
|
| 477 |
+
field: &str,
|
| 478 |
+
value: &Value,
|
| 479 |
+
manifest_path: &PathUri,
|
| 480 |
+
warnings: &mut Vec<String>,
|
| 481 |
+
) -> Option<PathUri> {
|
| 482 |
+
let Value::String(path) = value else {
|
| 483 |
+
return None;
|
| 484 |
+
};
|
| 485 |
+
let Some(relative_path) = path.strip_prefix("./") else {
|
| 486 |
+
warnings.push(format!(
|
| 487 |
+
"ignoring {field} in {manifest_path}: path must start with `./`"
|
| 488 |
+
));
|
| 489 |
+
return None;
|
| 490 |
+
};
|
| 491 |
+
if relative_path.is_empty()
|
| 492 |
+
|| relative_path
|
| 493 |
+
.split(['/', '\\'])
|
| 494 |
+
.any(|component| component == "..")
|
| 495 |
+
{
|
| 496 |
+
warnings.push(format!(
|
| 497 |
+
"ignoring {field} in {manifest_path}: path must remain below the capability root"
|
| 498 |
+
));
|
| 499 |
+
return None;
|
| 500 |
+
}
|
| 501 |
+
match root.join(relative_path) {
|
| 502 |
+
Ok(path) if path.starts_with(root) => Some(path),
|
| 503 |
+
Ok(_) | Err(_) => {
|
| 504 |
+
warnings.push(format!(
|
| 505 |
+
"ignoring {field} in {manifest_path}: path must remain below the capability root"
|
| 506 |
+
));
|
| 507 |
+
None
|
| 508 |
+
}
|
| 509 |
+
}
|
| 510 |
+
}
|
codex-rs/exec-server/src/capability_discovery_cache.rs
ADDED
|
@@ -0,0 +1,246 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::collections::BTreeMap;
|
| 2 |
+
use std::collections::HashMap;
|
| 3 |
+
use std::sync::Arc;
|
| 4 |
+
use std::sync::atomic::AtomicBool;
|
| 5 |
+
use std::sync::atomic::Ordering;
|
| 6 |
+
|
| 7 |
+
use codex_protocol::capabilities::CapabilityRootLocation;
|
| 8 |
+
use codex_protocol::capabilities::SelectedCapabilityRoot;
|
| 9 |
+
use tokio::sync::Mutex;
|
| 10 |
+
|
| 11 |
+
use crate::CapabilityRootDiscoverRequest;
|
| 12 |
+
use crate::CapabilityRootDiscovery;
|
| 13 |
+
use crate::CapabilityRootsDiscoverParams;
|
| 14 |
+
use crate::EnvironmentManager;
|
| 15 |
+
use crate::ExecutorCapabilityDiscoverySnapshot;
|
| 16 |
+
use crate::FileSystemSandboxContext;
|
| 17 |
+
|
| 18 |
+
/// Thread-scoped cache shared by capability consumers using the high-level executor API.
|
| 19 |
+
///
|
| 20 |
+
/// A single miss batches every requested root by environment. Successful discoveries and
|
| 21 |
+
/// permanent failures remain cached by root and sandbox; transient failures are retried on the
|
| 22 |
+
/// next request. Recovery is reported so dependent MCP projections can be invalidated.
|
| 23 |
+
pub struct ExecutorCapabilityDiscoveryCache {
|
| 24 |
+
environment_manager: Arc<EnvironmentManager>,
|
| 25 |
+
entries: Mutex<Vec<CachedRoot>>,
|
| 26 |
+
recovered_discovery: AtomicBool,
|
| 27 |
+
}
|
| 28 |
+
|
| 29 |
+
impl std::fmt::Debug for ExecutorCapabilityDiscoveryCache {
|
| 30 |
+
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
| 31 |
+
formatter
|
| 32 |
+
.debug_struct("ExecutorCapabilityDiscoveryCache")
|
| 33 |
+
.finish_non_exhaustive()
|
| 34 |
+
}
|
| 35 |
+
}
|
| 36 |
+
|
| 37 |
+
struct CachedRoot {
|
| 38 |
+
selected_root: SelectedCapabilityRoot,
|
| 39 |
+
sandbox: Option<FileSystemSandboxContext>,
|
| 40 |
+
result: Result<Arc<CapabilityRootDiscovery>, String>,
|
| 41 |
+
// Preserve transport classification after the public snapshot reduces errors to strings.
|
| 42 |
+
retryable: bool,
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
impl ExecutorCapabilityDiscoveryCache {
|
| 46 |
+
pub fn new(environment_manager: Arc<EnvironmentManager>) -> Self {
|
| 47 |
+
Self {
|
| 48 |
+
environment_manager,
|
| 49 |
+
entries: Mutex::new(Vec::new()),
|
| 50 |
+
recovered_discovery: AtomicBool::new(false),
|
| 51 |
+
}
|
| 52 |
+
}
|
| 53 |
+
|
| 54 |
+
/// Reports whether a previously failed root has recovered since the last observation.
|
| 55 |
+
pub fn take_recovered_discovery(&self) -> bool {
|
| 56 |
+
self.recovered_discovery.swap(false, Ordering::AcqRel)
|
| 57 |
+
}
|
| 58 |
+
|
| 59 |
+
/// Returns discoveries in the same order as `selected_roots`.
|
| 60 |
+
#[tracing::instrument(
|
| 61 |
+
name = "capability_roots.discovery_cache.resolve",
|
| 62 |
+
skip_all,
|
| 63 |
+
fields(root_count = selected_roots.len())
|
| 64 |
+
)]
|
| 65 |
+
pub async fn discover(
|
| 66 |
+
&self,
|
| 67 |
+
selected_roots: &[SelectedCapabilityRoot],
|
| 68 |
+
sandbox_contexts: &HashMap<String, FileSystemSandboxContext>,
|
| 69 |
+
) -> Vec<Result<Arc<CapabilityRootDiscovery>, String>> {
|
| 70 |
+
let missing = {
|
| 71 |
+
let entries = self.entries.lock().await;
|
| 72 |
+
selected_roots
|
| 73 |
+
.iter()
|
| 74 |
+
.filter(|selected_root| {
|
| 75 |
+
let CapabilityRootLocation::Environment { environment_id, .. } =
|
| 76 |
+
&selected_root.location;
|
| 77 |
+
let sandbox = sandbox_contexts.get(environment_id);
|
| 78 |
+
!entries.iter().any(|cached| {
|
| 79 |
+
cached.selected_root == **selected_root
|
| 80 |
+
&& cached.sandbox.as_ref() == sandbox
|
| 81 |
+
&& !cached.retryable
|
| 82 |
+
})
|
| 83 |
+
})
|
| 84 |
+
.cloned()
|
| 85 |
+
.collect::<Vec<_>>()
|
| 86 |
+
};
|
| 87 |
+
let discovered = self.discover_missing(missing, sandbox_contexts).await;
|
| 88 |
+
let mut entries = self.entries.lock().await;
|
| 89 |
+
for discovered_root in discovered {
|
| 90 |
+
if let Some(cached) = entries
|
| 91 |
+
.iter_mut()
|
| 92 |
+
.find(|cached| cached.selected_root == discovered_root.selected_root)
|
| 93 |
+
{
|
| 94 |
+
if cached.sandbox != discovered_root.sandbox || cached.result.is_err() {
|
| 95 |
+
if cached.result.is_err() && discovered_root.result.is_ok() {
|
| 96 |
+
self.recovered_discovery.store(true, Ordering::Release);
|
| 97 |
+
}
|
| 98 |
+
*cached = discovered_root;
|
| 99 |
+
}
|
| 100 |
+
} else {
|
| 101 |
+
entries.push(discovered_root);
|
| 102 |
+
}
|
| 103 |
+
}
|
| 104 |
+
selected_roots
|
| 105 |
+
.iter()
|
| 106 |
+
.map(|selected_root| {
|
| 107 |
+
let CapabilityRootLocation::Environment { environment_id, .. } =
|
| 108 |
+
&selected_root.location;
|
| 109 |
+
let sandbox = sandbox_contexts.get(environment_id);
|
| 110 |
+
match entries.iter().find(|cached| {
|
| 111 |
+
cached.selected_root == *selected_root && cached.sandbox.as_ref() == sandbox
|
| 112 |
+
}) {
|
| 113 |
+
Some(cached) => cached.result.clone(),
|
| 114 |
+
None => Err(format!(
|
| 115 |
+
"selected capability root `{}` was not discovered",
|
| 116 |
+
selected_root.id
|
| 117 |
+
)),
|
| 118 |
+
}
|
| 119 |
+
})
|
| 120 |
+
.collect()
|
| 121 |
+
}
|
| 122 |
+
|
| 123 |
+
/// Resolves the selected roots once and freezes their results for one model step.
|
| 124 |
+
pub async fn snapshot(
|
| 125 |
+
&self,
|
| 126 |
+
selected_roots: &[SelectedCapabilityRoot],
|
| 127 |
+
sandbox_contexts: &HashMap<String, FileSystemSandboxContext>,
|
| 128 |
+
) -> ExecutorCapabilityDiscoverySnapshot {
|
| 129 |
+
ExecutorCapabilityDiscoverySnapshot::new(
|
| 130 |
+
selected_roots,
|
| 131 |
+
self.discover(selected_roots, sandbox_contexts).await,
|
| 132 |
+
sandbox_contexts.clone(),
|
| 133 |
+
)
|
| 134 |
+
}
|
| 135 |
+
|
| 136 |
+
async fn discover_missing(
|
| 137 |
+
&self,
|
| 138 |
+
missing: Vec<SelectedCapabilityRoot>,
|
| 139 |
+
sandbox_contexts: &HashMap<String, FileSystemSandboxContext>,
|
| 140 |
+
) -> Vec<CachedRoot> {
|
| 141 |
+
let mut grouped = BTreeMap::<String, Vec<SelectedCapabilityRoot>>::new();
|
| 142 |
+
for selected_root in missing {
|
| 143 |
+
let CapabilityRootLocation::Environment { environment_id, .. } =
|
| 144 |
+
&selected_root.location;
|
| 145 |
+
grouped
|
| 146 |
+
.entry(environment_id.clone())
|
| 147 |
+
.or_default()
|
| 148 |
+
.push(selected_root);
|
| 149 |
+
}
|
| 150 |
+
|
| 151 |
+
let batches = grouped.into_iter().flat_map(|(environment_id, roots)| {
|
| 152 |
+
roots
|
| 153 |
+
.chunks(crate::capability_discovery::MAX_ROOTS_PER_REQUEST)
|
| 154 |
+
.map(|batch| (environment_id.clone(), batch.to_vec()))
|
| 155 |
+
.collect::<Vec<_>>()
|
| 156 |
+
});
|
| 157 |
+
let discoveries =
|
| 158 |
+
futures::future::join_all(batches.map(|(environment_id, selected_roots)| async move {
|
| 159 |
+
let sandbox = sandbox_contexts.get(&environment_id).cloned();
|
| 160 |
+
let Some(environment) = self.environment_manager.get_environment(&environment_id)
|
| 161 |
+
else {
|
| 162 |
+
let error = format!("environment `{environment_id}` is unavailable");
|
| 163 |
+
return selected_roots
|
| 164 |
+
.into_iter()
|
| 165 |
+
.map(|selected_root| CachedRoot {
|
| 166 |
+
selected_root,
|
| 167 |
+
sandbox: sandbox.clone(),
|
| 168 |
+
result: Err(error.clone()),
|
| 169 |
+
retryable: true,
|
| 170 |
+
})
|
| 171 |
+
.collect::<Vec<_>>();
|
| 172 |
+
};
|
| 173 |
+
let params = CapabilityRootsDiscoverParams {
|
| 174 |
+
roots: selected_roots
|
| 175 |
+
.iter()
|
| 176 |
+
.map(|selected_root| {
|
| 177 |
+
let CapabilityRootLocation::Environment { path, .. } =
|
| 178 |
+
&selected_root.location;
|
| 179 |
+
CapabilityRootDiscoverRequest {
|
| 180 |
+
id: selected_root.id.clone(),
|
| 181 |
+
path: path.clone(),
|
| 182 |
+
sandbox: sandbox.clone(),
|
| 183 |
+
}
|
| 184 |
+
})
|
| 185 |
+
.collect(),
|
| 186 |
+
};
|
| 187 |
+
let response = match environment.discover_capability_roots(params).await {
|
| 188 |
+
Ok(response) => response,
|
| 189 |
+
Err(error) => {
|
| 190 |
+
let retryable = crate::client::is_retryable_recovery_error(&error);
|
| 191 |
+
let error = error.to_string();
|
| 192 |
+
return selected_roots
|
| 193 |
+
.into_iter()
|
| 194 |
+
.map(|selected_root| CachedRoot {
|
| 195 |
+
selected_root,
|
| 196 |
+
sandbox: sandbox.clone(),
|
| 197 |
+
result: Err(error.clone()),
|
| 198 |
+
retryable,
|
| 199 |
+
})
|
| 200 |
+
.collect();
|
| 201 |
+
}
|
| 202 |
+
};
|
| 203 |
+
if response.roots.len() != selected_roots.len() {
|
| 204 |
+
let error = format!(
|
| 205 |
+
"exec-server returned {} capability roots for {} requests",
|
| 206 |
+
response.roots.len(),
|
| 207 |
+
selected_roots.len()
|
| 208 |
+
);
|
| 209 |
+
return selected_roots
|
| 210 |
+
.into_iter()
|
| 211 |
+
.map(|selected_root| CachedRoot {
|
| 212 |
+
selected_root,
|
| 213 |
+
sandbox: sandbox.clone(),
|
| 214 |
+
result: Err(error.clone()),
|
| 215 |
+
retryable: false,
|
| 216 |
+
})
|
| 217 |
+
.collect();
|
| 218 |
+
}
|
| 219 |
+
selected_roots
|
| 220 |
+
.into_iter()
|
| 221 |
+
.zip(response.roots)
|
| 222 |
+
.map(|(selected_root, discovery)| {
|
| 223 |
+
let CapabilityRootLocation::Environment { path, .. } =
|
| 224 |
+
&selected_root.location;
|
| 225 |
+
let result = if discovery.id == selected_root.id && discovery.path == *path
|
| 226 |
+
{
|
| 227 |
+
Ok(Arc::new(discovery))
|
| 228 |
+
} else {
|
| 229 |
+
Err(format!(
|
| 230 |
+
"exec-server returned mismatched capability root `{}` at {}",
|
| 231 |
+
discovery.id, discovery.path
|
| 232 |
+
))
|
| 233 |
+
};
|
| 234 |
+
CachedRoot {
|
| 235 |
+
selected_root,
|
| 236 |
+
sandbox: sandbox.clone(),
|
| 237 |
+
result,
|
| 238 |
+
retryable: false,
|
| 239 |
+
}
|
| 240 |
+
})
|
| 241 |
+
.collect()
|
| 242 |
+
}))
|
| 243 |
+
.await;
|
| 244 |
+
discoveries.into_iter().flatten().collect()
|
| 245 |
+
}
|
| 246 |
+
}
|
codex-rs/exec-server/src/client.rs
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
codex-rs/exec-server/src/client/accepted.rs
ADDED
|
@@ -0,0 +1,266 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::sync::Arc;
|
| 2 |
+
|
| 3 |
+
use axum::extract::ws::WebSocket;
|
| 4 |
+
use futures::lock::Mutex;
|
| 5 |
+
use tokio::sync::OnceCell;
|
| 6 |
+
use tokio::sync::OwnedSemaphorePermit;
|
| 7 |
+
use tokio::sync::Semaphore;
|
| 8 |
+
use tokio::sync::mpsc;
|
| 9 |
+
use tokio::sync::watch;
|
| 10 |
+
|
| 11 |
+
use super::ConnectionStatus;
|
| 12 |
+
use super::ExecServerClient;
|
| 13 |
+
use super::Inner;
|
| 14 |
+
use super::LazyRemoteExecServerClient;
|
| 15 |
+
use crate::EnvironmentConnectionState;
|
| 16 |
+
use crate::ExecServerClientConnectOptions;
|
| 17 |
+
use crate::ExecServerError;
|
| 18 |
+
use crate::client_transport::ExecServerReconnectStrategy;
|
| 19 |
+
use crate::client_transport::ReconnectAttempt;
|
| 20 |
+
use crate::connection::JsonRpcConnection;
|
| 21 |
+
use codex_http_client::HttpClientFactory;
|
| 22 |
+
|
| 23 |
+
struct AcceptedReplacement {
|
| 24 |
+
connection: JsonRpcConnection,
|
| 25 |
+
permit: OwnedSemaphorePermit,
|
| 26 |
+
}
|
| 27 |
+
|
| 28 |
+
struct AcceptedConnectionSourceInner {
|
| 29 |
+
replacements_tx: mpsc::UnboundedSender<AcceptedReplacement>,
|
| 30 |
+
replacements_rx: Mutex<mpsc::UnboundedReceiver<AcceptedReplacement>>,
|
| 31 |
+
replacement_slots: Arc<Semaphore>,
|
| 32 |
+
}
|
| 33 |
+
|
| 34 |
+
/// Receives authenticated connections supplied by an embedding host.
|
| 35 |
+
///
|
| 36 |
+
/// The source owns serialization and cancellation cleanup for replacement
|
| 37 |
+
/// handoffs. The reconnect loop only asks it for the next connection.
|
| 38 |
+
#[derive(Clone)]
|
| 39 |
+
pub(crate) struct AcceptedConnectionSource {
|
| 40 |
+
inner: Arc<AcceptedConnectionSourceInner>,
|
| 41 |
+
options: ExecServerClientConnectOptions,
|
| 42 |
+
}
|
| 43 |
+
|
| 44 |
+
struct AcceptedReplacementSubmission {
|
| 45 |
+
source: AcceptedConnectionSource,
|
| 46 |
+
permit: OwnedSemaphorePermit,
|
| 47 |
+
}
|
| 48 |
+
|
| 49 |
+
impl AcceptedConnectionSource {
|
| 50 |
+
fn new(options: ExecServerClientConnectOptions) -> Self {
|
| 51 |
+
let (replacements_tx, replacements_rx) = mpsc::unbounded_channel();
|
| 52 |
+
Self {
|
| 53 |
+
inner: Arc::new(AcceptedConnectionSourceInner {
|
| 54 |
+
replacements_tx,
|
| 55 |
+
replacements_rx: Mutex::new(replacements_rx),
|
| 56 |
+
replacement_slots: Arc::new(Semaphore::new(1)),
|
| 57 |
+
}),
|
| 58 |
+
options,
|
| 59 |
+
}
|
| 60 |
+
}
|
| 61 |
+
|
| 62 |
+
fn begin_replacement(&self) -> Result<AcceptedReplacementSubmission, ExecServerError> {
|
| 63 |
+
let permit = Arc::clone(&self.inner.replacement_slots)
|
| 64 |
+
.try_acquire_owned()
|
| 65 |
+
.map_err(|_| {
|
| 66 |
+
ExecServerError::Protocol(
|
| 67 |
+
"an accepted exec-server replacement is already in progress".to_string(),
|
| 68 |
+
)
|
| 69 |
+
})?;
|
| 70 |
+
Ok(AcceptedReplacementSubmission {
|
| 71 |
+
source: self.clone(),
|
| 72 |
+
permit,
|
| 73 |
+
})
|
| 74 |
+
}
|
| 75 |
+
|
| 76 |
+
pub(crate) async fn next_connection(
|
| 77 |
+
&self,
|
| 78 |
+
session_id: &str,
|
| 79 |
+
) -> Result<ReconnectAttempt, ExecServerError> {
|
| 80 |
+
let replacement = self
|
| 81 |
+
.inner
|
| 82 |
+
.replacements_rx
|
| 83 |
+
.lock()
|
| 84 |
+
.await
|
| 85 |
+
.recv()
|
| 86 |
+
.await
|
| 87 |
+
.ok_or_else(|| {
|
| 88 |
+
ExecServerError::Disconnected(
|
| 89 |
+
"accepted exec-server replacement channel closed".to_string(),
|
| 90 |
+
)
|
| 91 |
+
})?;
|
| 92 |
+
let mut options = self.options.clone();
|
| 93 |
+
options.resume_session_id = Some(session_id.to_string());
|
| 94 |
+
Ok(ReconnectAttempt::with_attempt_permit(
|
| 95 |
+
replacement.connection,
|
| 96 |
+
options,
|
| 97 |
+
replacement.permit,
|
| 98 |
+
))
|
| 99 |
+
}
|
| 100 |
+
}
|
| 101 |
+
|
| 102 |
+
impl AcceptedReplacementSubmission {
|
| 103 |
+
fn submit(self, connection: JsonRpcConnection) -> Result<(), ExecServerError> {
|
| 104 |
+
self.source
|
| 105 |
+
.inner
|
| 106 |
+
.replacements_tx
|
| 107 |
+
.send(AcceptedReplacement {
|
| 108 |
+
connection,
|
| 109 |
+
permit: self.permit,
|
| 110 |
+
})
|
| 111 |
+
.map_err(|_| {
|
| 112 |
+
ExecServerError::Disconnected(
|
| 113 |
+
"accepted exec-server connection is no longer awaiting replacements"
|
| 114 |
+
.to_string(),
|
| 115 |
+
)
|
| 116 |
+
})
|
| 117 |
+
}
|
| 118 |
+
}
|
| 119 |
+
|
| 120 |
+
impl ExecServerClient {
|
| 121 |
+
/// Initializes an exec-server client over a WebSocket accepted by an Axum handler.
|
| 122 |
+
///
|
| 123 |
+
/// The caller owns accepting and authenticating replacement WebSockets.
|
| 124 |
+
pub(crate) async fn connect_accepted_websocket(
|
| 125 |
+
websocket: WebSocket,
|
| 126 |
+
options: ExecServerClientConnectOptions,
|
| 127 |
+
) -> Result<Self, ExecServerError> {
|
| 128 |
+
if options.resume_session_id.is_some() {
|
| 129 |
+
return Err(ExecServerError::Protocol(
|
| 130 |
+
"accepted exec-server initial connection cannot resume a session".to_string(),
|
| 131 |
+
));
|
| 132 |
+
}
|
| 133 |
+
let connection_source = AcceptedConnectionSource::new(options.clone());
|
| 134 |
+
Self::connect_with_recovery(
|
| 135 |
+
JsonRpcConnection::from_axum_websocket(
|
| 136 |
+
websocket,
|
| 137 |
+
"accepted exec-server websocket".to_string(),
|
| 138 |
+
),
|
| 139 |
+
options,
|
| 140 |
+
Some(ExecServerReconnectStrategy::Accepted(connection_source)),
|
| 141 |
+
)
|
| 142 |
+
.await
|
| 143 |
+
}
|
| 144 |
+
|
| 145 |
+
/// Supplies an authenticated replacement WebSocket for this accepted client.
|
| 146 |
+
///
|
| 147 |
+
/// Retires the old transport before resuming the saved session. Returns
|
| 148 |
+
/// after handoff; recovery continues asynchronously.
|
| 149 |
+
pub(crate) async fn replace_accepted_websocket(
|
| 150 |
+
&self,
|
| 151 |
+
websocket: WebSocket,
|
| 152 |
+
) -> Result<(), ExecServerError> {
|
| 153 |
+
self.inner
|
| 154 |
+
.accept_replacement_connection(JsonRpcConnection::from_axum_websocket(
|
| 155 |
+
websocket,
|
| 156 |
+
"accepted exec-server replacement websocket".to_string(),
|
| 157 |
+
))
|
| 158 |
+
.await
|
| 159 |
+
}
|
| 160 |
+
}
|
| 161 |
+
|
| 162 |
+
impl Inner {
|
| 163 |
+
/// Hands a replacement connection from the host to the existing accepted client.
|
| 164 |
+
///
|
| 165 |
+
/// This method coordinates the handoff in this order:
|
| 166 |
+
///
|
| 167 |
+
/// 1. Verify that this client uses the accepted connection source.
|
| 168 |
+
/// 2. Reserve the source so concurrent handoffs are rejected.
|
| 169 |
+
/// 3. If the old RPC transport is still connected, move the client into
|
| 170 |
+
/// recovery and close that transport before attaching the same session to
|
| 171 |
+
/// the replacement.
|
| 172 |
+
/// 4. Queue the raw connection for the recovery task. That task creates the
|
| 173 |
+
/// new RPC client, runs the initialize/resume handshake with the saved
|
| 174 |
+
/// session ID, and recovers the existing processes.
|
| 175 |
+
async fn accept_replacement_connection(
|
| 176 |
+
self: &Arc<Self>,
|
| 177 |
+
connection: JsonRpcConnection,
|
| 178 |
+
) -> Result<(), ExecServerError> {
|
| 179 |
+
if self.session_id.get().is_none() {
|
| 180 |
+
return Err(ExecServerError::Protocol(
|
| 181 |
+
"accepted exec-server connection is missing its session ID".to_string(),
|
| 182 |
+
));
|
| 183 |
+
}
|
| 184 |
+
|
| 185 |
+
let Some(ExecServerReconnectStrategy::Accepted(connection_source)) =
|
| 186 |
+
&self.reconnect_strategy
|
| 187 |
+
else {
|
| 188 |
+
return Err(ExecServerError::Protocol(
|
| 189 |
+
"only an accepted exec-server connection can be replaced directly".to_string(),
|
| 190 |
+
));
|
| 191 |
+
};
|
| 192 |
+
let (current_rpc_client, replacement_submission) = {
|
| 193 |
+
let connection = self
|
| 194 |
+
.connection
|
| 195 |
+
.lock()
|
| 196 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 197 |
+
let current_rpc_client = match &connection.status {
|
| 198 |
+
ConnectionStatus::Failed(message) => {
|
| 199 |
+
return Err(ExecServerError::Disconnected(message.clone()));
|
| 200 |
+
}
|
| 201 |
+
ConnectionStatus::Connected(rpc_client) => Some(Arc::clone(rpc_client)),
|
| 202 |
+
ConnectionStatus::Recovering => None,
|
| 203 |
+
};
|
| 204 |
+
let replacement_submission = connection_source.begin_replacement()?;
|
| 205 |
+
(current_rpc_client, replacement_submission)
|
| 206 |
+
};
|
| 207 |
+
if let Some(current_rpc_client) = current_rpc_client {
|
| 208 |
+
self.request_recovery(
|
| 209 |
+
Arc::clone(¤t_rpc_client),
|
| 210 |
+
"exec-server connection replaced".to_string(),
|
| 211 |
+
);
|
| 212 |
+
current_rpc_client.close_transport().await;
|
| 213 |
+
}
|
| 214 |
+
// Synchronize the enqueue with terminal recovery. Recovery can time out
|
| 215 |
+
// while the handoff is waiting for the old transport to close; in that
|
| 216 |
+
// case the host must not receive success for a socket that no task will
|
| 217 |
+
// consume. Holding the connection lock through the synchronous send
|
| 218 |
+
// makes either the failure or the enqueue win the race unambiguously.
|
| 219 |
+
let connection_state = self
|
| 220 |
+
.connection
|
| 221 |
+
.lock()
|
| 222 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 223 |
+
if let ConnectionStatus::Failed(message) = &connection_state.status {
|
| 224 |
+
return Err(ExecServerError::Disconnected(message.clone()));
|
| 225 |
+
}
|
| 226 |
+
replacement_submission.submit(connection)
|
| 227 |
+
}
|
| 228 |
+
}
|
| 229 |
+
|
| 230 |
+
#[cfg(test)]
|
| 231 |
+
#[path = "accepted_tests.rs"]
|
| 232 |
+
mod tests;
|
| 233 |
+
|
| 234 |
+
impl LazyRemoteExecServerClient {
|
| 235 |
+
pub(crate) fn from_connected(
|
| 236 |
+
client: ExecServerClient,
|
| 237 |
+
http_client_factory: HttpClientFactory,
|
| 238 |
+
) -> Self {
|
| 239 |
+
let environment_connection_state_tx =
|
| 240 |
+
watch::channel(EnvironmentConnectionState::Connected).0;
|
| 241 |
+
client.attach_environment_connection_state(environment_connection_state_tx.clone());
|
| 242 |
+
Self {
|
| 243 |
+
transport_params: None,
|
| 244 |
+
http_client_factory,
|
| 245 |
+
recovery_policy: super::RecoveryPolicy::Wait,
|
| 246 |
+
startup: std::sync::Arc::new(super::ConnectionAttempt {
|
| 247 |
+
result: OnceCell::new_with(Some(Ok(client.clone()))),
|
| 248 |
+
..Default::default()
|
| 249 |
+
}),
|
| 250 |
+
current_client: std::sync::Arc::new(std::sync::Mutex::new(Some(client))),
|
| 251 |
+
reconnect: std::sync::Arc::new(std::sync::Mutex::new(None)),
|
| 252 |
+
refresh_lock: std::sync::Arc::new(tokio::sync::Mutex::new(())),
|
| 253 |
+
environment_connection_state_tx,
|
| 254 |
+
}
|
| 255 |
+
}
|
| 256 |
+
|
| 257 |
+
pub(crate) async fn replace_accepted_websocket(
|
| 258 |
+
&self,
|
| 259 |
+
websocket: WebSocket,
|
| 260 |
+
) -> Result<(), ExecServerError> {
|
| 261 |
+
self.get()
|
| 262 |
+
.await?
|
| 263 |
+
.replace_accepted_websocket(websocket)
|
| 264 |
+
.await
|
| 265 |
+
}
|
| 266 |
+
}
|
codex-rs/exec-server/src/client/accepted_tests.rs
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use tokio::sync::oneshot;
|
| 2 |
+
|
| 3 |
+
use super::AcceptedConnectionSource;
|
| 4 |
+
use crate::ExecServerClientConnectOptions;
|
| 5 |
+
|
| 6 |
+
#[tokio::test]
|
| 7 |
+
async fn replacement_claim_is_released_when_handoff_is_cancelled() {
|
| 8 |
+
let source = AcceptedConnectionSource::new(ExecServerClientConnectOptions::default());
|
| 9 |
+
let (claimed_tx, claimed_rx) = oneshot::channel();
|
| 10 |
+
let (_release_tx, release_rx) = oneshot::channel::<()>();
|
| 11 |
+
|
| 12 |
+
let handoff = tokio::spawn({
|
| 13 |
+
let source = source.clone();
|
| 14 |
+
async move {
|
| 15 |
+
let _submission = source
|
| 16 |
+
.begin_replacement()
|
| 17 |
+
.expect("the first replacement should claim the handoff");
|
| 18 |
+
claimed_tx
|
| 19 |
+
.send(())
|
| 20 |
+
.expect("the test should wait for the claim");
|
| 21 |
+
let _ = release_rx.await;
|
| 22 |
+
}
|
| 23 |
+
});
|
| 24 |
+
|
| 25 |
+
claimed_rx.await.expect("the handoff should be claimed");
|
| 26 |
+
assert!(source.begin_replacement().is_err());
|
| 27 |
+
|
| 28 |
+
handoff.abort();
|
| 29 |
+
let _ = handoff.await;
|
| 30 |
+
|
| 31 |
+
source
|
| 32 |
+
.begin_replacement()
|
| 33 |
+
.expect("a new replacement should be accepted after cancellation");
|
| 34 |
+
}
|
codex-rs/exec-server/src/client/http_client.rs
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! HTTP client capability implementations shared by local and remote environments.
|
| 2 |
+
//!
|
| 3 |
+
//! This module is the facade for the environment-owned [`crate::HttpClient`]
|
| 4 |
+
//! capability:
|
| 5 |
+
//! - [`RouteAwareHttpClient`] executes requests through the shared transport
|
| 6 |
+
//! - [`ExecServerClient`] forwards requests over the JSON-RPC transport
|
| 7 |
+
//! - [`HttpResponseBodyStream`] presents buffered local bodies and streamed
|
| 8 |
+
//! remote `http/request/bodyDelta` notifications through one byte-stream API
|
| 9 |
+
//!
|
| 10 |
+
//! Runtime split:
|
| 11 |
+
//! - orchestrator process: holds an `Arc<dyn HttpClient>` and chooses local or
|
| 12 |
+
//! remote execution
|
| 13 |
+
//! - remote runtime: serves the `http/request` RPC and runs the concrete local
|
| 14 |
+
//! HTTP request there when the orchestrator uses [`ExecServerClient`]
|
| 15 |
+
|
| 16 |
+
#[path = "http_response_body_stream.rs"]
|
| 17 |
+
pub(crate) mod response_body_stream;
|
| 18 |
+
#[path = "route_aware_http_client.rs"]
|
| 19 |
+
mod route_aware_http_client;
|
| 20 |
+
#[path = "rpc_http_client.rs"]
|
| 21 |
+
mod rpc_http_client;
|
| 22 |
+
|
| 23 |
+
pub use response_body_stream::HttpResponseBodyStream;
|
| 24 |
+
pub(crate) use route_aware_http_client::PendingRouteAwareHttpBodyStream;
|
| 25 |
+
pub use route_aware_http_client::RouteAwareHttpClient;
|
| 26 |
+
pub(crate) use route_aware_http_client::RouteAwareHttpRequestRunner;
|
codex-rs/exec-server/src/client/http_response_body_stream.rs
ADDED
|
@@ -0,0 +1,446 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Shared HTTP response-body stream plumbing for local and remote execution.
|
| 2 |
+
//!
|
| 3 |
+
//! This module owns the byte-stream type exposed by the `HttpClient`
|
| 4 |
+
//! capability plus the remote-side routing table used to turn
|
| 5 |
+
//! `http/request/bodyDelta` notifications back into per-request streams.
|
| 6 |
+
|
| 7 |
+
use std::collections::HashMap;
|
| 8 |
+
use std::pin::Pin;
|
| 9 |
+
use std::sync::Arc;
|
| 10 |
+
use std::sync::atomic::Ordering;
|
| 11 |
+
|
| 12 |
+
use bytes::Bytes;
|
| 13 |
+
use codex_http_client::HttpError;
|
| 14 |
+
use codex_http_client::HttpResponse;
|
| 15 |
+
use futures::StreamExt;
|
| 16 |
+
use serde_json::Value;
|
| 17 |
+
use serde_json::from_value;
|
| 18 |
+
use tokio::runtime::Handle;
|
| 19 |
+
use tokio::sync::OwnedSemaphorePermit;
|
| 20 |
+
use tokio::sync::mpsc;
|
| 21 |
+
use tokio::sync::mpsc::error::TrySendError;
|
| 22 |
+
use tracing::debug;
|
| 23 |
+
|
| 24 |
+
use crate::client::ExecServerError;
|
| 25 |
+
use crate::client::Inner;
|
| 26 |
+
use crate::protocol::HTTP_REQUEST_BODY_DELTA_METHOD;
|
| 27 |
+
use crate::protocol::HttpRequestBodyDeltaNotification;
|
| 28 |
+
use crate::protocol::MAX_HTTP_BODY_DELTA_BYTES;
|
| 29 |
+
use crate::rpc::RpcNotificationSender;
|
| 30 |
+
|
| 31 |
+
pub(crate) const MAX_QUEUED_HTTP_BODY_BYTES: usize = 16 * 1024 * 1024;
|
| 32 |
+
const MAX_ENCODED_HTTP_BODY_DELTA_BYTES: usize = MAX_HTTP_BODY_DELTA_BYTES.div_ceil(3) * 4;
|
| 33 |
+
|
| 34 |
+
pub(crate) struct QueuedHttpBodyDelta {
|
| 35 |
+
notification: HttpRequestBodyDeltaNotification,
|
| 36 |
+
_byte_permit: Option<OwnedSemaphorePermit>,
|
| 37 |
+
}
|
| 38 |
+
|
| 39 |
+
impl QueuedHttpBodyDelta {
|
| 40 |
+
pub(crate) fn new(
|
| 41 |
+
notification: HttpRequestBodyDeltaNotification,
|
| 42 |
+
byte_permit: Option<OwnedSemaphorePermit>,
|
| 43 |
+
) -> Self {
|
| 44 |
+
Self {
|
| 45 |
+
notification,
|
| 46 |
+
_byte_permit: byte_permit,
|
| 47 |
+
}
|
| 48 |
+
}
|
| 49 |
+
}
|
| 50 |
+
|
| 51 |
+
pub(super) struct HttpBodyStreamRegistration {
|
| 52 |
+
inner: Arc<Inner>,
|
| 53 |
+
request_id: String,
|
| 54 |
+
active: bool,
|
| 55 |
+
}
|
| 56 |
+
|
| 57 |
+
enum HttpResponseBodyStreamInner {
|
| 58 |
+
Local {
|
| 59 |
+
body: Pin<Box<dyn futures::Stream<Item = Result<Bytes, HttpError>> + Send>>,
|
| 60 |
+
},
|
| 61 |
+
Remote {
|
| 62 |
+
inner: Arc<Inner>,
|
| 63 |
+
request_id: String,
|
| 64 |
+
next_seq: u64,
|
| 65 |
+
rx: mpsc::Receiver<QueuedHttpBodyDelta>,
|
| 66 |
+
pending_eof: bool,
|
| 67 |
+
closed: bool,
|
| 68 |
+
},
|
| 69 |
+
}
|
| 70 |
+
|
| 71 |
+
/// Request-scoped stream of body chunks for an HTTP response.
|
| 72 |
+
///
|
| 73 |
+
/// The initial `http/request` call returns status and headers. This stream then
|
| 74 |
+
/// receives the ordered `http/request/bodyDelta` notifications for that request
|
| 75 |
+
/// id until EOF or a terminal error.
|
| 76 |
+
pub struct HttpResponseBodyStream {
|
| 77 |
+
inner: HttpResponseBodyStreamInner,
|
| 78 |
+
}
|
| 79 |
+
|
| 80 |
+
impl HttpResponseBodyStream {
|
| 81 |
+
/// Creates an in-memory response stream from pre-buffered chunks.
|
| 82 |
+
///
|
| 83 |
+
/// This is useful for [`crate::HttpClient`] implementations that already
|
| 84 |
+
/// own the response bytes, including lightweight test clients.
|
| 85 |
+
#[doc(hidden)]
|
| 86 |
+
pub fn from_chunks(chunks: Vec<Vec<u8>>) -> Self {
|
| 87 |
+
let body = futures::stream::iter(
|
| 88 |
+
chunks
|
| 89 |
+
.into_iter()
|
| 90 |
+
.map(|chunk| Ok::<Bytes, HttpError>(chunk.into())),
|
| 91 |
+
);
|
| 92 |
+
Self {
|
| 93 |
+
inner: HttpResponseBodyStreamInner::Local {
|
| 94 |
+
body: Box::pin(body),
|
| 95 |
+
},
|
| 96 |
+
}
|
| 97 |
+
}
|
| 98 |
+
|
| 99 |
+
pub(super) fn local(response: HttpResponse) -> Self {
|
| 100 |
+
Self {
|
| 101 |
+
inner: HttpResponseBodyStreamInner::Local {
|
| 102 |
+
body: Box::pin(response.bytes_stream()),
|
| 103 |
+
},
|
| 104 |
+
}
|
| 105 |
+
}
|
| 106 |
+
|
| 107 |
+
pub(super) fn remote(
|
| 108 |
+
inner: Arc<Inner>,
|
| 109 |
+
request_id: String,
|
| 110 |
+
rx: mpsc::Receiver<QueuedHttpBodyDelta>,
|
| 111 |
+
) -> Self {
|
| 112 |
+
Self {
|
| 113 |
+
inner: HttpResponseBodyStreamInner::Remote {
|
| 114 |
+
inner,
|
| 115 |
+
request_id,
|
| 116 |
+
next_seq: 1,
|
| 117 |
+
rx,
|
| 118 |
+
pending_eof: false,
|
| 119 |
+
closed: false,
|
| 120 |
+
},
|
| 121 |
+
}
|
| 122 |
+
}
|
| 123 |
+
|
| 124 |
+
/// Receives the next response-body chunk.
|
| 125 |
+
///
|
| 126 |
+
/// Returns `Ok(None)` at EOF and converts sequence gaps or stream-side
|
| 127 |
+
/// stream errors into protocol errors.
|
| 128 |
+
pub async fn recv(&mut self) -> Result<Option<Vec<u8>>, ExecServerError> {
|
| 129 |
+
match &mut self.inner {
|
| 130 |
+
HttpResponseBodyStreamInner::Local { body } => match body.next().await {
|
| 131 |
+
Some(chunk) => match chunk {
|
| 132 |
+
Ok(bytes) => Ok(Some(bytes.to_vec())),
|
| 133 |
+
Err(error) => Err(ExecServerError::HttpRequest(error.to_string())),
|
| 134 |
+
},
|
| 135 |
+
None => Ok(None),
|
| 136 |
+
},
|
| 137 |
+
HttpResponseBodyStreamInner::Remote {
|
| 138 |
+
inner,
|
| 139 |
+
request_id,
|
| 140 |
+
next_seq,
|
| 141 |
+
rx,
|
| 142 |
+
pending_eof,
|
| 143 |
+
closed,
|
| 144 |
+
} => {
|
| 145 |
+
if *pending_eof {
|
| 146 |
+
*pending_eof = false;
|
| 147 |
+
finish_remote_stream(inner, request_id, closed).await;
|
| 148 |
+
return Ok(None);
|
| 149 |
+
}
|
| 150 |
+
|
| 151 |
+
let Some(QueuedHttpBodyDelta {
|
| 152 |
+
notification: delta,
|
| 153 |
+
..
|
| 154 |
+
}) = rx.recv().await
|
| 155 |
+
else {
|
| 156 |
+
finish_remote_stream(inner, request_id, closed).await;
|
| 157 |
+
if let Some(error) = inner.take_http_body_stream_failure(request_id).await {
|
| 158 |
+
return Err(ExecServerError::Protocol(format!(
|
| 159 |
+
"http response stream `{request_id}` failed: {error}",
|
| 160 |
+
)));
|
| 161 |
+
}
|
| 162 |
+
return Ok(None);
|
| 163 |
+
};
|
| 164 |
+
if delta.seq != *next_seq {
|
| 165 |
+
finish_remote_stream(inner, request_id, closed).await;
|
| 166 |
+
return Err(ExecServerError::Protocol(format!(
|
| 167 |
+
"http response stream `{request_id}` received seq {}, expected {}",
|
| 168 |
+
delta.seq, *next_seq
|
| 169 |
+
)));
|
| 170 |
+
}
|
| 171 |
+
*next_seq += 1;
|
| 172 |
+
let chunk = delta.delta.into_inner();
|
| 173 |
+
|
| 174 |
+
if let Some(error) = delta.error {
|
| 175 |
+
finish_remote_stream(inner, request_id, closed).await;
|
| 176 |
+
return Err(ExecServerError::Protocol(format!(
|
| 177 |
+
"http response stream `{request_id}` failed: {error}",
|
| 178 |
+
)));
|
| 179 |
+
}
|
| 180 |
+
if delta.done {
|
| 181 |
+
finish_remote_stream(inner, request_id, closed).await;
|
| 182 |
+
if chunk.is_empty() {
|
| 183 |
+
return Ok(None);
|
| 184 |
+
}
|
| 185 |
+
*pending_eof = true;
|
| 186 |
+
}
|
| 187 |
+
Ok(Some(chunk))
|
| 188 |
+
}
|
| 189 |
+
}
|
| 190 |
+
}
|
| 191 |
+
}
|
| 192 |
+
|
| 193 |
+
impl Drop for HttpResponseBodyStream {
|
| 194 |
+
/// Schedules stream-route removal if the consumer drops before EOF.
|
| 195 |
+
fn drop(&mut self) {
|
| 196 |
+
if let HttpResponseBodyStreamInner::Remote {
|
| 197 |
+
inner,
|
| 198 |
+
request_id,
|
| 199 |
+
closed,
|
| 200 |
+
..
|
| 201 |
+
} = &mut self.inner
|
| 202 |
+
{
|
| 203 |
+
if *closed {
|
| 204 |
+
return;
|
| 205 |
+
}
|
| 206 |
+
*closed = true;
|
| 207 |
+
spawn_remove_http_body_stream(Arc::clone(inner), request_id.clone());
|
| 208 |
+
}
|
| 209 |
+
}
|
| 210 |
+
}
|
| 211 |
+
|
| 212 |
+
impl HttpBodyStreamRegistration {
|
| 213 |
+
pub(super) fn new(inner: Arc<Inner>, request_id: String) -> Self {
|
| 214 |
+
Self {
|
| 215 |
+
inner,
|
| 216 |
+
request_id,
|
| 217 |
+
active: true,
|
| 218 |
+
}
|
| 219 |
+
}
|
| 220 |
+
|
| 221 |
+
pub(super) fn disarm(&mut self) {
|
| 222 |
+
self.active = false;
|
| 223 |
+
}
|
| 224 |
+
}
|
| 225 |
+
|
| 226 |
+
impl Drop for HttpBodyStreamRegistration {
|
| 227 |
+
/// Removes the route if the stream request future is cancelled before headers return.
|
| 228 |
+
fn drop(&mut self) {
|
| 229 |
+
if self.active {
|
| 230 |
+
spawn_remove_http_body_stream(Arc::clone(&self.inner), self.request_id.clone());
|
| 231 |
+
}
|
| 232 |
+
}
|
| 233 |
+
}
|
| 234 |
+
|
| 235 |
+
async fn finish_remote_stream(inner: &Arc<Inner>, request_id: &str, closed: &mut bool) {
|
| 236 |
+
if *closed {
|
| 237 |
+
return;
|
| 238 |
+
}
|
| 239 |
+
*closed = true;
|
| 240 |
+
inner.remove_http_body_stream(request_id).await;
|
| 241 |
+
}
|
| 242 |
+
|
| 243 |
+
/// Schedules HTTP body route removal from synchronous drop paths.
|
| 244 |
+
fn spawn_remove_http_body_stream(inner: Arc<Inner>, request_id: String) {
|
| 245 |
+
if let Ok(handle) = Handle::try_current() {
|
| 246 |
+
handle.spawn(async move {
|
| 247 |
+
inner.remove_http_body_stream(&request_id).await;
|
| 248 |
+
});
|
| 249 |
+
}
|
| 250 |
+
}
|
| 251 |
+
|
| 252 |
+
pub(super) async fn send_body_delta(
|
| 253 |
+
notifications: &RpcNotificationSender,
|
| 254 |
+
delta: HttpRequestBodyDeltaNotification,
|
| 255 |
+
) -> bool {
|
| 256 |
+
notifications
|
| 257 |
+
.notify(HTTP_REQUEST_BODY_DELTA_METHOD, &delta)
|
| 258 |
+
.await
|
| 259 |
+
.is_ok()
|
| 260 |
+
}
|
| 261 |
+
|
| 262 |
+
impl Inner {
|
| 263 |
+
/// Routes one streamed HTTP body notification into its request-local receiver.
|
| 264 |
+
pub(crate) async fn handle_http_body_delta_notification(
|
| 265 |
+
&self,
|
| 266 |
+
params: Option<Value>,
|
| 267 |
+
) -> Result<(), ExecServerError> {
|
| 268 |
+
let params = params.unwrap_or(Value::Null);
|
| 269 |
+
if params
|
| 270 |
+
.get("deltaBase64")
|
| 271 |
+
.and_then(Value::as_str)
|
| 272 |
+
.is_some_and(|delta| delta.len() > MAX_ENCODED_HTTP_BODY_DELTA_BYTES)
|
| 273 |
+
{
|
| 274 |
+
return Err(ExecServerError::Protocol(format!(
|
| 275 |
+
"http response body delta exceeds {MAX_HTTP_BODY_DELTA_BYTES} bytes"
|
| 276 |
+
)));
|
| 277 |
+
}
|
| 278 |
+
let params: HttpRequestBodyDeltaNotification = from_value(params)?;
|
| 279 |
+
if params.delta.0.len() > MAX_HTTP_BODY_DELTA_BYTES {
|
| 280 |
+
return Err(ExecServerError::Protocol(format!(
|
| 281 |
+
"http response body delta exceeds {MAX_HTTP_BODY_DELTA_BYTES} bytes"
|
| 282 |
+
)));
|
| 283 |
+
}
|
| 284 |
+
// Unknown request ids are ignored intentionally: a stream may have already
|
| 285 |
+
// reached EOF and released its route.
|
| 286 |
+
if let Some(tx) = self
|
| 287 |
+
.http_body_streams
|
| 288 |
+
.load()
|
| 289 |
+
.get(¶ms.request_id)
|
| 290 |
+
.cloned()
|
| 291 |
+
{
|
| 292 |
+
let request_id = params.request_id.clone();
|
| 293 |
+
let terminal_delta = params.done || params.error.is_some();
|
| 294 |
+
let queued_bytes = params
|
| 295 |
+
.delta
|
| 296 |
+
.0
|
| 297 |
+
.len()
|
| 298 |
+
.saturating_add(params.error.as_deref().map_or(0, str::len));
|
| 299 |
+
let byte_permit = if queued_bytes == 0 {
|
| 300 |
+
None
|
| 301 |
+
} else {
|
| 302 |
+
u32::try_from(queued_bytes).ok().and_then(|queued_bytes| {
|
| 303 |
+
Arc::clone(&self.http_body_stream_byte_budget)
|
| 304 |
+
.try_acquire_many_owned(queued_bytes)
|
| 305 |
+
.ok()
|
| 306 |
+
})
|
| 307 |
+
};
|
| 308 |
+
if queued_bytes > 0 && byte_permit.is_none() {
|
| 309 |
+
self.record_http_body_stream_failure(
|
| 310 |
+
&request_id,
|
| 311 |
+
format!("queued body deltas exceed {MAX_QUEUED_HTTP_BODY_BYTES} bytes"),
|
| 312 |
+
)
|
| 313 |
+
.await;
|
| 314 |
+
self.remove_http_body_stream(&request_id).await;
|
| 315 |
+
debug!(
|
| 316 |
+
"closing http response stream `{request_id}` after exhausting the queued byte budget"
|
| 317 |
+
);
|
| 318 |
+
return Ok(());
|
| 319 |
+
}
|
| 320 |
+
match tx.try_send(QueuedHttpBodyDelta::new(params, byte_permit)) {
|
| 321 |
+
Ok(()) => {
|
| 322 |
+
if terminal_delta {
|
| 323 |
+
self.remove_http_body_stream(&request_id).await;
|
| 324 |
+
}
|
| 325 |
+
}
|
| 326 |
+
Err(TrySendError::Closed(_)) => {
|
| 327 |
+
self.remove_http_body_stream(&request_id).await;
|
| 328 |
+
debug!("http response stream receiver dropped before body delta delivery");
|
| 329 |
+
}
|
| 330 |
+
Err(TrySendError::Full(_)) => {
|
| 331 |
+
self.record_http_body_stream_failure(
|
| 332 |
+
&request_id,
|
| 333 |
+
"body delta channel filled before delivery".to_string(),
|
| 334 |
+
)
|
| 335 |
+
.await;
|
| 336 |
+
self.remove_http_body_stream(&request_id).await;
|
| 337 |
+
debug!(
|
| 338 |
+
"closing http response stream `{request_id}` after body delta backpressure"
|
| 339 |
+
);
|
| 340 |
+
}
|
| 341 |
+
}
|
| 342 |
+
}
|
| 343 |
+
Ok(())
|
| 344 |
+
}
|
| 345 |
+
|
| 346 |
+
/// Fails active streamed HTTP bodies so callers do not wait forever after a
|
| 347 |
+
/// transport disconnect or notification handling failure.
|
| 348 |
+
pub(crate) async fn fail_all_http_body_streams(&self, message: String) {
|
| 349 |
+
let _streams_write_guard = self.http_body_streams_write_lock.lock().await;
|
| 350 |
+
let streams = self.http_body_streams.load();
|
| 351 |
+
let streams = streams.as_ref().clone();
|
| 352 |
+
self.http_body_streams.store(Arc::new(HashMap::new()));
|
| 353 |
+
for (request_id, tx) in streams {
|
| 354 |
+
// Failure notifications must wake every stream even when no
|
| 355 |
+
// byte-budget permits remain.
|
| 356 |
+
if tx
|
| 357 |
+
.try_send(QueuedHttpBodyDelta::new(
|
| 358 |
+
HttpRequestBodyDeltaNotification {
|
| 359 |
+
request_id: request_id.clone(),
|
| 360 |
+
seq: 1,
|
| 361 |
+
delta: Vec::new().into(),
|
| 362 |
+
done: true,
|
| 363 |
+
error: Some(message.clone()),
|
| 364 |
+
},
|
| 365 |
+
/*byte_permit*/ None,
|
| 366 |
+
))
|
| 367 |
+
.is_err()
|
| 368 |
+
{
|
| 369 |
+
let mut next_failures = self.http_body_stream_failures.load().as_ref().clone();
|
| 370 |
+
next_failures.insert(request_id, message.clone());
|
| 371 |
+
self.http_body_stream_failures
|
| 372 |
+
.store(Arc::new(next_failures));
|
| 373 |
+
}
|
| 374 |
+
}
|
| 375 |
+
}
|
| 376 |
+
|
| 377 |
+
/// Allocates a connection-local streamed HTTP response id.
|
| 378 |
+
pub(super) fn next_http_body_stream_request_id(&self) -> String {
|
| 379 |
+
let id = self
|
| 380 |
+
.http_body_stream_next_id
|
| 381 |
+
.fetch_add(1, Ordering::Relaxed);
|
| 382 |
+
format!("http-{id}")
|
| 383 |
+
}
|
| 384 |
+
|
| 385 |
+
/// Registers a request id before issuing a streaming HTTP call.
|
| 386 |
+
pub(super) async fn insert_http_body_stream(
|
| 387 |
+
&self,
|
| 388 |
+
request_id: String,
|
| 389 |
+
tx: mpsc::Sender<QueuedHttpBodyDelta>,
|
| 390 |
+
) -> Result<(), ExecServerError> {
|
| 391 |
+
let _streams_write_guard = self.http_body_streams_write_lock.lock().await;
|
| 392 |
+
let streams = self.http_body_streams.load();
|
| 393 |
+
if streams.contains_key(&request_id) {
|
| 394 |
+
return Err(ExecServerError::Protocol(format!(
|
| 395 |
+
"http response stream already registered for request {request_id}"
|
| 396 |
+
)));
|
| 397 |
+
}
|
| 398 |
+
let mut next_streams = streams.as_ref().clone();
|
| 399 |
+
next_streams.insert(request_id.clone(), tx);
|
| 400 |
+
self.http_body_streams.store(Arc::new(next_streams));
|
| 401 |
+
let failures = self.http_body_stream_failures.load();
|
| 402 |
+
if failures.contains_key(&request_id) {
|
| 403 |
+
let mut next_failures = failures.as_ref().clone();
|
| 404 |
+
next_failures.remove(&request_id);
|
| 405 |
+
self.http_body_stream_failures
|
| 406 |
+
.store(Arc::new(next_failures));
|
| 407 |
+
}
|
| 408 |
+
Ok(())
|
| 409 |
+
}
|
| 410 |
+
|
| 411 |
+
/// Removes a request id after EOF, terminal error, or request failure.
|
| 412 |
+
pub(super) async fn remove_http_body_stream(
|
| 413 |
+
&self,
|
| 414 |
+
request_id: &str,
|
| 415 |
+
) -> Option<mpsc::Sender<QueuedHttpBodyDelta>> {
|
| 416 |
+
let _streams_write_guard = self.http_body_streams_write_lock.lock().await;
|
| 417 |
+
let streams = self.http_body_streams.load();
|
| 418 |
+
let stream = streams.get(request_id).cloned();
|
| 419 |
+
stream.as_ref()?;
|
| 420 |
+
let mut next_streams = streams.as_ref().clone();
|
| 421 |
+
next_streams.remove(request_id);
|
| 422 |
+
self.http_body_streams.store(Arc::new(next_streams));
|
| 423 |
+
stream
|
| 424 |
+
}
|
| 425 |
+
|
| 426 |
+
async fn record_http_body_stream_failure(&self, request_id: &str, message: String) {
|
| 427 |
+
let _streams_write_guard = self.http_body_streams_write_lock.lock().await;
|
| 428 |
+
let failures = self.http_body_stream_failures.load();
|
| 429 |
+
let mut next_failures = failures.as_ref().clone();
|
| 430 |
+
next_failures.insert(request_id.to_string(), message);
|
| 431 |
+
self.http_body_stream_failures
|
| 432 |
+
.store(Arc::new(next_failures));
|
| 433 |
+
}
|
| 434 |
+
|
| 435 |
+
async fn take_http_body_stream_failure(&self, request_id: &str) -> Option<String> {
|
| 436 |
+
let _streams_write_guard = self.http_body_streams_write_lock.lock().await;
|
| 437 |
+
let failures = self.http_body_stream_failures.load();
|
| 438 |
+
let error = failures.get(request_id).cloned();
|
| 439 |
+
error.as_ref()?;
|
| 440 |
+
let mut next_failures = failures.as_ref().clone();
|
| 441 |
+
next_failures.remove(request_id);
|
| 442 |
+
self.http_body_stream_failures
|
| 443 |
+
.store(Arc::new(next_failures));
|
| 444 |
+
error
|
| 445 |
+
}
|
| 446 |
+
}
|
codex-rs/exec-server/src/client/network_policy_audit.rs
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::NetworkPolicyAuditContext;
|
| 2 |
+
use crate::protocol::ExecServerNetworkProtocol;
|
| 3 |
+
use crate::protocol::MAX_NETWORK_POLICY_HOST_BYTES;
|
| 4 |
+
use crate::protocol::MAX_NETWORK_POLICY_PROCESS_ID_BYTES;
|
| 5 |
+
use crate::protocol::MAX_NETWORK_POLICY_REASON_BYTES;
|
| 6 |
+
use crate::protocol::NetworkPolicyDecisionNotification;
|
| 7 |
+
|
| 8 |
+
const MAX_NETWORK_POLICY_METHOD_BYTES: usize = 32;
|
| 9 |
+
const MAX_NETWORK_POLICY_CLIENT_BYTES: usize = 256;
|
| 10 |
+
const MAX_NETWORK_POLICY_TIMESTAMP_BYTES: usize = 64;
|
| 11 |
+
|
| 12 |
+
pub(super) fn emit_network_policy_decision(
|
| 13 |
+
context: &NetworkPolicyAuditContext,
|
| 14 |
+
decision: &NetworkPolicyDecisionNotification,
|
| 15 |
+
) -> bool {
|
| 16 |
+
if decision.process_id.is_empty()
|
| 17 |
+
|| decision.process_id.len() > MAX_NETWORK_POLICY_PROCESS_ID_BYTES
|
| 18 |
+
|| decision.host.is_empty()
|
| 19 |
+
|| decision.host.len() > MAX_NETWORK_POLICY_HOST_BYTES
|
| 20 |
+
|| decision.host.chars().any(char::is_control)
|
| 21 |
+
|| decision.host.chars().any(char::is_whitespace)
|
| 22 |
+
|| decision.reason.len() > MAX_NETWORK_POLICY_REASON_BYTES
|
| 23 |
+
|| decision.reason.chars().any(char::is_control)
|
| 24 |
+
|| !matches!(decision.scope.as_str(), "domain" | "non_domain")
|
| 25 |
+
|| !matches!(decision.decision.as_str(), "allow" | "deny" | "ask")
|
| 26 |
+
|| !matches!(
|
| 27 |
+
decision.source.as_str(),
|
| 28 |
+
"baseline_policy" | "mode_guard" | "proxy_state" | "decider"
|
| 29 |
+
)
|
| 30 |
+
|| decision.timestamp.is_empty()
|
| 31 |
+
|| decision.timestamp.len() > MAX_NETWORK_POLICY_TIMESTAMP_BYTES
|
| 32 |
+
|| decision.timestamp.chars().any(char::is_control)
|
| 33 |
+
|| decision.method.as_ref().is_some_and(|method| {
|
| 34 |
+
method.len() > MAX_NETWORK_POLICY_METHOD_BYTES
|
| 35 |
+
|| method.chars().any(char::is_control)
|
| 36 |
+
|| method.chars().any(char::is_whitespace)
|
| 37 |
+
})
|
| 38 |
+
|| decision.client.as_ref().is_some_and(|client| {
|
| 39 |
+
client.len() > MAX_NETWORK_POLICY_CLIENT_BYTES
|
| 40 |
+
|| client.chars().any(char::is_control)
|
| 41 |
+
|| client.chars().any(char::is_whitespace)
|
| 42 |
+
})
|
| 43 |
+
{
|
| 44 |
+
return false;
|
| 45 |
+
}
|
| 46 |
+
|
| 47 |
+
let protocol = match decision.protocol {
|
| 48 |
+
ExecServerNetworkProtocol::Http => "http",
|
| 49 |
+
ExecServerNetworkProtocol::HttpsConnect => "https_connect",
|
| 50 |
+
ExecServerNetworkProtocol::Socks5Tcp => "socks5_tcp",
|
| 51 |
+
ExecServerNetworkProtocol::Socks5Udp => "socks5_udp",
|
| 52 |
+
};
|
| 53 |
+
let metadata = &context.metadata;
|
| 54 |
+
tracing::event!(
|
| 55 |
+
target: "codex_otel.log_only",
|
| 56 |
+
tracing::Level::INFO,
|
| 57 |
+
event.name = "codex.network_proxy.policy_decision",
|
| 58 |
+
event.timestamp = decision.timestamp,
|
| 59 |
+
conversation.id = metadata.conversation_id.as_deref(),
|
| 60 |
+
app.version = metadata.app_version.as_deref(),
|
| 61 |
+
auth_mode = metadata.auth_mode.as_deref(),
|
| 62 |
+
originator = metadata.originator.as_deref(),
|
| 63 |
+
user.account_id = metadata.user_account_id.as_deref(),
|
| 64 |
+
user.email = metadata.user_email.as_deref(),
|
| 65 |
+
terminal.type = metadata.terminal_type.as_deref(),
|
| 66 |
+
model = metadata.model.as_deref(),
|
| 67 |
+
slug = metadata.slug.as_deref(),
|
| 68 |
+
network.policy.scope = decision.scope,
|
| 69 |
+
network.policy.decision = decision.decision,
|
| 70 |
+
network.policy.source = decision.source,
|
| 71 |
+
network.policy.reason = decision.reason,
|
| 72 |
+
network.transport.protocol = protocol,
|
| 73 |
+
server.address = decision.host,
|
| 74 |
+
server.port = decision.port,
|
| 75 |
+
http.request.method = decision.method.as_deref().unwrap_or("none"),
|
| 76 |
+
client.address = decision.client.as_deref().unwrap_or("unknown"),
|
| 77 |
+
execution.id = context.execution_id.as_deref(),
|
| 78 |
+
network.policy.override = decision.policy_override,
|
| 79 |
+
);
|
| 80 |
+
true
|
| 81 |
+
}
|
codex-rs/exec-server/src/client/route_aware_http_client.rs
ADDED
|
@@ -0,0 +1,378 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Route-aware local HTTP capability implementation.
|
| 2 |
+
//!
|
| 3 |
+
//! This code runs wherever the real network request should originate:
|
| 4 |
+
//! - in a local environment, that means the orchestrator process
|
| 5 |
+
//! - in a remote environment, that means the remote runtime after the
|
| 6 |
+
//! orchestrator has forwarded `http/request` over JSON-RPC
|
| 7 |
+
|
| 8 |
+
use std::time::Duration;
|
| 9 |
+
|
| 10 |
+
use codex_exec_server_protocol::JSONRPCErrorError;
|
| 11 |
+
use codex_http_client::ClientRouteClass;
|
| 12 |
+
use codex_http_client::HttpClientFactory;
|
| 13 |
+
use codex_http_client::RouteAwareClientPool;
|
| 14 |
+
use codex_http_client::RouteAwareRequestError;
|
| 15 |
+
use codex_protocol::shell_environment::CODEX_EXEC_SERVER_NOISE_AUTH_TOKEN_ENV_VAR;
|
| 16 |
+
use codex_protocol::shell_environment::OPENAI_FEDERATION_RULE_ID_ENV_VAR;
|
| 17 |
+
use codex_protocol::shell_environment::OPENAI_IDENTITY_TOKEN_FILE_ENV_VAR;
|
| 18 |
+
use codex_protocol::shell_environment::OPENAI_WORKLOAD_IDENTITY_CONTEXT_ENV_VAR;
|
| 19 |
+
use futures::FutureExt;
|
| 20 |
+
use futures::StreamExt;
|
| 21 |
+
use futures::future::BoxFuture;
|
| 22 |
+
use http::HeaderMap;
|
| 23 |
+
use http::HeaderName;
|
| 24 |
+
use http::HeaderValue;
|
| 25 |
+
use http::Method;
|
| 26 |
+
use tracing::Instrument;
|
| 27 |
+
use url::Url;
|
| 28 |
+
|
| 29 |
+
use super::HttpResponseBodyStream;
|
| 30 |
+
use super::response_body_stream::send_body_delta;
|
| 31 |
+
use crate::HttpClient;
|
| 32 |
+
use crate::client::ExecServerError;
|
| 33 |
+
use crate::protocol::HttpHeader;
|
| 34 |
+
use crate::protocol::HttpRedirectPolicy;
|
| 35 |
+
use crate::protocol::HttpRequestBodyDeltaNotification;
|
| 36 |
+
use crate::protocol::HttpRequestParams;
|
| 37 |
+
use crate::protocol::HttpRequestResponse;
|
| 38 |
+
use crate::protocol::MAX_HTTP_BODY_DELTA_BYTES;
|
| 39 |
+
use crate::rpc::RpcNotificationSender;
|
| 40 |
+
use crate::rpc::internal_error;
|
| 41 |
+
use crate::rpc::invalid_params;
|
| 42 |
+
|
| 43 |
+
const HTTP_HEADER_ENV_DENYLIST: &[&str] = &[
|
| 44 |
+
CODEX_EXEC_SERVER_NOISE_AUTH_TOKEN_ENV_VAR,
|
| 45 |
+
OPENAI_FEDERATION_RULE_ID_ENV_VAR,
|
| 46 |
+
OPENAI_IDENTITY_TOKEN_FILE_ENV_VAR,
|
| 47 |
+
OPENAI_WORKLOAD_IDENTITY_CONTEXT_ENV_VAR,
|
| 48 |
+
"OPENAI_API_KEY",
|
| 49 |
+
"CODEX_API_KEY",
|
| 50 |
+
"CODEX_ACCESS_TOKEN",
|
| 51 |
+
"CODEX_CONNECTORS_TOKEN",
|
| 52 |
+
"AWS_ACCESS_KEY_ID",
|
| 53 |
+
"AWS_SECRET_ACCESS_KEY",
|
| 54 |
+
"AWS_SESSION_TOKEN",
|
| 55 |
+
"AZURE_CLIENT_SECRET",
|
| 56 |
+
"AZURE_FEDERATED_TOKEN_FILE",
|
| 57 |
+
"GOOGLE_APPLICATION_CREDENTIALS",
|
| 58 |
+
];
|
| 59 |
+
|
| 60 |
+
/// HTTP capability implementation backed by the shared route-aware transport.
|
| 61 |
+
#[derive(Clone)]
|
| 62 |
+
pub struct RouteAwareHttpClient {
|
| 63 |
+
follow_redirects: RouteAwareClientPool,
|
| 64 |
+
stop_redirects: RouteAwareClientPool,
|
| 65 |
+
}
|
| 66 |
+
|
| 67 |
+
/// Streaming response state held between the initial HTTP response and
|
| 68 |
+
/// downstream body-delta forwarding.
|
| 69 |
+
pub(crate) struct PendingRouteAwareHttpBodyStream {
|
| 70 |
+
pub(crate) request_id: String,
|
| 71 |
+
pub(crate) response: codex_http_client::HttpResponse,
|
| 72 |
+
}
|
| 73 |
+
|
| 74 |
+
/// Validates `http/request` parameters and runs the actual HTTP call used
|
| 75 |
+
/// by the exec-server route and the local [`HttpClient`] backend.
|
| 76 |
+
pub(crate) struct RouteAwareHttpRequestRunner {
|
| 77 |
+
client: RouteAwareClientPool,
|
| 78 |
+
}
|
| 79 |
+
|
| 80 |
+
impl RouteAwareHttpClient {
|
| 81 |
+
pub fn new(http_client_factory: HttpClientFactory) -> Self {
|
| 82 |
+
Self {
|
| 83 |
+
follow_redirects: RouteAwareClientPool::with_chatgpt_cloudflare_cookies_without_request_logging(
|
| 84 |
+
http_client_factory.clone(),
|
| 85 |
+
// Delegated HTTP targets arbitrary endpoints; route class only labels diagnostics.
|
| 86 |
+
ClientRouteClass::Other,
|
| 87 |
+
),
|
| 88 |
+
stop_redirects:
|
| 89 |
+
RouteAwareClientPool::with_chatgpt_cloudflare_cookies_without_redirects_or_request_logging(
|
| 90 |
+
http_client_factory,
|
| 91 |
+
// Proxy routing comes from the factory, not this diagnostic-only route class.
|
| 92 |
+
ClientRouteClass::Other,
|
| 93 |
+
),
|
| 94 |
+
}
|
| 95 |
+
}
|
| 96 |
+
|
| 97 |
+
/// Enables narrowly scoped TLS-backend fallback for both redirect policies.
|
| 98 |
+
pub fn with_tls_backend_fallback(mut self) -> Self {
|
| 99 |
+
self.follow_redirects = self.follow_redirects.with_tls_backend_fallback();
|
| 100 |
+
self.stop_redirects = self.stop_redirects.with_tls_backend_fallback();
|
| 101 |
+
self
|
| 102 |
+
}
|
| 103 |
+
|
| 104 |
+
pub(crate) fn runner(
|
| 105 |
+
&self,
|
| 106 |
+
redirect_policy: HttpRedirectPolicy,
|
| 107 |
+
) -> RouteAwareHttpRequestRunner {
|
| 108 |
+
let client = match redirect_policy {
|
| 109 |
+
HttpRedirectPolicy::Follow => self.follow_redirects.clone(),
|
| 110 |
+
HttpRedirectPolicy::Stop => self.stop_redirects.clone(),
|
| 111 |
+
};
|
| 112 |
+
RouteAwareHttpRequestRunner { client }
|
| 113 |
+
}
|
| 114 |
+
}
|
| 115 |
+
|
| 116 |
+
impl HttpClient for RouteAwareHttpClient {
|
| 117 |
+
fn http_request(
|
| 118 |
+
&self,
|
| 119 |
+
params: HttpRequestParams,
|
| 120 |
+
) -> BoxFuture<'_, Result<HttpRequestResponse, ExecServerError>> {
|
| 121 |
+
async move {
|
| 122 |
+
let runner = self.runner(params.redirect_policy);
|
| 123 |
+
let (response, _) = runner
|
| 124 |
+
.run(HttpRequestParams {
|
| 125 |
+
stream_response: false,
|
| 126 |
+
..params
|
| 127 |
+
})
|
| 128 |
+
.await
|
| 129 |
+
.map_err(|error| ExecServerError::HttpRequest(error.message))?;
|
| 130 |
+
Ok(response)
|
| 131 |
+
}
|
| 132 |
+
.boxed()
|
| 133 |
+
}
|
| 134 |
+
|
| 135 |
+
fn http_request_stream(
|
| 136 |
+
&self,
|
| 137 |
+
params: HttpRequestParams,
|
| 138 |
+
) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> {
|
| 139 |
+
async move {
|
| 140 |
+
let runner = self.runner(params.redirect_policy);
|
| 141 |
+
let (response, pending_stream) = runner
|
| 142 |
+
.run(HttpRequestParams {
|
| 143 |
+
stream_response: true,
|
| 144 |
+
..params
|
| 145 |
+
})
|
| 146 |
+
.await
|
| 147 |
+
.map_err(|error| ExecServerError::HttpRequest(error.message))?;
|
| 148 |
+
let pending_stream = pending_stream.ok_or_else(|| {
|
| 149 |
+
ExecServerError::Protocol(
|
| 150 |
+
"http request stream did not return a response body stream".to_string(),
|
| 151 |
+
)
|
| 152 |
+
})?;
|
| 153 |
+
Ok((
|
| 154 |
+
response,
|
| 155 |
+
HttpResponseBodyStream::local(pending_stream.response),
|
| 156 |
+
))
|
| 157 |
+
}
|
| 158 |
+
.boxed()
|
| 159 |
+
}
|
| 160 |
+
}
|
| 161 |
+
|
| 162 |
+
impl RouteAwareHttpRequestRunner {
|
| 163 |
+
pub(crate) async fn run(
|
| 164 |
+
&self,
|
| 165 |
+
params: HttpRequestParams,
|
| 166 |
+
) -> Result<(HttpRequestResponse, Option<PendingRouteAwareHttpBodyStream>), JSONRPCErrorError>
|
| 167 |
+
{
|
| 168 |
+
let method = Method::from_bytes(params.method.as_bytes())
|
| 169 |
+
.map_err(|error| invalid_params(format!("http/request method is invalid: {error}")))?;
|
| 170 |
+
let url = Url::parse(¶ms.url)
|
| 171 |
+
.map_err(|error| invalid_params(format!("http/request url is invalid: {error}")))?;
|
| 172 |
+
match url.scheme() {
|
| 173 |
+
"http" | "https" => {}
|
| 174 |
+
scheme => {
|
| 175 |
+
return Err(invalid_params(format!(
|
| 176 |
+
"http/request only supports http and https URLs, got {scheme}"
|
| 177 |
+
)));
|
| 178 |
+
}
|
| 179 |
+
}
|
| 180 |
+
|
| 181 |
+
let request_span = tracing::info_span!(
|
| 182 |
+
"codex.exec_server.http_request",
|
| 183 |
+
otel.kind = "client",
|
| 184 |
+
http.request.method = method.as_str(),
|
| 185 |
+
server.address = url.host_str().unwrap_or_default(),
|
| 186 |
+
server.port = u64::from(url.port_or_known_default().unwrap_or_default()),
|
| 187 |
+
http.response.status_code = tracing::field::Empty,
|
| 188 |
+
error.type = tracing::field::Empty,
|
| 189 |
+
);
|
| 190 |
+
let mut headers = Self::build_headers(params.headers)?;
|
| 191 |
+
codex_otel::inject_span_w3c_trace_headers(&request_span, &mut headers);
|
| 192 |
+
let mut request = self.client.request(method.clone(), url).headers(headers);
|
| 193 |
+
if let Some(body) = params.body {
|
| 194 |
+
request = request.body(body.into_inner());
|
| 195 |
+
}
|
| 196 |
+
if let Some(timeout_ms) = params.timeout_ms {
|
| 197 |
+
request = request.timeout(Duration::from_millis(timeout_ms));
|
| 198 |
+
}
|
| 199 |
+
|
| 200 |
+
let response = match request.send().instrument(request_span.clone()).await {
|
| 201 |
+
Ok(response) => response,
|
| 202 |
+
Err(error) => {
|
| 203 |
+
request_span.record("error.type", "request");
|
| 204 |
+
let error_message = error.to_string();
|
| 205 |
+
log_send_error(&method, error);
|
| 206 |
+
return Err(internal_error(format!(
|
| 207 |
+
"http/request failed: {error_message}"
|
| 208 |
+
)));
|
| 209 |
+
}
|
| 210 |
+
};
|
| 211 |
+
let status = response.status().as_u16();
|
| 212 |
+
request_span.record("http.response.status_code", u64::from(status));
|
| 213 |
+
let headers = Self::response_headers(response.headers());
|
| 214 |
+
|
| 215 |
+
if params.stream_response {
|
| 216 |
+
return Ok((
|
| 217 |
+
HttpRequestResponse {
|
| 218 |
+
status,
|
| 219 |
+
headers,
|
| 220 |
+
body: Vec::new().into(),
|
| 221 |
+
},
|
| 222 |
+
Some(PendingRouteAwareHttpBodyStream {
|
| 223 |
+
request_id: params.request_id,
|
| 224 |
+
response,
|
| 225 |
+
}),
|
| 226 |
+
));
|
| 227 |
+
}
|
| 228 |
+
|
| 229 |
+
let body = response.bytes().await.map_err(|error| {
|
| 230 |
+
internal_error(format!(
|
| 231 |
+
"failed to read http/request response body: {error}"
|
| 232 |
+
))
|
| 233 |
+
})?;
|
| 234 |
+
|
| 235 |
+
Ok((
|
| 236 |
+
HttpRequestResponse {
|
| 237 |
+
status,
|
| 238 |
+
headers,
|
| 239 |
+
body: body.to_vec().into(),
|
| 240 |
+
},
|
| 241 |
+
None,
|
| 242 |
+
))
|
| 243 |
+
}
|
| 244 |
+
|
| 245 |
+
pub(crate) async fn stream_body(
|
| 246 |
+
pending_stream: PendingRouteAwareHttpBodyStream,
|
| 247 |
+
notifications: RpcNotificationSender,
|
| 248 |
+
) {
|
| 249 |
+
let PendingRouteAwareHttpBodyStream {
|
| 250 |
+
request_id,
|
| 251 |
+
response,
|
| 252 |
+
} = pending_stream;
|
| 253 |
+
let mut seq = 1;
|
| 254 |
+
let mut body = response.bytes_stream();
|
| 255 |
+
while let Some(chunk) = body.next().await {
|
| 256 |
+
match chunk {
|
| 257 |
+
Ok(bytes) => {
|
| 258 |
+
for chunk in bytes.chunks(MAX_HTTP_BODY_DELTA_BYTES) {
|
| 259 |
+
if !send_body_delta(
|
| 260 |
+
¬ifications,
|
| 261 |
+
HttpRequestBodyDeltaNotification {
|
| 262 |
+
request_id: request_id.clone(),
|
| 263 |
+
seq,
|
| 264 |
+
delta: chunk.to_vec().into(),
|
| 265 |
+
done: false,
|
| 266 |
+
error: None,
|
| 267 |
+
},
|
| 268 |
+
)
|
| 269 |
+
.await
|
| 270 |
+
{
|
| 271 |
+
return;
|
| 272 |
+
}
|
| 273 |
+
seq += 1;
|
| 274 |
+
}
|
| 275 |
+
}
|
| 276 |
+
Err(error) => {
|
| 277 |
+
let _ = send_body_delta(
|
| 278 |
+
¬ifications,
|
| 279 |
+
HttpRequestBodyDeltaNotification {
|
| 280 |
+
request_id,
|
| 281 |
+
seq,
|
| 282 |
+
delta: Vec::new().into(),
|
| 283 |
+
done: true,
|
| 284 |
+
error: Some(error.to_string()),
|
| 285 |
+
},
|
| 286 |
+
)
|
| 287 |
+
.await;
|
| 288 |
+
return;
|
| 289 |
+
}
|
| 290 |
+
}
|
| 291 |
+
}
|
| 292 |
+
|
| 293 |
+
let _ = send_body_delta(
|
| 294 |
+
¬ifications,
|
| 295 |
+
HttpRequestBodyDeltaNotification {
|
| 296 |
+
request_id,
|
| 297 |
+
seq,
|
| 298 |
+
delta: Vec::new().into(),
|
| 299 |
+
done: true,
|
| 300 |
+
error: None,
|
| 301 |
+
},
|
| 302 |
+
)
|
| 303 |
+
.await;
|
| 304 |
+
}
|
| 305 |
+
|
| 306 |
+
fn build_headers(headers: Vec<HttpHeader>) -> Result<HeaderMap, JSONRPCErrorError> {
|
| 307 |
+
let mut header_map = HeaderMap::new();
|
| 308 |
+
for header in headers {
|
| 309 |
+
let name = HeaderName::from_bytes(header.name.as_bytes()).map_err(|error| {
|
| 310 |
+
invalid_params(format!("http/request header name is invalid: {error}"))
|
| 311 |
+
})?;
|
| 312 |
+
let value = match header.value_env_var {
|
| 313 |
+
Some(env_var) => {
|
| 314 |
+
if HTTP_HEADER_ENV_DENYLIST
|
| 315 |
+
.iter()
|
| 316 |
+
.any(|denied| denied.eq_ignore_ascii_case(&env_var))
|
| 317 |
+
{
|
| 318 |
+
return Err(invalid_params(format!(
|
| 319 |
+
"http/request header {} cannot use executor environment variable {env_var}",
|
| 320 |
+
header.name
|
| 321 |
+
)));
|
| 322 |
+
}
|
| 323 |
+
let env_value = std::env::var(&env_var).map_err(|_| {
|
| 324 |
+
invalid_params(format!(
|
| 325 |
+
"http/request header {} requires executor environment variable {env_var}",
|
| 326 |
+
header.name
|
| 327 |
+
))
|
| 328 |
+
})?;
|
| 329 |
+
if env_value.is_empty() {
|
| 330 |
+
return Err(invalid_params(format!(
|
| 331 |
+
"http/request header {} requires a non-empty executor environment variable {env_var}",
|
| 332 |
+
header.name
|
| 333 |
+
)));
|
| 334 |
+
}
|
| 335 |
+
format!("{}{env_value}", header.value)
|
| 336 |
+
}
|
| 337 |
+
None => header.value,
|
| 338 |
+
};
|
| 339 |
+
let value = HeaderValue::from_str(&value).map_err(|error| {
|
| 340 |
+
invalid_params(format!(
|
| 341 |
+
"http/request header value is invalid for {}: {error}",
|
| 342 |
+
header.name
|
| 343 |
+
))
|
| 344 |
+
})?;
|
| 345 |
+
header_map.append(name, value);
|
| 346 |
+
}
|
| 347 |
+
Ok(header_map)
|
| 348 |
+
}
|
| 349 |
+
|
| 350 |
+
fn response_headers(headers: &HeaderMap) -> Vec<HttpHeader> {
|
| 351 |
+
headers
|
| 352 |
+
.iter()
|
| 353 |
+
.filter_map(|(name, value)| {
|
| 354 |
+
Some(HttpHeader {
|
| 355 |
+
name: name.as_str().to_string(),
|
| 356 |
+
value: value.to_str().ok()?.to_string(),
|
| 357 |
+
value_env_var: None,
|
| 358 |
+
})
|
| 359 |
+
})
|
| 360 |
+
.collect()
|
| 361 |
+
}
|
| 362 |
+
}
|
| 363 |
+
|
| 364 |
+
fn log_send_error(method: &Method, error: RouteAwareRequestError) {
|
| 365 |
+
let error_is_timeout = error.is_timeout();
|
| 366 |
+
let error_is_connect = error.is_connect();
|
| 367 |
+
let error = match error {
|
| 368 |
+
RouteAwareRequestError::Request(error) => error.without_url().to_string(),
|
| 369 |
+
error => error.to_string(),
|
| 370 |
+
};
|
| 371 |
+
tracing::warn!(
|
| 372 |
+
http_method = method.as_str(),
|
| 373 |
+
error_is_timeout,
|
| 374 |
+
error_is_connect,
|
| 375 |
+
error = %error,
|
| 376 |
+
"http/request send failed"
|
| 377 |
+
);
|
| 378 |
+
}
|
codex-rs/exec-server/src/client/rpc_http_client.rs
ADDED
|
@@ -0,0 +1,92 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! JSON-RPC-backed `HttpClient` implementation.
|
| 2 |
+
//!
|
| 3 |
+
//! This code runs in the orchestrator process. It does not issue network
|
| 4 |
+
//! requests directly; instead it forwards `http/request` to the remote runtime
|
| 5 |
+
//! and then reconstructs streamed bodies from `http/request/bodyDelta`
|
| 6 |
+
//! notifications on the shared connection.
|
| 7 |
+
|
| 8 |
+
use std::sync::Arc;
|
| 9 |
+
|
| 10 |
+
use futures::FutureExt;
|
| 11 |
+
use futures::future::BoxFuture;
|
| 12 |
+
use tokio::sync::mpsc;
|
| 13 |
+
|
| 14 |
+
use super::HttpResponseBodyStream;
|
| 15 |
+
use super::response_body_stream::HttpBodyStreamRegistration;
|
| 16 |
+
use crate::HttpClient;
|
| 17 |
+
use crate::client::ExecServerClient;
|
| 18 |
+
use crate::client::ExecServerError;
|
| 19 |
+
use crate::protocol::HTTP_REQUEST_METHOD;
|
| 20 |
+
use crate::protocol::HttpRequestParams;
|
| 21 |
+
use crate::protocol::HttpRequestResponse;
|
| 22 |
+
|
| 23 |
+
/// Maximum queued body frames per streamed HTTP response.
|
| 24 |
+
const HTTP_BODY_DELTA_CHANNEL_CAPACITY: usize = 256;
|
| 25 |
+
|
| 26 |
+
impl ExecServerClient {
|
| 27 |
+
/// Performs an HTTP request and buffers the response body.
|
| 28 |
+
pub async fn http_request(
|
| 29 |
+
&self,
|
| 30 |
+
mut params: HttpRequestParams,
|
| 31 |
+
) -> Result<HttpRequestResponse, ExecServerError> {
|
| 32 |
+
params.stream_response = false;
|
| 33 |
+
self.call(HTTP_REQUEST_METHOD, ¶ms).await
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
/// Performs an HTTP request and returns a body stream.
|
| 37 |
+
///
|
| 38 |
+
/// The method sets `stream_response` and replaces any caller-supplied
|
| 39 |
+
/// `request_id` with a connection-local id, so late deltas from abandoned
|
| 40 |
+
/// streams cannot be confused with later requests.
|
| 41 |
+
pub async fn http_request_stream(
|
| 42 |
+
&self,
|
| 43 |
+
mut params: HttpRequestParams,
|
| 44 |
+
) -> Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError> {
|
| 45 |
+
let rpc_client = self.rpc_client().await?;
|
| 46 |
+
params.stream_response = true;
|
| 47 |
+
let request_id = self.inner.next_http_body_stream_request_id();
|
| 48 |
+
params.request_id = request_id.clone();
|
| 49 |
+
let (tx, rx) = mpsc::channel(HTTP_BODY_DELTA_CHANNEL_CAPACITY);
|
| 50 |
+
self.inner
|
| 51 |
+
.insert_http_body_stream(request_id.clone(), tx)
|
| 52 |
+
.await?;
|
| 53 |
+
let mut registration =
|
| 54 |
+
HttpBodyStreamRegistration::new(Arc::clone(&self.inner), request_id.clone());
|
| 55 |
+
let response = match self
|
| 56 |
+
.call_rpc(&rpc_client, HTTP_REQUEST_METHOD, ¶ms)
|
| 57 |
+
.await
|
| 58 |
+
{
|
| 59 |
+
Ok(response) => response,
|
| 60 |
+
Err(error) => {
|
| 61 |
+
self.inner.remove_http_body_stream(&request_id).await;
|
| 62 |
+
registration.disarm();
|
| 63 |
+
return Err(error);
|
| 64 |
+
}
|
| 65 |
+
};
|
| 66 |
+
registration.disarm();
|
| 67 |
+
Ok((
|
| 68 |
+
response,
|
| 69 |
+
HttpResponseBodyStream::remote(Arc::clone(&self.inner), request_id, rx),
|
| 70 |
+
))
|
| 71 |
+
}
|
| 72 |
+
}
|
| 73 |
+
|
| 74 |
+
impl HttpClient for ExecServerClient {
|
| 75 |
+
/// Orchestrator-side adapter that forwards buffered HTTP requests to the
|
| 76 |
+
/// remote runtime over the shared JSON-RPC connection.
|
| 77 |
+
fn http_request(
|
| 78 |
+
&self,
|
| 79 |
+
params: HttpRequestParams,
|
| 80 |
+
) -> BoxFuture<'_, Result<HttpRequestResponse, ExecServerError>> {
|
| 81 |
+
async move { ExecServerClient::http_request(self, params).await }.boxed()
|
| 82 |
+
}
|
| 83 |
+
|
| 84 |
+
/// Orchestrator-side adapter that forwards streamed HTTP requests to the
|
| 85 |
+
/// remote runtime and exposes body deltas as a byte stream.
|
| 86 |
+
fn http_request_stream(
|
| 87 |
+
&self,
|
| 88 |
+
params: HttpRequestParams,
|
| 89 |
+
) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> {
|
| 90 |
+
async move { ExecServerClient::http_request_stream(self, params).await }.boxed()
|
| 91 |
+
}
|
| 92 |
+
}
|
codex-rs/exec-server/src/client/tests/network_policy_tests.rs
ADDED
|
@@ -0,0 +1,544 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::sync::Arc;
|
| 2 |
+
use std::sync::Mutex;
|
| 3 |
+
use std::time::Duration;
|
| 4 |
+
|
| 5 |
+
use codex_exec_server_protocol::JSONRPCMessage;
|
| 6 |
+
use codex_exec_server_protocol::JSONRPCNotification;
|
| 7 |
+
use codex_exec_server_protocol::JSONRPCRequest;
|
| 8 |
+
use codex_exec_server_protocol::JSONRPCResponse;
|
| 9 |
+
use codex_exec_server_protocol::RequestId;
|
| 10 |
+
use codex_http_client::HttpClientFactory;
|
| 11 |
+
use codex_http_client::OutboundProxyPolicy;
|
| 12 |
+
use codex_network_proxy::NetworkDecision;
|
| 13 |
+
use codex_network_proxy::NetworkPolicyDecider;
|
| 14 |
+
use codex_network_proxy::NetworkPolicyRequest;
|
| 15 |
+
use codex_network_proxy::NetworkProxyAuditMetadata;
|
| 16 |
+
use codex_utils_path_uri::PathUri;
|
| 17 |
+
use http::HeaderMap;
|
| 18 |
+
use opentelemetry::trace::TracerProvider as _;
|
| 19 |
+
use opentelemetry_sdk::trace::InMemorySpanExporter;
|
| 20 |
+
use opentelemetry_sdk::trace::SdkTracerProvider;
|
| 21 |
+
use pretty_assertions::assert_eq;
|
| 22 |
+
use tokio::net::TcpListener;
|
| 23 |
+
use tokio::sync::mpsc;
|
| 24 |
+
use tokio::sync::oneshot;
|
| 25 |
+
use tokio::time::timeout;
|
| 26 |
+
use tracing::instrument::WithSubscriber;
|
| 27 |
+
use tracing_subscriber::filter::filter_fn;
|
| 28 |
+
use tracing_subscriber::prelude::*;
|
| 29 |
+
|
| 30 |
+
use super::super::LazyRemoteExecServerClient;
|
| 31 |
+
use super::super::NetworkPolicyAuditContext;
|
| 32 |
+
use super::super::NetworkPolicyDecisionController;
|
| 33 |
+
use super::super::SessionState;
|
| 34 |
+
use super::super::handle_server_notification;
|
| 35 |
+
use super::accept_websocket;
|
| 36 |
+
use super::complete_websocket_initialize;
|
| 37 |
+
use super::read_jsonrpc_websocket;
|
| 38 |
+
use super::write_jsonrpc_websocket;
|
| 39 |
+
use crate::ProcessId;
|
| 40 |
+
use crate::client_api::ExecServerTransportParams;
|
| 41 |
+
use crate::protocol::EXEC_METHOD;
|
| 42 |
+
use crate::protocol::EXEC_TERMINATE_METHOD;
|
| 43 |
+
use crate::protocol::ExecParams;
|
| 44 |
+
use crate::protocol::ExecServerNetworkPolicyDecision;
|
| 45 |
+
use crate::protocol::ExecServerNetworkPolicyRequest;
|
| 46 |
+
use crate::protocol::ExecServerNetworkProtocol;
|
| 47 |
+
use crate::protocol::NETWORK_POLICY_DECISION_METHOD;
|
| 48 |
+
use crate::protocol::NETWORK_POLICY_REQUEST_METHOD;
|
| 49 |
+
use crate::protocol::NetworkPolicyDecisionNotification;
|
| 50 |
+
use crate::protocol::NetworkPolicyRequestParams;
|
| 51 |
+
use crate::protocol::NetworkPolicyRequestResponse;
|
| 52 |
+
use crate::rpc_server_requests::MAX_IN_FLIGHT_SERVER_CALLS;
|
| 53 |
+
|
| 54 |
+
struct PendingDecisionGuard(mpsc::UnboundedSender<()>);
|
| 55 |
+
|
| 56 |
+
impl Drop for PendingDecisionGuard {
|
| 57 |
+
fn drop(&mut self) {
|
| 58 |
+
let _ = self.0.send(());
|
| 59 |
+
}
|
| 60 |
+
}
|
| 61 |
+
|
| 62 |
+
fn policy_request(request_id: i64, process_id: ProcessId, host: &str) -> JSONRPCMessage {
|
| 63 |
+
JSONRPCMessage::Request(JSONRPCRequest {
|
| 64 |
+
id: RequestId::Integer(request_id),
|
| 65 |
+
method: NETWORK_POLICY_REQUEST_METHOD.to_string(),
|
| 66 |
+
params: Some(
|
| 67 |
+
serde_json::to_value(NetworkPolicyRequestParams {
|
| 68 |
+
process_id,
|
| 69 |
+
request: ExecServerNetworkPolicyRequest {
|
| 70 |
+
protocol: ExecServerNetworkProtocol::HttpsConnect,
|
| 71 |
+
host: host.to_string(),
|
| 72 |
+
port: 443,
|
| 73 |
+
},
|
| 74 |
+
})
|
| 75 |
+
.expect("policy request should serialize"),
|
| 76 |
+
),
|
| 77 |
+
trace: None,
|
| 78 |
+
})
|
| 79 |
+
}
|
| 80 |
+
|
| 81 |
+
async fn read_decision(
|
| 82 |
+
websocket: &mut tokio_tungstenite::WebSocketStream<tokio::net::TcpStream>,
|
| 83 |
+
request_id: i64,
|
| 84 |
+
) -> ExecServerNetworkPolicyDecision {
|
| 85 |
+
let JSONRPCMessage::Response(response) = read_jsonrpc_websocket(websocket).await else {
|
| 86 |
+
panic!("expected network policy response");
|
| 87 |
+
};
|
| 88 |
+
assert_eq!(response.id, RequestId::Integer(request_id));
|
| 89 |
+
serde_json::from_value::<NetworkPolicyRequestResponse>(response.result)
|
| 90 |
+
.expect("policy response should deserialize")
|
| 91 |
+
.decision
|
| 92 |
+
}
|
| 93 |
+
|
| 94 |
+
#[tokio::test(flavor = "current_thread")]
|
| 95 |
+
async fn policy_decisions_reject_forged_process_and_use_trusted_controller_metadata() {
|
| 96 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 97 |
+
.await
|
| 98 |
+
.expect("listener should bind");
|
| 99 |
+
let websocket_url = format!("ws://{}", listener.local_addr().expect("listener address"));
|
| 100 |
+
let (release_tx, release_rx) = oneshot::channel();
|
| 101 |
+
let (initialized_tx, initialized_rx) = oneshot::channel();
|
| 102 |
+
let server = tokio::spawn(async move {
|
| 103 |
+
let mut websocket = accept_websocket(&listener).await;
|
| 104 |
+
complete_websocket_initialize(
|
| 105 |
+
&mut websocket,
|
| 106 |
+
"audit-session",
|
| 107 |
+
/*expected_resume_session_id*/ None,
|
| 108 |
+
)
|
| 109 |
+
.await;
|
| 110 |
+
initialized_tx
|
| 111 |
+
.send(())
|
| 112 |
+
.expect("client should await completed WebSocket initialization");
|
| 113 |
+
release_rx.await.expect("server should be released");
|
| 114 |
+
});
|
| 115 |
+
|
| 116 |
+
let logs = Arc::new(Mutex::new(Vec::new()));
|
| 117 |
+
let writer_logs = Arc::clone(&logs);
|
| 118 |
+
let subscriber = tracing_subscriber::registry().with(
|
| 119 |
+
tracing_subscriber::fmt::layer()
|
| 120 |
+
.with_ansi(false)
|
| 121 |
+
.with_writer(move || AuditLogWriter(Arc::clone(&writer_logs))),
|
| 122 |
+
);
|
| 123 |
+
async move {
|
| 124 |
+
let client = LazyRemoteExecServerClient::new(
|
| 125 |
+
ExecServerTransportParams::websocket_url(websocket_url, Duration::from_secs(1)),
|
| 126 |
+
HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault),
|
| 127 |
+
)
|
| 128 |
+
.get()
|
| 129 |
+
.await
|
| 130 |
+
.expect("client should connect");
|
| 131 |
+
initialized_rx
|
| 132 |
+
.await
|
| 133 |
+
.expect("server should complete WebSocket initialization");
|
| 134 |
+
let mut state = SessionState::new(/*recoverable*/ true);
|
| 135 |
+
state.network_policy.audit = Some(NetworkPolicyAuditContext {
|
| 136 |
+
metadata: NetworkProxyAuditMetadata {
|
| 137 |
+
conversation_id: Some("trusted-conversation".to_string()),
|
| 138 |
+
user_account_id: Some("trusted-account".to_string()),
|
| 139 |
+
..NetworkProxyAuditMetadata::default()
|
| 140 |
+
},
|
| 141 |
+
execution_id: Some("trusted-execution".to_string()),
|
| 142 |
+
});
|
| 143 |
+
client
|
| 144 |
+
.inner
|
| 145 |
+
.insert_session(&ProcessId::from("trusted-process"), Arc::new(state))
|
| 146 |
+
.expect("trusted process should register");
|
| 147 |
+
for (process_id, host) in [
|
| 148 |
+
("forged-process", "forged.example"),
|
| 149 |
+
("trusted-process", "trusted.example"),
|
| 150 |
+
] {
|
| 151 |
+
handle_server_notification(
|
| 152 |
+
&client.inner,
|
| 153 |
+
JSONRPCNotification {
|
| 154 |
+
method: NETWORK_POLICY_DECISION_METHOD.to_string(),
|
| 155 |
+
params: Some(
|
| 156 |
+
serde_json::to_value(NetworkPolicyDecisionNotification {
|
| 157 |
+
process_id: ProcessId::from(process_id),
|
| 158 |
+
timestamp: "2026-08-11T12:00:00.000Z".to_string(),
|
| 159 |
+
scope: "domain".to_string(),
|
| 160 |
+
decision: "deny".to_string(),
|
| 161 |
+
source: "baseline_policy".to_string(),
|
| 162 |
+
reason: "not_allowed".to_string(),
|
| 163 |
+
protocol: ExecServerNetworkProtocol::HttpsConnect,
|
| 164 |
+
host: host.to_string(),
|
| 165 |
+
port: 443,
|
| 166 |
+
method: None,
|
| 167 |
+
client: None,
|
| 168 |
+
policy_override: false,
|
| 169 |
+
})
|
| 170 |
+
.expect("network policy decision should serialize"),
|
| 171 |
+
),
|
| 172 |
+
},
|
| 173 |
+
)
|
| 174 |
+
.await
|
| 175 |
+
.expect("controller should handle network policy notification");
|
| 176 |
+
}
|
| 177 |
+
let output = String::from_utf8(
|
| 178 |
+
logs.lock()
|
| 179 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner)
|
| 180 |
+
.clone(),
|
| 181 |
+
)
|
| 182 |
+
.expect("audit log should be UTF-8");
|
| 183 |
+
assert!(!output.contains("forged.example"));
|
| 184 |
+
for expected in [
|
| 185 |
+
"codex_otel.log_only",
|
| 186 |
+
"trusted-conversation",
|
| 187 |
+
"trusted-account",
|
| 188 |
+
"trusted-execution",
|
| 189 |
+
] {
|
| 190 |
+
assert!(
|
| 191 |
+
output.contains(expected),
|
| 192 |
+
"missing `{expected}` in {output}"
|
| 193 |
+
);
|
| 194 |
+
}
|
| 195 |
+
release_tx.send(()).expect("server should be released");
|
| 196 |
+
}
|
| 197 |
+
.with_subscriber(subscriber)
|
| 198 |
+
.await;
|
| 199 |
+
server.await.expect("server should finish");
|
| 200 |
+
}
|
| 201 |
+
|
| 202 |
+
struct AuditLogWriter(Arc<Mutex<Vec<u8>>>);
|
| 203 |
+
|
| 204 |
+
impl std::io::Write for AuditLogWriter {
|
| 205 |
+
fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
|
| 206 |
+
self.0
|
| 207 |
+
.lock()
|
| 208 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner)
|
| 209 |
+
.extend_from_slice(bytes);
|
| 210 |
+
Ok(bytes.len())
|
| 211 |
+
}
|
| 212 |
+
|
| 213 |
+
fn flush(&mut self) -> std::io::Result<()> {
|
| 214 |
+
Ok(())
|
| 215 |
+
}
|
| 216 |
+
}
|
| 217 |
+
|
| 218 |
+
#[tokio::test]
|
| 219 |
+
async fn abandoned_process_start_unregisters_and_cleans_up() {
|
| 220 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 221 |
+
.await
|
| 222 |
+
.expect("listener should bind");
|
| 223 |
+
let websocket_url = format!("ws://{}", listener.local_addr().expect("listener address"));
|
| 224 |
+
let (start_seen_tx, start_seen_rx) = oneshot::channel();
|
| 225 |
+
let (finish_start_tx, finish_start_rx) = oneshot::channel();
|
| 226 |
+
let server = tokio::spawn(async move {
|
| 227 |
+
let mut websocket = accept_websocket(&listener).await;
|
| 228 |
+
complete_websocket_initialize(&mut websocket, "p", Default::default()).await;
|
| 229 |
+
let JSONRPCMessage::Request(start) = read_jsonrpc_websocket(&mut websocket).await else {
|
| 230 |
+
panic!("expected process start request");
|
| 231 |
+
};
|
| 232 |
+
assert_eq!(start.method, EXEC_METHOD);
|
| 233 |
+
start_seen_tx.send(()).expect("start should be observed");
|
| 234 |
+
finish_start_rx.await.expect("start should be released");
|
| 235 |
+
write_jsonrpc_websocket(
|
| 236 |
+
&mut websocket,
|
| 237 |
+
JSONRPCMessage::Response(JSONRPCResponse {
|
| 238 |
+
id: start.id,
|
| 239 |
+
result: serde_json::json!({"processId": "pending-start"}),
|
| 240 |
+
}),
|
| 241 |
+
)
|
| 242 |
+
.await;
|
| 243 |
+
let JSONRPCMessage::Request(terminate) = read_jsonrpc_websocket(&mut websocket).await
|
| 244 |
+
else {
|
| 245 |
+
panic!("expected process terminate request");
|
| 246 |
+
};
|
| 247 |
+
assert_eq!(terminate.method, EXEC_TERMINATE_METHOD);
|
| 248 |
+
write_jsonrpc_websocket(
|
| 249 |
+
&mut websocket,
|
| 250 |
+
JSONRPCMessage::Response(JSONRPCResponse {
|
| 251 |
+
id: terminate.id,
|
| 252 |
+
result: serde_json::json!({"running": true}),
|
| 253 |
+
}),
|
| 254 |
+
)
|
| 255 |
+
.await;
|
| 256 |
+
});
|
| 257 |
+
|
| 258 |
+
let client = LazyRemoteExecServerClient::new(
|
| 259 |
+
ExecServerTransportParams::websocket_url(websocket_url, Duration::from_secs(1)),
|
| 260 |
+
HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault),
|
| 261 |
+
)
|
| 262 |
+
.get()
|
| 263 |
+
.await
|
| 264 |
+
.expect("client should connect");
|
| 265 |
+
let process_id = ProcessId::from("pending-start");
|
| 266 |
+
let start_client = client.clone();
|
| 267 |
+
let start_process_id = process_id.clone();
|
| 268 |
+
let start = tokio::spawn(async move {
|
| 269 |
+
let params = ExecParams {
|
| 270 |
+
metadata: Default::default(),
|
| 271 |
+
process_id: start_process_id,
|
| 272 |
+
argv: vec!["true".to_string()],
|
| 273 |
+
cwd: PathUri::from_host_native_path(std::env::current_dir().expect("cwd"))
|
| 274 |
+
.expect("cwd URI"),
|
| 275 |
+
shell_snapshot: None,
|
| 276 |
+
env_policy: None,
|
| 277 |
+
env: Default::default(),
|
| 278 |
+
tty: false,
|
| 279 |
+
pipe_stdin: false,
|
| 280 |
+
arg0: None,
|
| 281 |
+
sandbox: None,
|
| 282 |
+
enforce_managed_network: false,
|
| 283 |
+
managed_network: None,
|
| 284 |
+
network_proxy: None,
|
| 285 |
+
};
|
| 286 |
+
start_client
|
| 287 |
+
.start_process(params, /*network_policy_decider*/ None)
|
| 288 |
+
.await
|
| 289 |
+
});
|
| 290 |
+
start_seen_rx.await.expect("start should be observed");
|
| 291 |
+
let state = client
|
| 292 |
+
.inner
|
| 293 |
+
.get_session(&process_id)
|
| 294 |
+
.expect("pending process should be registered");
|
| 295 |
+
let decider: Arc<dyn NetworkPolicyDecider> =
|
| 296 |
+
Arc::new(|_request: NetworkPolicyRequest| async { NetworkDecision::Allow });
|
| 297 |
+
let decider_weak = Arc::downgrade(&decider);
|
| 298 |
+
state
|
| 299 |
+
.network_policy
|
| 300 |
+
.controller
|
| 301 |
+
.store(Some(Arc::new(NetworkPolicyDecisionController {
|
| 302 |
+
decider,
|
| 303 |
+
timeout: Duration::from_secs(30),
|
| 304 |
+
})));
|
| 305 |
+
|
| 306 |
+
start.abort();
|
| 307 |
+
assert!(start.await.is_err_and(|error| error.is_cancelled()));
|
| 308 |
+
assert!(state.network_policy.cancelled.is_cancelled());
|
| 309 |
+
assert!(client.inner.get_session(&process_id).is_none());
|
| 310 |
+
assert!(decider_weak.upgrade().is_none());
|
| 311 |
+
|
| 312 |
+
finish_start_tx.send(()).expect("start should be released");
|
| 313 |
+
server.await.expect("server task should finish");
|
| 314 |
+
}
|
| 315 |
+
|
| 316 |
+
#[tokio::test]
|
| 317 |
+
async fn policy_requests_use_process_decider_and_cancel_on_unregister() {
|
| 318 |
+
let span_exporter = InMemorySpanExporter::default();
|
| 319 |
+
let tracer_provider = SdkTracerProvider::builder()
|
| 320 |
+
.with_simple_exporter(span_exporter.clone())
|
| 321 |
+
.build();
|
| 322 |
+
let subscriber = tracing_subscriber::registry().with(
|
| 323 |
+
tracing_opentelemetry::layer()
|
| 324 |
+
.with_tracer(tracer_provider.tracer("exec-server-test"))
|
| 325 |
+
.with_filter(filter_fn(codex_otel::OtelProvider::trace_export_filter)),
|
| 326 |
+
);
|
| 327 |
+
let _subscriber = tracing::subscriber::set_default(subscriber);
|
| 328 |
+
tracing::callsite::rebuild_interest_cache();
|
| 329 |
+
|
| 330 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 331 |
+
.await
|
| 332 |
+
.expect("listener should bind");
|
| 333 |
+
let websocket_url = format!(
|
| 334 |
+
"ws://{}",
|
| 335 |
+
listener.local_addr().expect("listener should have address")
|
| 336 |
+
);
|
| 337 |
+
let process_id = ProcessId::from("policy-process");
|
| 338 |
+
let server_process_id = process_id.clone();
|
| 339 |
+
let (ready_tx, ready_rx) = oneshot::channel();
|
| 340 |
+
let (overflow_checked_tx, overflow_checked_rx) = oneshot::channel();
|
| 341 |
+
let (unregistered_tx, unregistered_rx) = oneshot::channel();
|
| 342 |
+
let server = tokio::spawn(async move {
|
| 343 |
+
let mut websocket = accept_websocket(&listener).await;
|
| 344 |
+
complete_websocket_initialize(
|
| 345 |
+
&mut websocket,
|
| 346 |
+
"policy-session",
|
| 347 |
+
/*expected_resume_session_id*/ None,
|
| 348 |
+
)
|
| 349 |
+
.await;
|
| 350 |
+
ready_rx.await.expect("process should be registered");
|
| 351 |
+
|
| 352 |
+
for (request_id, host, expected) in [
|
| 353 |
+
(0, "allowed.example", ExecServerNetworkPolicyDecision::Allow),
|
| 354 |
+
(
|
| 355 |
+
2,
|
| 356 |
+
"denied.example",
|
| 357 |
+
ExecServerNetworkPolicyDecision::Deny {
|
| 358 |
+
reason: "blocked".to_string(),
|
| 359 |
+
},
|
| 360 |
+
),
|
| 361 |
+
(
|
| 362 |
+
3,
|
| 363 |
+
"invalid host",
|
| 364 |
+
ExecServerNetworkPolicyDecision::Deny {
|
| 365 |
+
reason: "not_allowed".to_string(),
|
| 366 |
+
},
|
| 367 |
+
),
|
| 368 |
+
] {
|
| 369 |
+
write_jsonrpc_websocket(
|
| 370 |
+
&mut websocket,
|
| 371 |
+
policy_request(request_id, server_process_id.clone(), host),
|
| 372 |
+
)
|
| 373 |
+
.await;
|
| 374 |
+
assert_eq!(read_decision(&mut websocket, request_id).await, expected);
|
| 375 |
+
}
|
| 376 |
+
|
| 377 |
+
let first_pending_request_id = 100;
|
| 378 |
+
for offset in 0..MAX_IN_FLIGHT_SERVER_CALLS {
|
| 379 |
+
write_jsonrpc_websocket(
|
| 380 |
+
&mut websocket,
|
| 381 |
+
policy_request(
|
| 382 |
+
first_pending_request_id + offset as i64,
|
| 383 |
+
server_process_id.clone(),
|
| 384 |
+
"pending.example",
|
| 385 |
+
),
|
| 386 |
+
)
|
| 387 |
+
.await;
|
| 388 |
+
}
|
| 389 |
+
let overflow_request_id = first_pending_request_id + MAX_IN_FLIGHT_SERVER_CALLS as i64;
|
| 390 |
+
write_jsonrpc_websocket(
|
| 391 |
+
&mut websocket,
|
| 392 |
+
policy_request(
|
| 393 |
+
overflow_request_id,
|
| 394 |
+
server_process_id.clone(),
|
| 395 |
+
"pending.example",
|
| 396 |
+
),
|
| 397 |
+
)
|
| 398 |
+
.await;
|
| 399 |
+
assert_eq!(
|
| 400 |
+
read_decision(&mut websocket, overflow_request_id).await,
|
| 401 |
+
ExecServerNetworkPolicyDecision::Deny {
|
| 402 |
+
reason: "not_allowed".to_string(),
|
| 403 |
+
}
|
| 404 |
+
);
|
| 405 |
+
overflow_checked_tx.send(()).expect("overflow observed");
|
| 406 |
+
|
| 407 |
+
unregistered_rx
|
| 408 |
+
.await
|
| 409 |
+
.expect("process should be unregistered");
|
| 410 |
+
|
| 411 |
+
let post_unregister_request_id = 900;
|
| 412 |
+
write_jsonrpc_websocket(
|
| 413 |
+
&mut websocket,
|
| 414 |
+
policy_request(
|
| 415 |
+
post_unregister_request_id,
|
| 416 |
+
server_process_id,
|
| 417 |
+
"allowed.example",
|
| 418 |
+
),
|
| 419 |
+
)
|
| 420 |
+
.await;
|
| 421 |
+
assert_eq!(
|
| 422 |
+
read_decision(&mut websocket, post_unregister_request_id).await,
|
| 423 |
+
ExecServerNetworkPolicyDecision::Deny {
|
| 424 |
+
reason: "not_allowed".to_string(),
|
| 425 |
+
}
|
| 426 |
+
);
|
| 427 |
+
});
|
| 428 |
+
|
| 429 |
+
let client = LazyRemoteExecServerClient::new(
|
| 430 |
+
ExecServerTransportParams::WebSocketUrl {
|
| 431 |
+
websocket_url,
|
| 432 |
+
connect_timeout: Duration::from_secs(1),
|
| 433 |
+
initialize_timeout: Duration::from_secs(1),
|
| 434 |
+
http_headers: HeaderMap::new(),
|
| 435 |
+
},
|
| 436 |
+
HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault),
|
| 437 |
+
)
|
| 438 |
+
.get()
|
| 439 |
+
.await
|
| 440 |
+
.expect("client should connect");
|
| 441 |
+
let session = client
|
| 442 |
+
.register_session(&process_id)
|
| 443 |
+
.await
|
| 444 |
+
.expect("session should register");
|
| 445 |
+
let (started_tx, mut started_rx) = mpsc::unbounded_channel();
|
| 446 |
+
let (dropped_tx, mut dropped_rx) = mpsc::unbounded_channel();
|
| 447 |
+
let decider: Arc<dyn NetworkPolicyDecider> = Arc::new(move |request: NetworkPolicyRequest| {
|
| 448 |
+
let started_tx = started_tx.clone();
|
| 449 |
+
let dropped_tx = dropped_tx.clone();
|
| 450 |
+
async move {
|
| 451 |
+
assert_eq!(
|
| 452 |
+
tracing::Span::current()
|
| 453 |
+
.metadata()
|
| 454 |
+
.map(tracing::Metadata::name),
|
| 455 |
+
Some("codex.exec_server.request"),
|
| 456 |
+
"network policy decisions must run inside the inbound request span"
|
| 457 |
+
);
|
| 458 |
+
match request.host.as_str() {
|
| 459 |
+
"allowed.example" => NetworkDecision::Allow,
|
| 460 |
+
"denied.example" => NetworkDecision::deny("blocked"),
|
| 461 |
+
"pending.example" => {
|
| 462 |
+
started_tx.send(()).expect("decision should start");
|
| 463 |
+
let _drop_guard = PendingDecisionGuard(dropped_tx);
|
| 464 |
+
std::future::pending().await
|
| 465 |
+
}
|
| 466 |
+
host => panic!("unexpected policy host: {host}"),
|
| 467 |
+
}
|
| 468 |
+
}
|
| 469 |
+
});
|
| 470 |
+
session.state.network_policy.controller.store(Some(Arc::new(
|
| 471 |
+
NetworkPolicyDecisionController {
|
| 472 |
+
decider,
|
| 473 |
+
timeout: Duration::from_secs(30),
|
| 474 |
+
},
|
| 475 |
+
)));
|
| 476 |
+
ready_tx.send(()).expect("server should be waiting");
|
| 477 |
+
timeout(Duration::from_secs(5), async {
|
| 478 |
+
for _ in 0..MAX_IN_FLIGHT_SERVER_CALLS {
|
| 479 |
+
started_rx
|
| 480 |
+
.recv()
|
| 481 |
+
.await
|
| 482 |
+
.expect("pending decision should start");
|
| 483 |
+
}
|
| 484 |
+
})
|
| 485 |
+
.await
|
| 486 |
+
.expect("pending decisions should start");
|
| 487 |
+
overflow_checked_rx
|
| 488 |
+
.await
|
| 489 |
+
.expect("overflow should be observed");
|
| 490 |
+
session.unregister().await;
|
| 491 |
+
timeout(Duration::from_secs(5), async {
|
| 492 |
+
for _ in 0..MAX_IN_FLIGHT_SERVER_CALLS {
|
| 493 |
+
dropped_rx
|
| 494 |
+
.recv()
|
| 495 |
+
.await
|
| 496 |
+
.expect("unregistered decision should be dropped");
|
| 497 |
+
}
|
| 498 |
+
})
|
| 499 |
+
.await
|
| 500 |
+
.expect("unregistered decisions should be cancelled");
|
| 501 |
+
unregistered_tx
|
| 502 |
+
.send(())
|
| 503 |
+
.expect("server should verify late responses");
|
| 504 |
+
timeout(Duration::from_secs(2), server)
|
| 505 |
+
.await
|
| 506 |
+
.expect("policy routing should finish")
|
| 507 |
+
.expect("server task should finish");
|
| 508 |
+
|
| 509 |
+
tracer_provider.force_flush().expect("flush traces");
|
| 510 |
+
let spans = span_exporter.get_finished_spans().expect("span export");
|
| 511 |
+
let policy_spans = spans
|
| 512 |
+
.iter()
|
| 513 |
+
.filter(|span| span.name.as_ref() == NETWORK_POLICY_REQUEST_METHOD)
|
| 514 |
+
.collect::<Vec<_>>();
|
| 515 |
+
assert!(
|
| 516 |
+
!policy_spans.is_empty(),
|
| 517 |
+
"network policy requests should export server spans"
|
| 518 |
+
);
|
| 519 |
+
let outcomes = policy_spans
|
| 520 |
+
.iter()
|
| 521 |
+
.map(|span| {
|
| 522 |
+
span.attributes
|
| 523 |
+
.iter()
|
| 524 |
+
.find(|attribute| attribute.key.as_str() == "result")
|
| 525 |
+
.map(|attribute| attribute.value.as_str().into_owned())
|
| 526 |
+
})
|
| 527 |
+
.collect::<Vec<_>>();
|
| 528 |
+
assert!(
|
| 529 |
+
outcomes.iter().all(Option::is_some),
|
| 530 |
+
"completed, rejected, and cancelled policy requests must all record an outcome"
|
| 531 |
+
);
|
| 532 |
+
assert!(
|
| 533 |
+
outcomes
|
| 534 |
+
.iter()
|
| 535 |
+
.any(|outcome| outcome.as_deref() == Some("success")),
|
| 536 |
+
"completed and capacity-rejected requests should record successful responses"
|
| 537 |
+
);
|
| 538 |
+
assert!(
|
| 539 |
+
outcomes
|
| 540 |
+
.iter()
|
| 541 |
+
.any(|outcome| outcome.as_deref() == Some("disconnected")),
|
| 542 |
+
"cancelled requests should record disconnection"
|
| 543 |
+
);
|
| 544 |
+
}
|
codex-rs/exec-server/src/client_api.rs
ADDED
|
@@ -0,0 +1,187 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::collections::HashMap;
|
| 2 |
+
use std::path::PathBuf;
|
| 3 |
+
use std::sync::Arc;
|
| 4 |
+
use std::time::Duration;
|
| 5 |
+
|
| 6 |
+
use codex_http_client::HttpClientFactory;
|
| 7 |
+
use futures::future::BoxFuture;
|
| 8 |
+
use http::HeaderMap;
|
| 9 |
+
use tokio::sync::watch;
|
| 10 |
+
|
| 11 |
+
use crate::ExecServerError;
|
| 12 |
+
use crate::HttpRequestParams;
|
| 13 |
+
use crate::HttpRequestResponse;
|
| 14 |
+
use crate::HttpResponseBodyStream;
|
| 15 |
+
use crate::NoiseChannelIdentity;
|
| 16 |
+
use crate::NoiseChannelPublicKey;
|
| 17 |
+
|
| 18 |
+
pub(crate) const DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
|
| 19 |
+
pub(crate) const DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT: Duration = Duration::from_secs(10);
|
| 20 |
+
|
| 21 |
+
/// Connection options for any exec-server client transport.
|
| 22 |
+
#[derive(Debug, Clone, PartialEq, Eq)]
|
| 23 |
+
pub struct ExecServerClientConnectOptions {
|
| 24 |
+
pub client_name: String,
|
| 25 |
+
pub initialize_timeout: Duration,
|
| 26 |
+
pub resume_session_id: Option<String>,
|
| 27 |
+
}
|
| 28 |
+
|
| 29 |
+
/// WebSocket connection arguments for a remote exec-server.
|
| 30 |
+
#[derive(Debug, Clone, PartialEq, Eq)]
|
| 31 |
+
pub struct RemoteExecServerConnectArgs {
|
| 32 |
+
pub websocket_url: String,
|
| 33 |
+
pub client_name: String,
|
| 34 |
+
pub connect_timeout: Duration,
|
| 35 |
+
pub initialize_timeout: Duration,
|
| 36 |
+
pub resume_session_id: Option<String>,
|
| 37 |
+
pub http_client_factory: HttpClientFactory,
|
| 38 |
+
}
|
| 39 |
+
|
| 40 |
+
/// Registry-authorized material for one Noise rendezvous connection attempt.
|
| 41 |
+
///
|
| 42 |
+
/// Treat this as an atomic, single-use bundle. The URL authorization, executor
|
| 43 |
+
/// registration, pinned executor key, and harness-key authorization describe one
|
| 44 |
+
/// physical connection attempt and must not be mixed with values from another
|
| 45 |
+
/// registry response.
|
| 46 |
+
pub struct NoiseRendezvousConnectBundle {
|
| 47 |
+
pub websocket_url: String,
|
| 48 |
+
pub environment_id: String,
|
| 49 |
+
pub executor_registration_id: String,
|
| 50 |
+
pub executor_public_key: NoiseChannelPublicKey,
|
| 51 |
+
pub harness_key_authorization: String,
|
| 52 |
+
}
|
| 53 |
+
|
| 54 |
+
/// Connection arguments for an authenticated Noise rendezvous exec-server.
|
| 55 |
+
///
|
| 56 |
+
/// `harness_identity` identifies the logical harness endpoint and may be reused
|
| 57 |
+
/// across reconnects. In contrast, callers must supply a fresh
|
| 58 |
+
/// [`NoiseRendezvousConnectBundle`] for each physical connection attempt.
|
| 59 |
+
pub struct NoiseRendezvousConnectArgs {
|
| 60 |
+
pub bundle: NoiseRendezvousConnectBundle,
|
| 61 |
+
pub harness_identity: NoiseChannelIdentity,
|
| 62 |
+
pub client_name: String,
|
| 63 |
+
pub connect_timeout: Duration,
|
| 64 |
+
pub initialize_timeout: Duration,
|
| 65 |
+
pub resume_session_id: Option<String>,
|
| 66 |
+
pub http_client_factory: HttpClientFactory,
|
| 67 |
+
}
|
| 68 |
+
|
| 69 |
+
/// Supplies fresh registry-authorized material for Noise rendezvous connections.
|
| 70 |
+
pub trait NoiseRendezvousConnectProvider: Send + Sync {
|
| 71 |
+
/// Fetch a bundle authorizing this harness key for one physical connection.
|
| 72 |
+
fn connect_bundle(
|
| 73 |
+
&self,
|
| 74 |
+
harness_public_key: NoiseChannelPublicKey,
|
| 75 |
+
) -> BoxFuture<'_, Result<NoiseRendezvousConnectBundle, ExecServerError>>;
|
| 76 |
+
}
|
| 77 |
+
|
| 78 |
+
/// Stdio connection arguments for a command-backed exec-server.
|
| 79 |
+
#[derive(Debug, Clone, PartialEq, Eq)]
|
| 80 |
+
pub(crate) struct StdioExecServerConnectArgs {
|
| 81 |
+
pub command: StdioExecServerCommand,
|
| 82 |
+
pub client_name: String,
|
| 83 |
+
pub initialize_timeout: Duration,
|
| 84 |
+
pub resume_session_id: Option<String>,
|
| 85 |
+
}
|
| 86 |
+
|
| 87 |
+
/// Structured process command used to start an exec-server over stdio.
|
| 88 |
+
#[derive(Debug, Clone, PartialEq, Eq)]
|
| 89 |
+
pub(crate) struct StdioExecServerCommand {
|
| 90 |
+
pub program: String,
|
| 91 |
+
pub args: Vec<String>,
|
| 92 |
+
pub env: HashMap<String, String>,
|
| 93 |
+
pub cwd: Option<PathBuf>,
|
| 94 |
+
}
|
| 95 |
+
|
| 96 |
+
pub(crate) type DeferredEnvironmentReadiness = watch::Receiver<Option<Result<(), String>>>;
|
| 97 |
+
|
| 98 |
+
#[derive(Clone)]
|
| 99 |
+
pub(crate) struct Deferred<T> {
|
| 100 |
+
pub readiness: DeferredEnvironmentReadiness,
|
| 101 |
+
pub transport: T,
|
| 102 |
+
}
|
| 103 |
+
|
| 104 |
+
/// Parameters used to connect to a remote exec-server environment.
|
| 105 |
+
#[derive(Clone)]
|
| 106 |
+
pub(crate) enum ExecServerTransportParams {
|
| 107 |
+
Deferred(Box<Deferred<ExecServerTransportParams>>),
|
| 108 |
+
WebSocketUrl {
|
| 109 |
+
websocket_url: String,
|
| 110 |
+
connect_timeout: Duration,
|
| 111 |
+
initialize_timeout: Duration,
|
| 112 |
+
http_headers: HeaderMap,
|
| 113 |
+
},
|
| 114 |
+
NoiseRendezvous {
|
| 115 |
+
provider: Arc<dyn NoiseRendezvousConnectProvider>,
|
| 116 |
+
identity: NoiseChannelIdentity,
|
| 117 |
+
},
|
| 118 |
+
#[allow(dead_code)]
|
| 119 |
+
StdioCommand {
|
| 120 |
+
command: StdioExecServerCommand,
|
| 121 |
+
initialize_timeout: Duration,
|
| 122 |
+
},
|
| 123 |
+
}
|
| 124 |
+
|
| 125 |
+
impl std::fmt::Debug for ExecServerTransportParams {
|
| 126 |
+
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
| 127 |
+
match self {
|
| 128 |
+
Self::Deferred(deferred) => f
|
| 129 |
+
.debug_struct("Deferred")
|
| 130 |
+
.field("transport", &deferred.transport)
|
| 131 |
+
.finish_non_exhaustive(),
|
| 132 |
+
Self::WebSocketUrl {
|
| 133 |
+
websocket_url,
|
| 134 |
+
connect_timeout,
|
| 135 |
+
initialize_timeout,
|
| 136 |
+
..
|
| 137 |
+
} => f
|
| 138 |
+
.debug_struct("WebSocketUrl")
|
| 139 |
+
.field("websocket_url", websocket_url)
|
| 140 |
+
.field("connect_timeout", connect_timeout)
|
| 141 |
+
.field("initialize_timeout", initialize_timeout)
|
| 142 |
+
.field("http_headers", &"<redacted>")
|
| 143 |
+
.finish(),
|
| 144 |
+
Self::NoiseRendezvous { .. } => {
|
| 145 |
+
f.debug_struct("NoiseRendezvous").finish_non_exhaustive()
|
| 146 |
+
}
|
| 147 |
+
Self::StdioCommand {
|
| 148 |
+
command,
|
| 149 |
+
initialize_timeout,
|
| 150 |
+
} => f
|
| 151 |
+
.debug_struct("StdioCommand")
|
| 152 |
+
.field("command", command)
|
| 153 |
+
.field("initialize_timeout", initialize_timeout)
|
| 154 |
+
.finish(),
|
| 155 |
+
}
|
| 156 |
+
}
|
| 157 |
+
}
|
| 158 |
+
|
| 159 |
+
impl ExecServerTransportParams {
|
| 160 |
+
pub(crate) fn websocket_url(websocket_url: String, connect_timeout: Duration) -> Self {
|
| 161 |
+
Self::WebSocketUrl {
|
| 162 |
+
websocket_url,
|
| 163 |
+
connect_timeout,
|
| 164 |
+
initialize_timeout: DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT,
|
| 165 |
+
http_headers: HeaderMap::new(),
|
| 166 |
+
}
|
| 167 |
+
}
|
| 168 |
+
}
|
| 169 |
+
|
| 170 |
+
/// Sends HTTP requests through a runtime-selected transport.
|
| 171 |
+
///
|
| 172 |
+
/// This is the HTTP capability counterpart to [`crate::ExecBackend`]. Callers
|
| 173 |
+
/// use it when they need environment-owned network requests but should not
|
| 174 |
+
/// depend on the concrete connection type or how that connection is established.
|
| 175 |
+
pub trait HttpClient: Send + Sync {
|
| 176 |
+
/// Perform an HTTP request and buffer the response body.
|
| 177 |
+
fn http_request(
|
| 178 |
+
&self,
|
| 179 |
+
params: HttpRequestParams,
|
| 180 |
+
) -> BoxFuture<'_, Result<HttpRequestResponse, ExecServerError>>;
|
| 181 |
+
|
| 182 |
+
/// Perform an HTTP request and return a streamed body handle.
|
| 183 |
+
fn http_request_stream(
|
| 184 |
+
&self,
|
| 185 |
+
params: HttpRequestParams,
|
| 186 |
+
) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>>;
|
| 187 |
+
}
|
codex-rs/exec-server/src/client_recovery.rs
ADDED
|
@@ -0,0 +1,900 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::collections::hash_map::DefaultHasher;
|
| 2 |
+
use std::hash::Hash;
|
| 3 |
+
use std::hash::Hasher;
|
| 4 |
+
use std::sync::Arc;
|
| 5 |
+
use std::sync::atomic::Ordering;
|
| 6 |
+
use std::time::Duration;
|
| 7 |
+
|
| 8 |
+
use codex_network_proxy::NetworkDecision;
|
| 9 |
+
use codex_network_proxy::NetworkPolicyDecision;
|
| 10 |
+
use codex_network_proxy::NetworkPolicyRequest;
|
| 11 |
+
use codex_network_proxy::NetworkProtocol;
|
| 12 |
+
use codex_network_proxy::NetworkRequestCancellation;
|
| 13 |
+
use codex_network_proxy::NetworkRequestCancellationReason;
|
| 14 |
+
use serde_json::Value;
|
| 15 |
+
use tokio::sync::mpsc;
|
| 16 |
+
use tokio::time::Instant;
|
| 17 |
+
use tokio::time::sleep;
|
| 18 |
+
use tokio::time::timeout;
|
| 19 |
+
use tokio::time::timeout_at;
|
| 20 |
+
use tokio_util::sync::CancellationToken;
|
| 21 |
+
use tracing::Instrument;
|
| 22 |
+
use tracing::debug;
|
| 23 |
+
|
| 24 |
+
use super::ConnectionStatus;
|
| 25 |
+
use super::ExecServerClient;
|
| 26 |
+
use super::ExecServerError;
|
| 27 |
+
use super::Inner;
|
| 28 |
+
use super::OrderedSessionEvents;
|
| 29 |
+
use super::RecoveryPolicy;
|
| 30 |
+
use super::SessionState;
|
| 31 |
+
use super::disconnected_message;
|
| 32 |
+
use super::fail_all_in_flight_work;
|
| 33 |
+
use super::handle_server_notification;
|
| 34 |
+
use super::is_transport_closed_error;
|
| 35 |
+
use crate::client_transport::ExecServerReconnectStrategy;
|
| 36 |
+
use crate::process::ExecProcessEvent;
|
| 37 |
+
use crate::protocol::EXEC_READ_METHOD;
|
| 38 |
+
use crate::protocol::EXEC_TERMINATE_METHOD;
|
| 39 |
+
use crate::protocol::ExecServerNetworkPolicyDecision;
|
| 40 |
+
use crate::protocol::ExecServerNetworkProtocol;
|
| 41 |
+
use crate::protocol::MAX_NETWORK_POLICY_HOST_BYTES;
|
| 42 |
+
use crate::protocol::MAX_NETWORK_POLICY_PROCESS_ID_BYTES;
|
| 43 |
+
use crate::protocol::MAX_NETWORK_POLICY_REASON_BYTES;
|
| 44 |
+
use crate::protocol::NETWORK_POLICY_REQUEST_METHOD;
|
| 45 |
+
use crate::protocol::NetworkPolicyRequestParams;
|
| 46 |
+
use crate::protocol::NetworkPolicyRequestResponse;
|
| 47 |
+
use crate::protocol::ReadParams;
|
| 48 |
+
use crate::protocol::ReadResponse;
|
| 49 |
+
use crate::protocol::TerminateParams;
|
| 50 |
+
use crate::protocol::TerminateResponse;
|
| 51 |
+
use crate::rpc::RpcClient;
|
| 52 |
+
use crate::rpc::RpcClientEvent;
|
| 53 |
+
use crate::rpc::RpcInboundRequestAdmissionError;
|
| 54 |
+
use crate::rpc::SESSION_ALREADY_ATTACHED_ERROR_CODE;
|
| 55 |
+
use crate::rpc::invalid_params;
|
| 56 |
+
use crate::rpc::method_not_found;
|
| 57 |
+
|
| 58 |
+
#[cfg(test)]
|
| 59 |
+
const SESSION_RECOVERY_TIMEOUT: Duration = Duration::from_millis(500);
|
| 60 |
+
#[cfg(not(test))]
|
| 61 |
+
// Leave margin inside the server's 30-second retention windows because the
|
| 62 |
+
// client and server start their disconnect clocks independently.
|
| 63 |
+
const SESSION_RECOVERY_TIMEOUT: Duration = Duration::from_secs(25);
|
| 64 |
+
const SESSION_RECOVERY_RETRY_INTERVAL: Duration = Duration::from_millis(100);
|
| 65 |
+
const REGISTRY_RECOVERY_INITIAL_RETRY_INTERVAL: Duration = Duration::from_millis(500);
|
| 66 |
+
const REGISTRY_RECOVERY_MAX_RETRY_INTERVAL: Duration = Duration::from_secs(5);
|
| 67 |
+
const NETWORK_POLICY_DENIAL_REASON: &str = "not_allowed";
|
| 68 |
+
|
| 69 |
+
struct ClientRequestOutcome {
|
| 70 |
+
span: tracing::Span,
|
| 71 |
+
result: &'static str,
|
| 72 |
+
}
|
| 73 |
+
|
| 74 |
+
impl ClientRequestOutcome {
|
| 75 |
+
fn complete(&mut self, result: &'static str) {
|
| 76 |
+
self.result = result;
|
| 77 |
+
}
|
| 78 |
+
}
|
| 79 |
+
|
| 80 |
+
impl Drop for ClientRequestOutcome {
|
| 81 |
+
fn drop(&mut self) {
|
| 82 |
+
self.span.record("result", self.result);
|
| 83 |
+
}
|
| 84 |
+
}
|
| 85 |
+
|
| 86 |
+
impl SessionState {
|
| 87 |
+
fn last_published_seq(&self) -> u64 {
|
| 88 |
+
self.ordered_events
|
| 89 |
+
.lock()
|
| 90 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner)
|
| 91 |
+
.last_published_seq
|
| 92 |
+
}
|
| 93 |
+
|
| 94 |
+
fn recover_events(&self, response: ReadResponse) -> Result<bool, ExecServerError> {
|
| 95 |
+
let ReadResponse {
|
| 96 |
+
chunks,
|
| 97 |
+
next_seq,
|
| 98 |
+
exited,
|
| 99 |
+
exit_code,
|
| 100 |
+
closed,
|
| 101 |
+
failure,
|
| 102 |
+
sandbox_denied,
|
| 103 |
+
} = response;
|
| 104 |
+
if let Some(message) = failure {
|
| 105 |
+
return Err(ExecServerError::Protocol(format!(
|
| 106 |
+
"process failed while recovering: {message}"
|
| 107 |
+
)));
|
| 108 |
+
}
|
| 109 |
+
|
| 110 |
+
let target_seq = next_seq.saturating_sub(1);
|
| 111 |
+
let published_closed = {
|
| 112 |
+
let mut ordered_events = self
|
| 113 |
+
.ordered_events
|
| 114 |
+
.lock()
|
| 115 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 116 |
+
if ordered_events.failure.is_some()
|
| 117 |
+
|| ordered_events.closed_published
|
| 118 |
+
|| target_seq <= ordered_events.last_published_seq
|
| 119 |
+
{
|
| 120 |
+
return Ok(false);
|
| 121 |
+
}
|
| 122 |
+
let pending_exit = ordered_events.pending.range_mut(..=target_seq).find_map(
|
| 123 |
+
|(_, event)| match event {
|
| 124 |
+
ExecProcessEvent::Exited {
|
| 125 |
+
sandbox_denied: pending_sandbox_denied,
|
| 126 |
+
..
|
| 127 |
+
} => Some(pending_sandbox_denied),
|
| 128 |
+
_ => None,
|
| 129 |
+
},
|
| 130 |
+
);
|
| 131 |
+
let exit_pending = pending_exit.is_some();
|
| 132 |
+
if let Some(pending_sandbox_denied) = pending_exit {
|
| 133 |
+
*pending_sandbox_denied =
|
| 134 |
+
Some(pending_sandbox_denied.unwrap_or(false) || sandbox_denied);
|
| 135 |
+
}
|
| 136 |
+
let mut exit_known = ordered_events.exit_published || exit_pending;
|
| 137 |
+
if closed
|
| 138 |
+
&& (matches!(
|
| 139 |
+
ordered_events.pending.get(&target_seq),
|
| 140 |
+
Some(event) if !matches!(event, ExecProcessEvent::Closed { .. })
|
| 141 |
+
) || chunks.iter().any(|chunk| chunk.seq == target_seq))
|
| 142 |
+
{
|
| 143 |
+
return Err(ExecServerError::Protocol(format!(
|
| 144 |
+
"process close sequence {target_seq} conflicts with recovered output"
|
| 145 |
+
)));
|
| 146 |
+
}
|
| 147 |
+
let mut published_closed = false;
|
| 148 |
+
for chunk in chunks {
|
| 149 |
+
if chunk.seq > target_seq {
|
| 150 |
+
return Err(ExecServerError::Protocol(format!(
|
| 151 |
+
"recovered process output sequence {} exceeds target sequence {target_seq}",
|
| 152 |
+
chunk.seq
|
| 153 |
+
)));
|
| 154 |
+
}
|
| 155 |
+
let next_seq = ordered_events.last_published_seq.saturating_add(1);
|
| 156 |
+
if exited && !exit_known && chunk.seq > next_seq {
|
| 157 |
+
let exit_code = exit_code.ok_or_else(|| {
|
| 158 |
+
ExecServerError::Protocol(
|
| 159 |
+
"recovering exited process did not include its exit code".to_string(),
|
| 160 |
+
)
|
| 161 |
+
})?;
|
| 162 |
+
ordered_events
|
| 163 |
+
.insert_pending(ExecProcessEvent::Exited {
|
| 164 |
+
seq: next_seq,
|
| 165 |
+
exit_code,
|
| 166 |
+
sandbox_denied: Some(sandbox_denied),
|
| 167 |
+
})
|
| 168 |
+
.map_err(ExecServerError::Protocol)?;
|
| 169 |
+
published_closed |= self.publish_ready(&mut ordered_events);
|
| 170 |
+
exit_known = true;
|
| 171 |
+
}
|
| 172 |
+
if chunk.seq > ordered_events.last_published_seq {
|
| 173 |
+
ordered_events
|
| 174 |
+
.insert_pending(ExecProcessEvent::Output(chunk))
|
| 175 |
+
.map_err(ExecServerError::Protocol)?;
|
| 176 |
+
published_closed |= self.publish_ready(&mut ordered_events);
|
| 177 |
+
}
|
| 178 |
+
}
|
| 179 |
+
if closed
|
| 180 |
+
&& !ordered_events.closed_published
|
| 181 |
+
&& !matches!(
|
| 182 |
+
ordered_events.pending.get(&target_seq),
|
| 183 |
+
Some(ExecProcessEvent::Closed { .. })
|
| 184 |
+
)
|
| 185 |
+
{
|
| 186 |
+
ordered_events
|
| 187 |
+
.insert_pending(ExecProcessEvent::Closed { seq: target_seq })
|
| 188 |
+
.map_err(ExecServerError::Protocol)?;
|
| 189 |
+
}
|
| 190 |
+
|
| 191 |
+
let event_count = target_seq.saturating_sub(ordered_events.last_published_seq);
|
| 192 |
+
let first_unpublished_seq = ordered_events.last_published_seq.saturating_add(1);
|
| 193 |
+
let retained_count = if first_unpublished_seq <= target_seq {
|
| 194 |
+
ordered_events
|
| 195 |
+
.pending
|
| 196 |
+
.range(first_unpublished_seq..=target_seq)
|
| 197 |
+
.count() as u64
|
| 198 |
+
} else {
|
| 199 |
+
0
|
| 200 |
+
};
|
| 201 |
+
let missing_count = event_count.saturating_sub(retained_count);
|
| 202 |
+
if exited && !exit_known {
|
| 203 |
+
if missing_count != 1 {
|
| 204 |
+
return Err(recovery_gap_error(target_seq));
|
| 205 |
+
}
|
| 206 |
+
let seq = first_missing_seq(&ordered_events, target_seq);
|
| 207 |
+
let exit_code = exit_code.ok_or_else(|| {
|
| 208 |
+
ExecServerError::Protocol(
|
| 209 |
+
"recovering exited process did not include its exit code".to_string(),
|
| 210 |
+
)
|
| 211 |
+
})?;
|
| 212 |
+
ordered_events
|
| 213 |
+
.insert_pending(ExecProcessEvent::Exited {
|
| 214 |
+
seq,
|
| 215 |
+
exit_code,
|
| 216 |
+
sandbox_denied: Some(sandbox_denied),
|
| 217 |
+
})
|
| 218 |
+
.map_err(ExecServerError::Protocol)?;
|
| 219 |
+
} else if missing_count != 0 {
|
| 220 |
+
return Err(recovery_gap_error(target_seq));
|
| 221 |
+
}
|
| 222 |
+
published_closed |= self.publish_ready(&mut ordered_events);
|
| 223 |
+
published_closed
|
| 224 |
+
};
|
| 225 |
+
|
| 226 |
+
self.note_change(target_seq);
|
| 227 |
+
Ok(published_closed)
|
| 228 |
+
}
|
| 229 |
+
}
|
| 230 |
+
|
| 231 |
+
fn first_missing_seq(events: &OrderedSessionEvents, target_seq: u64) -> u64 {
|
| 232 |
+
let mut expected = events.last_published_seq.saturating_add(1);
|
| 233 |
+
for seq in events
|
| 234 |
+
.pending
|
| 235 |
+
.range(expected..=target_seq)
|
| 236 |
+
.map(|(seq, _)| *seq)
|
| 237 |
+
{
|
| 238 |
+
if seq != expected {
|
| 239 |
+
break;
|
| 240 |
+
}
|
| 241 |
+
expected = expected.saturating_add(1);
|
| 242 |
+
}
|
| 243 |
+
expected
|
| 244 |
+
}
|
| 245 |
+
|
| 246 |
+
fn recovery_gap_error(target_seq: u64) -> ExecServerError {
|
| 247 |
+
ExecServerError::Protocol(format!(
|
| 248 |
+
"process events are no longer retained while recovering through sequence {target_seq}"
|
| 249 |
+
))
|
| 250 |
+
}
|
| 251 |
+
|
| 252 |
+
impl Inner {
|
| 253 |
+
pub(super) async fn rpc_client(self: &Arc<Self>) -> Result<Arc<RpcClient>, ExecServerError> {
|
| 254 |
+
let mut connection_changed = self.connection_changed.subscribe();
|
| 255 |
+
loop {
|
| 256 |
+
if let Some(message) = self.failure_message() {
|
| 257 |
+
return Err(ExecServerError::Disconnected(message));
|
| 258 |
+
}
|
| 259 |
+
|
| 260 |
+
let rpc_client = {
|
| 261 |
+
let connection = self
|
| 262 |
+
.connection
|
| 263 |
+
.lock()
|
| 264 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 265 |
+
match &connection.status {
|
| 266 |
+
ConnectionStatus::Connected(rpc_client) => Some(Arc::clone(rpc_client)),
|
| 267 |
+
ConnectionStatus::Recovering | ConnectionStatus::Failed(_) => None,
|
| 268 |
+
}
|
| 269 |
+
};
|
| 270 |
+
let Some(rpc_client) = rpc_client else {
|
| 271 |
+
let _ = connection_changed.changed().await;
|
| 272 |
+
continue;
|
| 273 |
+
};
|
| 274 |
+
if !rpc_client.is_disconnected() {
|
| 275 |
+
return Ok(rpc_client);
|
| 276 |
+
}
|
| 277 |
+
|
| 278 |
+
let _ = connection_changed.changed().await;
|
| 279 |
+
}
|
| 280 |
+
}
|
| 281 |
+
|
| 282 |
+
pub(super) fn begin_process_start(&self, expected: &Arc<RpcClient>) -> bool {
|
| 283 |
+
let mut connection = self
|
| 284 |
+
.connection
|
| 285 |
+
.lock()
|
| 286 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 287 |
+
let ConnectionStatus::Connected(current) = &connection.status else {
|
| 288 |
+
return false;
|
| 289 |
+
};
|
| 290 |
+
if !Arc::ptr_eq(current, expected) || expected.is_disconnected() {
|
| 291 |
+
return false;
|
| 292 |
+
}
|
| 293 |
+
connection.active_process_starts += 1;
|
| 294 |
+
true
|
| 295 |
+
}
|
| 296 |
+
|
| 297 |
+
pub(super) fn finish_process_start(&self) {
|
| 298 |
+
{
|
| 299 |
+
let mut connection = self
|
| 300 |
+
.connection
|
| 301 |
+
.lock()
|
| 302 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 303 |
+
if connection.active_process_starts == 0 {
|
| 304 |
+
tracing::error!("finished an exec-server process start that was not active");
|
| 305 |
+
return;
|
| 306 |
+
}
|
| 307 |
+
connection.active_process_starts -= 1;
|
| 308 |
+
}
|
| 309 |
+
self.notify_connection_changed();
|
| 310 |
+
}
|
| 311 |
+
|
| 312 |
+
pub(super) fn is_failed(&self) -> bool {
|
| 313 |
+
self.failure_message().is_some()
|
| 314 |
+
}
|
| 315 |
+
|
| 316 |
+
pub(super) fn failure_message(&self) -> Option<String> {
|
| 317 |
+
let connection = self
|
| 318 |
+
.connection
|
| 319 |
+
.lock()
|
| 320 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 321 |
+
match &connection.status {
|
| 322 |
+
ConnectionStatus::Failed(message) => Some(message.clone()),
|
| 323 |
+
ConnectionStatus::Connected(_) | ConnectionStatus::Recovering => None,
|
| 324 |
+
}
|
| 325 |
+
}
|
| 326 |
+
|
| 327 |
+
pub(super) fn request_recovery(
|
| 328 |
+
self: &Arc<Self>,
|
| 329 |
+
failed_rpc_client: Arc<RpcClient>,
|
| 330 |
+
disconnect_message: String,
|
| 331 |
+
) {
|
| 332 |
+
let should_recover = {
|
| 333 |
+
let mut connection = self
|
| 334 |
+
.connection
|
| 335 |
+
.lock()
|
| 336 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 337 |
+
match &connection.status {
|
| 338 |
+
ConnectionStatus::Connected(current)
|
| 339 |
+
if Arc::ptr_eq(current, &failed_rpc_client) =>
|
| 340 |
+
{
|
| 341 |
+
connection.set_status(ConnectionStatus::Recovering);
|
| 342 |
+
true
|
| 343 |
+
}
|
| 344 |
+
ConnectionStatus::Connected(_)
|
| 345 |
+
| ConnectionStatus::Recovering
|
| 346 |
+
| ConnectionStatus::Failed(_) => false,
|
| 347 |
+
}
|
| 348 |
+
};
|
| 349 |
+
if !should_recover {
|
| 350 |
+
return;
|
| 351 |
+
}
|
| 352 |
+
|
| 353 |
+
self.notify_connection_changed();
|
| 354 |
+
let inner = Arc::clone(self);
|
| 355 |
+
tokio::spawn(async move {
|
| 356 |
+
tokio::select! {
|
| 357 |
+
biased;
|
| 358 |
+
_ = inner.retired.cancelled() => {},
|
| 359 |
+
_ = inner.recover(disconnect_message) => {},
|
| 360 |
+
}
|
| 361 |
+
});
|
| 362 |
+
}
|
| 363 |
+
|
| 364 |
+
async fn recover(self: &Arc<Self>, disconnect_message: String) {
|
| 365 |
+
let deadline = Instant::now() + SESSION_RECOVERY_TIMEOUT;
|
| 366 |
+
self.fail_all_http_body_streams(disconnect_message.clone())
|
| 367 |
+
.await;
|
| 368 |
+
if timeout_at(deadline, self.wait_for_process_starts())
|
| 369 |
+
.await
|
| 370 |
+
.is_err()
|
| 371 |
+
{
|
| 372 |
+
let message = format!(
|
| 373 |
+
"{disconnect_message}; failed to resume exec-server session: recovery timed out after {SESSION_RECOVERY_TIMEOUT:?}"
|
| 374 |
+
);
|
| 375 |
+
self.fail(message).await;
|
| 376 |
+
return;
|
| 377 |
+
}
|
| 378 |
+
if self.reconnect_strategy.is_none() {
|
| 379 |
+
self.fail(disconnect_message).await;
|
| 380 |
+
return;
|
| 381 |
+
}
|
| 382 |
+
|
| 383 |
+
let Some(session_id) = self.session_id.get().cloned() else {
|
| 384 |
+
let message = format!(
|
| 385 |
+
"{disconnect_message}; failed to resume exec-server session: missing session id"
|
| 386 |
+
);
|
| 387 |
+
self.fail(message).await;
|
| 388 |
+
return;
|
| 389 |
+
};
|
| 390 |
+
let uses_registry_backoff = matches!(
|
| 391 |
+
self.reconnect_strategy.as_ref(),
|
| 392 |
+
Some(ExecServerReconnectStrategy::NoiseRendezvous { .. })
|
| 393 |
+
);
|
| 394 |
+
let mut registry_retry_attempt = 0;
|
| 395 |
+
let last_error = loop {
|
| 396 |
+
match timeout_at(deadline, self.resume_once(&session_id)).await {
|
| 397 |
+
Ok(Ok((rpc_client, _attempt))) => {
|
| 398 |
+
if !rpc_client.is_disconnected() && self.install_recovered_client(rpc_client) {
|
| 399 |
+
return;
|
| 400 |
+
}
|
| 401 |
+
}
|
| 402 |
+
Ok(Err(error)) if !is_retryable_recovery_error(&error) => {
|
| 403 |
+
break error.to_string();
|
| 404 |
+
}
|
| 405 |
+
Ok(Err(_)) => {}
|
| 406 |
+
Err(_) => {
|
| 407 |
+
break format!("recovery timed out after {SESSION_RECOVERY_TIMEOUT:?}");
|
| 408 |
+
}
|
| 409 |
+
}
|
| 410 |
+
|
| 411 |
+
let retry_delay = if uses_registry_backoff {
|
| 412 |
+
let delay = registry_recovery_retry_delay(&session_id, registry_retry_attempt);
|
| 413 |
+
registry_retry_attempt = registry_retry_attempt.saturating_add(1);
|
| 414 |
+
delay
|
| 415 |
+
} else {
|
| 416 |
+
SESSION_RECOVERY_RETRY_INTERVAL
|
| 417 |
+
};
|
| 418 |
+
|
| 419 |
+
let now = Instant::now();
|
| 420 |
+
if now >= deadline {
|
| 421 |
+
break format!("recovery timed out after {SESSION_RECOVERY_TIMEOUT:?}");
|
| 422 |
+
}
|
| 423 |
+
sleep(retry_delay.min(deadline - now)).await;
|
| 424 |
+
};
|
| 425 |
+
|
| 426 |
+
let message =
|
| 427 |
+
format!("{disconnect_message}; failed to resume exec-server session: {last_error}");
|
| 428 |
+
self.fail(message).await;
|
| 429 |
+
}
|
| 430 |
+
|
| 431 |
+
async fn wait_for_process_starts(&self) {
|
| 432 |
+
let mut connection_changed = self.connection_changed.subscribe();
|
| 433 |
+
loop {
|
| 434 |
+
let starts_are_done = self
|
| 435 |
+
.connection
|
| 436 |
+
.lock()
|
| 437 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner)
|
| 438 |
+
.active_process_starts
|
| 439 |
+
== 0;
|
| 440 |
+
if starts_are_done {
|
| 441 |
+
return;
|
| 442 |
+
}
|
| 443 |
+
let _ = connection_changed.changed().await;
|
| 444 |
+
}
|
| 445 |
+
}
|
| 446 |
+
|
| 447 |
+
fn install_recovered_client(&self, rpc_client: Arc<RpcClient>) -> bool {
|
| 448 |
+
let installed = {
|
| 449 |
+
let mut connection = self
|
| 450 |
+
.connection
|
| 451 |
+
.lock()
|
| 452 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 453 |
+
if !matches!(connection.status, ConnectionStatus::Recovering)
|
| 454 |
+
|| rpc_client.is_disconnected()
|
| 455 |
+
{
|
| 456 |
+
false
|
| 457 |
+
} else {
|
| 458 |
+
connection.set_status(ConnectionStatus::Connected(rpc_client));
|
| 459 |
+
true
|
| 460 |
+
}
|
| 461 |
+
};
|
| 462 |
+
if installed {
|
| 463 |
+
self.notify_connection_changed();
|
| 464 |
+
}
|
| 465 |
+
installed
|
| 466 |
+
}
|
| 467 |
+
|
| 468 |
+
fn notify_connection_changed(&self) {
|
| 469 |
+
self.connection_changed.send_replace(());
|
| 470 |
+
}
|
| 471 |
+
|
| 472 |
+
async fn resume_once(
|
| 473 |
+
self: &Arc<Self>,
|
| 474 |
+
session_id: &str,
|
| 475 |
+
) -> Result<(Arc<RpcClient>, Option<tokio::sync::OwnedSemaphorePermit>), ExecServerError> {
|
| 476 |
+
let reconnect_strategy = self
|
| 477 |
+
.reconnect_strategy
|
| 478 |
+
.as_ref()
|
| 479 |
+
.ok_or_else(|| ExecServerError::Protocol("missing reconnect strategy".to_string()))?;
|
| 480 |
+
let attempt = reconnect_strategy.resume(session_id).await?;
|
| 481 |
+
let (connection, options, attempt_permit, noise_context) = attempt.into_parts();
|
| 482 |
+
let (rpc_client, events_rx) = RpcClient::new(connection);
|
| 483 |
+
let rpc_client = Arc::new(rpc_client);
|
| 484 |
+
let client = ExecServerClient {
|
| 485 |
+
inner: Arc::clone(self),
|
| 486 |
+
recovery_policy: RecoveryPolicy::Wait,
|
| 487 |
+
};
|
| 488 |
+
// Resuming a session redirects notifications from its running processes
|
| 489 |
+
// to this connection during initialize. Drain them immediately so a
|
| 490 |
+
// burst cannot fill the bounded event channel and block the initialize
|
| 491 |
+
// response behind it.
|
| 492 |
+
client.spawn_rpc_reader(&rpc_client, events_rx);
|
| 493 |
+
client
|
| 494 |
+
.initialize_rpc(&rpc_client, options, noise_context)
|
| 495 |
+
.await?;
|
| 496 |
+
|
| 497 |
+
self.recover_processes(&rpc_client).await?;
|
| 498 |
+
Ok((rpc_client, attempt_permit))
|
| 499 |
+
}
|
| 500 |
+
|
| 501 |
+
async fn recover_processes(
|
| 502 |
+
self: &Arc<Self>,
|
| 503 |
+
rpc_client: &RpcClient,
|
| 504 |
+
) -> Result<(), ExecServerError> {
|
| 505 |
+
let sessions = self.sessions.load_full();
|
| 506 |
+
for (process_id, session) in sessions.iter() {
|
| 507 |
+
if !session.recoverable.load(Ordering::Acquire) {
|
| 508 |
+
continue;
|
| 509 |
+
}
|
| 510 |
+
let response = rpc_client
|
| 511 |
+
.call::<_, ReadResponse>(
|
| 512 |
+
EXEC_READ_METHOD,
|
| 513 |
+
&ReadParams {
|
| 514 |
+
process_id: process_id.clone(),
|
| 515 |
+
after_seq: Some(session.last_published_seq()),
|
| 516 |
+
max_bytes: None,
|
| 517 |
+
wait_ms: Some(0),
|
| 518 |
+
},
|
| 519 |
+
)
|
| 520 |
+
.await
|
| 521 |
+
.map_err(ExecServerError::from);
|
| 522 |
+
let recovered = match response {
|
| 523 |
+
Ok(response) => session.recover_events(response),
|
| 524 |
+
Err(error) if is_transport_closed_error(&error) => return Err(error),
|
| 525 |
+
Err(error) => Err(error),
|
| 526 |
+
};
|
| 527 |
+
match recovered {
|
| 528 |
+
Ok(true) => self.remove_session_if(process_id, session),
|
| 529 |
+
Ok(false) => {}
|
| 530 |
+
Err(error) => {
|
| 531 |
+
session
|
| 532 |
+
.network_policy
|
| 533 |
+
.cancellation
|
| 534 |
+
.record(NetworkRequestCancellationReason::ProcessCancelled);
|
| 535 |
+
let terminated: Result<TerminateResponse, ExecServerError> = rpc_client
|
| 536 |
+
.call_for_cleanup(
|
| 537 |
+
EXEC_TERMINATE_METHOD,
|
| 538 |
+
&TerminateParams {
|
| 539 |
+
process_id: process_id.clone(),
|
| 540 |
+
},
|
| 541 |
+
)
|
| 542 |
+
.await
|
| 543 |
+
.map_err(ExecServerError::from);
|
| 544 |
+
if let Err(terminate_error) = terminated
|
| 545 |
+
&& is_transport_closed_error(&terminate_error)
|
| 546 |
+
{
|
| 547 |
+
return Err(terminate_error);
|
| 548 |
+
}
|
| 549 |
+
self.remove_session_if(process_id, session);
|
| 550 |
+
session.set_failure(format!("failed to recover process {process_id}: {error}"));
|
| 551 |
+
}
|
| 552 |
+
}
|
| 553 |
+
}
|
| 554 |
+
Ok(())
|
| 555 |
+
}
|
| 556 |
+
|
| 557 |
+
async fn fail(self: &Arc<Self>, message: String) {
|
| 558 |
+
let (message, newly_failed) = {
|
| 559 |
+
let mut connection = self
|
| 560 |
+
.connection
|
| 561 |
+
.lock()
|
| 562 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 563 |
+
match &connection.status {
|
| 564 |
+
ConnectionStatus::Failed(existing) => (existing.clone(), false),
|
| 565 |
+
ConnectionStatus::Connected(_) | ConnectionStatus::Recovering => {
|
| 566 |
+
connection.set_status(ConnectionStatus::Failed(message.clone()));
|
| 567 |
+
(message, true)
|
| 568 |
+
}
|
| 569 |
+
}
|
| 570 |
+
};
|
| 571 |
+
if newly_failed {
|
| 572 |
+
self.notify_connection_changed();
|
| 573 |
+
fail_all_in_flight_work(self, message.clone()).await;
|
| 574 |
+
}
|
| 575 |
+
}
|
| 576 |
+
}
|
| 577 |
+
|
| 578 |
+
impl ExecServerClient {
|
| 579 |
+
pub(super) fn spawn_rpc_reader(
|
| 580 |
+
&self,
|
| 581 |
+
rpc_client: &Arc<RpcClient>,
|
| 582 |
+
mut events_rx: mpsc::Receiver<RpcClientEvent>,
|
| 583 |
+
) {
|
| 584 |
+
let inner = Arc::downgrade(&self.inner);
|
| 585 |
+
let rpc_inbound_request_slots = Arc::clone(&self.inner.rpc_inbound_request_slots);
|
| 586 |
+
let rpc_client = Arc::downgrade(rpc_client);
|
| 587 |
+
let connection_cancelled = CancellationToken::new();
|
| 588 |
+
let connection_cancel_guard = connection_cancelled.clone().drop_guard();
|
| 589 |
+
tokio::spawn(async move {
|
| 590 |
+
let _connection_cancel_guard = connection_cancel_guard;
|
| 591 |
+
while let Some(event) = events_rx.recv().await {
|
| 592 |
+
let (Some(inner), Some(rpc_client)) = (inner.upgrade(), rpc_client.upgrade())
|
| 593 |
+
else {
|
| 594 |
+
return;
|
| 595 |
+
};
|
| 596 |
+
match event {
|
| 597 |
+
RpcClientEvent::Request {
|
| 598 |
+
request,
|
| 599 |
+
request_span,
|
| 600 |
+
} => {
|
| 601 |
+
let mut request_outcome = ClientRequestOutcome {
|
| 602 |
+
span: request_span,
|
| 603 |
+
result: "disconnected",
|
| 604 |
+
};
|
| 605 |
+
if request.method != NETWORK_POLICY_REQUEST_METHOD {
|
| 606 |
+
let error = method_not_found(format!(
|
| 607 |
+
"exec-server client does not implement `{}` yet",
|
| 608 |
+
request.method
|
| 609 |
+
));
|
| 610 |
+
if rpc_client.respond_error(request.id, error).await.is_err() {
|
| 611 |
+
inner.request_recovery(
|
| 612 |
+
rpc_client,
|
| 613 |
+
disconnected_message(/*reason*/ None),
|
| 614 |
+
);
|
| 615 |
+
return;
|
| 616 |
+
}
|
| 617 |
+
request_outcome.complete("error");
|
| 618 |
+
continue;
|
| 619 |
+
}
|
| 620 |
+
request_outcome
|
| 621 |
+
.span
|
| 622 |
+
.record("otel.name", NETWORK_POLICY_REQUEST_METHOD);
|
| 623 |
+
|
| 624 |
+
let request_guard = match rpc_client
|
| 625 |
+
.admit_inbound_request(&request.id, &rpc_inbound_request_slots)
|
| 626 |
+
{
|
| 627 |
+
Ok(request_guard) => request_guard,
|
| 628 |
+
Err(RpcInboundRequestAdmissionError::InvalidRequestId) => {
|
| 629 |
+
rpc_client.close_transport().await;
|
| 630 |
+
inner.request_recovery(
|
| 631 |
+
rpc_client,
|
| 632 |
+
"exec-server sent an invalid request ID".to_string(),
|
| 633 |
+
);
|
| 634 |
+
return;
|
| 635 |
+
}
|
| 636 |
+
Err(RpcInboundRequestAdmissionError::DuplicateRequestId) => {
|
| 637 |
+
rpc_client.close_transport().await;
|
| 638 |
+
inner.request_recovery(
|
| 639 |
+
rpc_client,
|
| 640 |
+
"exec-server reused an in-flight request ID".to_string(),
|
| 641 |
+
);
|
| 642 |
+
return;
|
| 643 |
+
}
|
| 644 |
+
Err(RpcInboundRequestAdmissionError::AtCapacity) => {
|
| 645 |
+
let response = NetworkPolicyRequestResponse {
|
| 646 |
+
decision: ExecServerNetworkPolicyDecision::Deny {
|
| 647 |
+
reason: NETWORK_POLICY_DENIAL_REASON.to_string(),
|
| 648 |
+
},
|
| 649 |
+
};
|
| 650 |
+
if rpc_client.respond(request.id, &response).await.is_err() {
|
| 651 |
+
inner.request_recovery(
|
| 652 |
+
rpc_client,
|
| 653 |
+
disconnected_message(/*reason*/ None),
|
| 654 |
+
);
|
| 655 |
+
return;
|
| 656 |
+
}
|
| 657 |
+
request_outcome.complete("success");
|
| 658 |
+
continue;
|
| 659 |
+
}
|
| 660 |
+
};
|
| 661 |
+
let request_id = request.id;
|
| 662 |
+
let params: NetworkPolicyRequestParams =
|
| 663 |
+
match serde_json::from_value(request.params.unwrap_or(Value::Null)) {
|
| 664 |
+
Ok(params) => params,
|
| 665 |
+
Err(_) => {
|
| 666 |
+
let error = invalid_params(
|
| 667 |
+
"invalid network policy request params".to_string(),
|
| 668 |
+
);
|
| 669 |
+
if rpc_client.respond_error(request_id, error).await.is_err() {
|
| 670 |
+
inner.request_recovery(
|
| 671 |
+
rpc_client,
|
| 672 |
+
disconnected_message(/*reason*/ None),
|
| 673 |
+
);
|
| 674 |
+
return;
|
| 675 |
+
}
|
| 676 |
+
request_outcome.complete("error");
|
| 677 |
+
continue;
|
| 678 |
+
}
|
| 679 |
+
};
|
| 680 |
+
let process_id = params.process_id;
|
| 681 |
+
let request = params.request;
|
| 682 |
+
let process_id_valid = !process_id.is_empty()
|
| 683 |
+
&& process_id.len() <= MAX_NETWORK_POLICY_PROCESS_ID_BYTES;
|
| 684 |
+
let host_valid = !request.host.is_empty()
|
| 685 |
+
&& request.host.len() <= MAX_NETWORK_POLICY_HOST_BYTES
|
| 686 |
+
&& !request.host.chars().any(char::is_control)
|
| 687 |
+
&& !request.host.chars().any(char::is_whitespace);
|
| 688 |
+
let session = (process_id_valid && host_valid)
|
| 689 |
+
.then(|| inner.get_session(&process_id))
|
| 690 |
+
.flatten();
|
| 691 |
+
let controller = session
|
| 692 |
+
.as_ref()
|
| 693 |
+
.and_then(|session| session.network_policy.controller.load_full());
|
| 694 |
+
let process_cancelled = session
|
| 695 |
+
.as_ref()
|
| 696 |
+
.map(|session| session.network_policy.cancelled.clone());
|
| 697 |
+
let process_cancellation = session
|
| 698 |
+
.as_ref()
|
| 699 |
+
.map(|session| session.network_policy.cancellation.clone());
|
| 700 |
+
let cancellation = NetworkRequestCancellation::default();
|
| 701 |
+
let expected_session = session.as_ref().map(Arc::downgrade);
|
| 702 |
+
let policy_request =
|
| 703 |
+
(process_id_valid && host_valid).then_some(NetworkPolicyRequest {
|
| 704 |
+
protocol: match request.protocol {
|
| 705 |
+
ExecServerNetworkProtocol::Http => NetworkProtocol::Http,
|
| 706 |
+
ExecServerNetworkProtocol::HttpsConnect => {
|
| 707 |
+
NetworkProtocol::HttpsConnect
|
| 708 |
+
}
|
| 709 |
+
ExecServerNetworkProtocol::Socks5Tcp => {
|
| 710 |
+
NetworkProtocol::Socks5Tcp
|
| 711 |
+
}
|
| 712 |
+
ExecServerNetworkProtocol::Socks5Udp => {
|
| 713 |
+
NetworkProtocol::Socks5Udp
|
| 714 |
+
}
|
| 715 |
+
},
|
| 716 |
+
host: request.host,
|
| 717 |
+
port: request.port,
|
| 718 |
+
environment_id: None,
|
| 719 |
+
client_addr: None,
|
| 720 |
+
method: None,
|
| 721 |
+
command: None,
|
| 722 |
+
exec_policy_hint: None,
|
| 723 |
+
execution_id: None,
|
| 724 |
+
disconnect: None,
|
| 725 |
+
cancellation: Some(cancellation.clone()),
|
| 726 |
+
});
|
| 727 |
+
let inner = Arc::downgrade(&inner);
|
| 728 |
+
let rpc_client = Arc::downgrade(&rpc_client);
|
| 729 |
+
let connection_cancelled = connection_cancelled.clone();
|
| 730 |
+
let task_span = request_outcome.span.clone();
|
| 731 |
+
let task = async move {
|
| 732 |
+
let _request_guard = request_guard;
|
| 733 |
+
let decision = match (controller, policy_request, process_cancelled) {
|
| 734 |
+
(Some(controller), Some(request), Some(process_cancelled)) => {
|
| 735 |
+
// Keep the decision future outside select/timeout so its
|
| 736 |
+
// guard sees the cancellation cause before it is dropped.
|
| 737 |
+
let mut decision = controller.decider.decide(request);
|
| 738 |
+
tokio::select! {
|
| 739 |
+
biased;
|
| 740 |
+
_ = connection_cancelled.cancelled() => {
|
| 741 |
+
cancellation.record(NetworkRequestCancellationReason::ConnectionClosed);
|
| 742 |
+
return;
|
| 743 |
+
},
|
| 744 |
+
_ = process_cancelled.cancelled() => {
|
| 745 |
+
cancellation.record(process_cancellation.as_ref()
|
| 746 |
+
.and_then(NetworkRequestCancellation::reason)
|
| 747 |
+
.unwrap_or(NetworkRequestCancellationReason::ProcessCancelled));
|
| 748 |
+
NetworkDecision::deny(NETWORK_POLICY_DENIAL_REASON)
|
| 749 |
+
}
|
| 750 |
+
result = timeout(
|
| 751 |
+
controller.timeout,
|
| 752 |
+
&mut decision,
|
| 753 |
+
) => result.unwrap_or_else(|_| {
|
| 754 |
+
cancellation.record(NetworkRequestCancellationReason::TimedOut);
|
| 755 |
+
NetworkDecision::deny(NETWORK_POLICY_DENIAL_REASON)
|
| 756 |
+
}),
|
| 757 |
+
}
|
| 758 |
+
}
|
| 759 |
+
(None, _, _) | (_, None, _) | (_, _, None) => {
|
| 760 |
+
NetworkDecision::deny(NETWORK_POLICY_DENIAL_REASON)
|
| 761 |
+
}
|
| 762 |
+
};
|
| 763 |
+
if let Some(expected_session) = expected_session {
|
| 764 |
+
let (Some(inner), Some(expected_session)) =
|
| 765 |
+
(inner.upgrade(), expected_session.upgrade())
|
| 766 |
+
else {
|
| 767 |
+
return;
|
| 768 |
+
};
|
| 769 |
+
if !inner
|
| 770 |
+
.get_session(&process_id)
|
| 771 |
+
.is_some_and(|session| Arc::ptr_eq(&session, &expected_session))
|
| 772 |
+
{
|
| 773 |
+
return;
|
| 774 |
+
}
|
| 775 |
+
}
|
| 776 |
+
let Some(rpc_client) = rpc_client.upgrade() else {
|
| 777 |
+
return;
|
| 778 |
+
};
|
| 779 |
+
let decision = match decision {
|
| 780 |
+
NetworkDecision::Allow => ExecServerNetworkPolicyDecision::Allow,
|
| 781 |
+
NetworkDecision::Deny {
|
| 782 |
+
reason, decision, ..
|
| 783 |
+
} if reason.len() <= MAX_NETWORK_POLICY_REASON_BYTES
|
| 784 |
+
&& !reason.chars().any(char::is_control) =>
|
| 785 |
+
{
|
| 786 |
+
match decision {
|
| 787 |
+
NetworkPolicyDecision::Deny => {
|
| 788 |
+
ExecServerNetworkPolicyDecision::Deny { reason }
|
| 789 |
+
}
|
| 790 |
+
NetworkPolicyDecision::Ask => {
|
| 791 |
+
ExecServerNetworkPolicyDecision::Ask { reason }
|
| 792 |
+
}
|
| 793 |
+
}
|
| 794 |
+
}
|
| 795 |
+
NetworkDecision::Deny { .. } => {
|
| 796 |
+
ExecServerNetworkPolicyDecision::Deny {
|
| 797 |
+
reason: NETWORK_POLICY_DENIAL_REASON.to_string(),
|
| 798 |
+
}
|
| 799 |
+
}
|
| 800 |
+
};
|
| 801 |
+
if let Err(error) = rpc_client
|
| 802 |
+
.respond(request_id, &NetworkPolicyRequestResponse { decision })
|
| 803 |
+
.await
|
| 804 |
+
{
|
| 805 |
+
debug!(
|
| 806 |
+
?error,
|
| 807 |
+
"failed to send network policy decision to exec-server"
|
| 808 |
+
);
|
| 809 |
+
} else {
|
| 810 |
+
request_outcome.complete("success");
|
| 811 |
+
}
|
| 812 |
+
};
|
| 813 |
+
tokio::spawn(task.instrument(task_span));
|
| 814 |
+
}
|
| 815 |
+
RpcClientEvent::Notification(notification) => {
|
| 816 |
+
if let Err(error) = handle_server_notification(&inner, notification).await {
|
| 817 |
+
rpc_client.close_transport().await;
|
| 818 |
+
inner.request_recovery(
|
| 819 |
+
rpc_client,
|
| 820 |
+
format!("exec-server notification handling failed: {error}"),
|
| 821 |
+
);
|
| 822 |
+
return;
|
| 823 |
+
}
|
| 824 |
+
}
|
| 825 |
+
RpcClientEvent::Disconnected { reason } => {
|
| 826 |
+
inner.request_recovery(rpc_client, disconnected_message(reason.as_deref()));
|
| 827 |
+
return;
|
| 828 |
+
}
|
| 829 |
+
}
|
| 830 |
+
}
|
| 831 |
+
});
|
| 832 |
+
}
|
| 833 |
+
}
|
| 834 |
+
|
| 835 |
+
pub(crate) fn is_retryable_recovery_error(error: &ExecServerError) -> bool {
|
| 836 |
+
if let ExecServerError::ConnectionAttempt(error) = error {
|
| 837 |
+
return is_retryable_recovery_error(error.as_ref());
|
| 838 |
+
}
|
| 839 |
+
is_transport_closed_error(error)
|
| 840 |
+
|| matches!(
|
| 841 |
+
error,
|
| 842 |
+
ExecServerError::ProvisioningFailed(_)
|
| 843 |
+
| ExecServerError::WebSocketConnectTimeout { .. }
|
| 844 |
+
| ExecServerError::WebSocketConnect { .. }
|
| 845 |
+
| ExecServerError::InitializeTimedOut { .. }
|
| 846 |
+
)
|
| 847 |
+
|| is_retryable_registry_error(error)
|
| 848 |
+
|| matches!(
|
| 849 |
+
error,
|
| 850 |
+
ExecServerError::Server { code, .. }
|
| 851 |
+
if *code == SESSION_ALREADY_ATTACHED_ERROR_CODE
|
| 852 |
+
)
|
| 853 |
+
}
|
| 854 |
+
|
| 855 |
+
pub(crate) fn is_retryable_registry_error(error: &ExecServerError) -> bool {
|
| 856 |
+
matches!(
|
| 857 |
+
error,
|
| 858 |
+
ExecServerError::EnvironmentRegistryRequest(error)
|
| 859 |
+
if error.is_connect()
|
| 860 |
+
|| error.is_timeout()
|
| 861 |
+
|| error.is_body()
|
| 862 |
+
|| matches!(
|
| 863 |
+
error,
|
| 864 |
+
codex_http_client::RouteAwareRequestError::Request(error)
|
| 865 |
+
if error.is_decode()
|
| 866 |
+
)
|
| 867 |
+
) || matches!(
|
| 868 |
+
error,
|
| 869 |
+
ExecServerError::EnvironmentRegistryHttp { status, .. }
|
| 870 |
+
if status.is_server_error()
|
| 871 |
+
|| *status == http::StatusCode::REQUEST_TIMEOUT
|
| 872 |
+
|| *status == http::StatusCode::TOO_MANY_REQUESTS
|
| 873 |
+
) || is_environment_offline_error(error)
|
| 874 |
+
}
|
| 875 |
+
|
| 876 |
+
pub(crate) fn is_environment_offline_error(error: &ExecServerError) -> bool {
|
| 877 |
+
matches!(
|
| 878 |
+
error,
|
| 879 |
+
ExecServerError::EnvironmentRegistryHttp { status, code, .. }
|
| 880 |
+
if *status == http::StatusCode::CONFLICT
|
| 881 |
+
&& code.as_deref() == Some("environment_offline")
|
| 882 |
+
)
|
| 883 |
+
}
|
| 884 |
+
|
| 885 |
+
pub(crate) fn registry_recovery_retry_delay(retry_key: &str, attempt: u32) -> Duration {
|
| 886 |
+
let multiplier = 1_u32.checked_shl(attempt.min(4)).unwrap_or(u32::MAX);
|
| 887 |
+
let base_delay = REGISTRY_RECOVERY_INITIAL_RETRY_INTERVAL
|
| 888 |
+
.saturating_mul(multiplier)
|
| 889 |
+
.min(REGISTRY_RECOVERY_MAX_RETRY_INTERVAL);
|
| 890 |
+
let base_millis = base_delay.as_millis() as u64;
|
| 891 |
+
let mut hasher = DefaultHasher::new();
|
| 892 |
+
retry_key.hash(&mut hasher);
|
| 893 |
+
attempt.hash(&mut hasher);
|
| 894 |
+
|
| 895 |
+
Duration::from_millis(base_millis + hasher.finish() % (base_millis / 2 + 1))
|
| 896 |
+
}
|
| 897 |
+
|
| 898 |
+
#[cfg(test)]
|
| 899 |
+
#[path = "client_recovery_tests.rs"]
|
| 900 |
+
mod tests;
|
codex-rs/exec-server/src/client_recovery_tests.rs
ADDED
|
@@ -0,0 +1,262 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::time::Duration;
|
| 2 |
+
|
| 3 |
+
use pretty_assertions::assert_eq;
|
| 4 |
+
|
| 5 |
+
use super::*;
|
| 6 |
+
use crate::protocol::ExecOutputStream;
|
| 7 |
+
use crate::protocol::ProcessOutputChunk;
|
| 8 |
+
|
| 9 |
+
fn registry_error(status: http::StatusCode, code: Option<&str>) -> ExecServerError {
|
| 10 |
+
ExecServerError::EnvironmentRegistryHttp {
|
| 11 |
+
status,
|
| 12 |
+
code: code.map(str::to_string),
|
| 13 |
+
message: "registry unavailable".to_string(),
|
| 14 |
+
}
|
| 15 |
+
}
|
| 16 |
+
|
| 17 |
+
#[test]
|
| 18 |
+
fn registry_recovery_retry_delay_exponentially_backs_off_and_caps() {
|
| 19 |
+
let cases = [
|
| 20 |
+
(0, Duration::from_millis(500)),
|
| 21 |
+
(1, Duration::from_secs(1)),
|
| 22 |
+
(2, Duration::from_secs(2)),
|
| 23 |
+
(3, Duration::from_secs(4)),
|
| 24 |
+
(4, Duration::from_secs(5)),
|
| 25 |
+
(20, Duration::from_secs(5)),
|
| 26 |
+
];
|
| 27 |
+
|
| 28 |
+
for (attempt, base) in cases {
|
| 29 |
+
let delay = registry_recovery_retry_delay("session-1", attempt);
|
| 30 |
+
assert!(delay >= base, "delay {delay:?} for attempt {attempt}");
|
| 31 |
+
assert!(
|
| 32 |
+
delay <= base + base / 2,
|
| 33 |
+
"delay {delay:?} for attempt {attempt}"
|
| 34 |
+
);
|
| 35 |
+
}
|
| 36 |
+
}
|
| 37 |
+
|
| 38 |
+
#[test]
|
| 39 |
+
fn recovery_retries_transient_registry_errors() {
|
| 40 |
+
for status in [
|
| 41 |
+
http::StatusCode::REQUEST_TIMEOUT,
|
| 42 |
+
http::StatusCode::TOO_MANY_REQUESTS,
|
| 43 |
+
http::StatusCode::INTERNAL_SERVER_ERROR,
|
| 44 |
+
http::StatusCode::BAD_GATEWAY,
|
| 45 |
+
http::StatusCode::SERVICE_UNAVAILABLE,
|
| 46 |
+
] {
|
| 47 |
+
let error = registry_error(status, /*code*/ None);
|
| 48 |
+
|
| 49 |
+
assert!(is_retryable_registry_error(&error));
|
| 50 |
+
assert!(is_retryable_recovery_error(&error));
|
| 51 |
+
assert!(is_retryable_recovery_error(
|
| 52 |
+
&ExecServerError::ConnectionAttempt(Arc::new(error))
|
| 53 |
+
));
|
| 54 |
+
}
|
| 55 |
+
}
|
| 56 |
+
|
| 57 |
+
#[test]
|
| 58 |
+
fn recovery_retries_registry_request_timeouts() {
|
| 59 |
+
let error = ExecServerError::EnvironmentRegistryRequest(
|
| 60 |
+
codex_http_client::RouteAwareRequestError::Timeout,
|
| 61 |
+
);
|
| 62 |
+
|
| 63 |
+
assert!(is_retryable_registry_error(&error));
|
| 64 |
+
assert!(is_retryable_recovery_error(&error));
|
| 65 |
+
}
|
| 66 |
+
|
| 67 |
+
#[test]
|
| 68 |
+
fn recovery_retries_environment_offline_conflicts() {
|
| 69 |
+
let error = registry_error(http::StatusCode::CONFLICT, Some("environment_offline"));
|
| 70 |
+
|
| 71 |
+
assert!(is_retryable_registry_error(&error));
|
| 72 |
+
assert!(is_retryable_recovery_error(&error));
|
| 73 |
+
}
|
| 74 |
+
|
| 75 |
+
#[test]
|
| 76 |
+
fn recovery_does_not_retry_other_registry_conflicts() {
|
| 77 |
+
let error = registry_error(http::StatusCode::CONFLICT, Some("registration_conflict"));
|
| 78 |
+
|
| 79 |
+
assert!(!is_retryable_registry_error(&error));
|
| 80 |
+
assert!(!is_retryable_recovery_error(&error));
|
| 81 |
+
assert!(!is_retryable_recovery_error(
|
| 82 |
+
&ExecServerError::ConnectionAttempt(Arc::new(error))
|
| 83 |
+
));
|
| 84 |
+
}
|
| 85 |
+
|
| 86 |
+
#[test]
|
| 87 |
+
fn process_event_reorder_rejects_oversized_output() {
|
| 88 |
+
let state = SessionState::new(/*recoverable*/ true);
|
| 89 |
+
|
| 90 |
+
let error = state
|
| 91 |
+
.publish_ordered_event(ExecProcessEvent::Output(ProcessOutputChunk {
|
| 92 |
+
seq: 1,
|
| 93 |
+
stream: ExecOutputStream::Stdout,
|
| 94 |
+
chunk: vec![0; super::super::MAX_PENDING_PROCESS_EVENT_BYTES + 1].into(),
|
| 95 |
+
}))
|
| 96 |
+
.expect_err("oversized pending process output should be rejected");
|
| 97 |
+
|
| 98 |
+
assert!(error.contains("bytes"));
|
| 99 |
+
}
|
| 100 |
+
|
| 101 |
+
#[test]
|
| 102 |
+
fn process_event_reorder_accepts_gap_closing_event_at_limits() {
|
| 103 |
+
let state = SessionState::new(/*recoverable*/ true);
|
| 104 |
+
let chunk_size =
|
| 105 |
+
super::super::MAX_PENDING_PROCESS_EVENT_BYTES / super::super::MAX_PENDING_PROCESS_EVENTS;
|
| 106 |
+
let last_seq = super::super::MAX_PENDING_PROCESS_EVENTS as u64 + 1;
|
| 107 |
+
|
| 108 |
+
for seq in 2..=last_seq {
|
| 109 |
+
assert!(
|
| 110 |
+
!state
|
| 111 |
+
.publish_ordered_event(ExecProcessEvent::Output(ProcessOutputChunk {
|
| 112 |
+
seq,
|
| 113 |
+
stream: ExecOutputStream::Stdout,
|
| 114 |
+
chunk: vec![0; chunk_size].into(),
|
| 115 |
+
}))
|
| 116 |
+
.expect("future output should fit within reorder limits")
|
| 117 |
+
);
|
| 118 |
+
}
|
| 119 |
+
assert!(
|
| 120 |
+
!state
|
| 121 |
+
.publish_ordered_event(ExecProcessEvent::Output(ProcessOutputChunk {
|
| 122 |
+
seq: 1,
|
| 123 |
+
stream: ExecOutputStream::Stdout,
|
| 124 |
+
chunk: b"x".to_vec().into(),
|
| 125 |
+
}))
|
| 126 |
+
.expect("gap-closing output should drain the reorder buffer")
|
| 127 |
+
);
|
| 128 |
+
|
| 129 |
+
let ordered_events = state
|
| 130 |
+
.ordered_events
|
| 131 |
+
.lock()
|
| 132 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 133 |
+
assert_eq!(
|
| 134 |
+
(
|
| 135 |
+
ordered_events.last_published_seq,
|
| 136 |
+
ordered_events.pending.len(),
|
| 137 |
+
ordered_events.pending_bytes,
|
| 138 |
+
),
|
| 139 |
+
(last_seq, 0, 0)
|
| 140 |
+
);
|
| 141 |
+
}
|
| 142 |
+
|
| 143 |
+
#[test]
|
| 144 |
+
fn recovery_handles_dense_tail_output_and_newer_notification() {
|
| 145 |
+
let state = SessionState::new(/*recoverable*/ true);
|
| 146 |
+
let last_seq = super::super::MAX_PENDING_PROCESS_EVENTS as u64 + 2;
|
| 147 |
+
let live_seq = last_seq + 1;
|
| 148 |
+
assert!(
|
| 149 |
+
!state
|
| 150 |
+
.publish_ordered_event(ExecProcessEvent::Output(ProcessOutputChunk {
|
| 151 |
+
seq: live_seq,
|
| 152 |
+
stream: ExecOutputStream::Stdout,
|
| 153 |
+
chunk: b"live".to_vec().into(),
|
| 154 |
+
}))
|
| 155 |
+
.expect("live output should remain bounded while recovery fills the gap")
|
| 156 |
+
);
|
| 157 |
+
let chunks = (2..=last_seq)
|
| 158 |
+
.map(|seq| ProcessOutputChunk {
|
| 159 |
+
seq,
|
| 160 |
+
stream: ExecOutputStream::Stdout,
|
| 161 |
+
chunk: b"x".to_vec().into(),
|
| 162 |
+
})
|
| 163 |
+
.collect();
|
| 164 |
+
|
| 165 |
+
assert!(
|
| 166 |
+
!state
|
| 167 |
+
.recover_events(ReadResponse {
|
| 168 |
+
chunks,
|
| 169 |
+
next_seq: last_seq + 1,
|
| 170 |
+
exited: true,
|
| 171 |
+
exit_code: Some(17),
|
| 172 |
+
closed: false,
|
| 173 |
+
failure: None,
|
| 174 |
+
sandbox_denied: false,
|
| 175 |
+
})
|
| 176 |
+
.expect("dense retained output should recover")
|
| 177 |
+
);
|
| 178 |
+
|
| 179 |
+
let ordered_events = state
|
| 180 |
+
.ordered_events
|
| 181 |
+
.lock()
|
| 182 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 183 |
+
assert_eq!(
|
| 184 |
+
(
|
| 185 |
+
ordered_events.last_published_seq,
|
| 186 |
+
ordered_events.pending.len(),
|
| 187 |
+
ordered_events.pending_bytes,
|
| 188 |
+
),
|
| 189 |
+
(live_seq, 0, 0)
|
| 190 |
+
);
|
| 191 |
+
}
|
| 192 |
+
|
| 193 |
+
#[test]
|
| 194 |
+
fn recovery_rejects_output_at_closed_sequence() {
|
| 195 |
+
let state = SessionState::new(/*recoverable*/ true);
|
| 196 |
+
|
| 197 |
+
let error = state
|
| 198 |
+
.recover_events(ReadResponse {
|
| 199 |
+
chunks: vec![ProcessOutputChunk {
|
| 200 |
+
seq: 1,
|
| 201 |
+
stream: ExecOutputStream::Stdout,
|
| 202 |
+
chunk: b"output".to_vec().into(),
|
| 203 |
+
}],
|
| 204 |
+
next_seq: 2,
|
| 205 |
+
exited: false,
|
| 206 |
+
exit_code: None,
|
| 207 |
+
closed: true,
|
| 208 |
+
failure: None,
|
| 209 |
+
sandbox_denied: false,
|
| 210 |
+
})
|
| 211 |
+
.expect_err("output should not occupy the closed sequence");
|
| 212 |
+
|
| 213 |
+
assert!(
|
| 214 |
+
error
|
| 215 |
+
.to_string()
|
| 216 |
+
.contains("conflicts with recovered output")
|
| 217 |
+
);
|
| 218 |
+
}
|
| 219 |
+
|
| 220 |
+
#[tokio::test]
|
| 221 |
+
async fn recovery_adds_sandbox_denial_to_pending_exit_event() {
|
| 222 |
+
let state = SessionState::new(/*recoverable*/ true);
|
| 223 |
+
assert!(
|
| 224 |
+
!state
|
| 225 |
+
.publish_ordered_event(ExecProcessEvent::Exited {
|
| 226 |
+
seq: 2,
|
| 227 |
+
exit_code: 1,
|
| 228 |
+
sandbox_denied: None,
|
| 229 |
+
})
|
| 230 |
+
.expect("pending exit should fit within reorder limits")
|
| 231 |
+
);
|
| 232 |
+
|
| 233 |
+
state
|
| 234 |
+
.recover_events(ReadResponse {
|
| 235 |
+
chunks: vec![ProcessOutputChunk {
|
| 236 |
+
seq: 1,
|
| 237 |
+
stream: ExecOutputStream::Stderr,
|
| 238 |
+
chunk: b"sandbox denied".to_vec().into(),
|
| 239 |
+
}],
|
| 240 |
+
next_seq: 3,
|
| 241 |
+
exited: true,
|
| 242 |
+
exit_code: Some(1),
|
| 243 |
+
closed: false,
|
| 244 |
+
failure: None,
|
| 245 |
+
sandbox_denied: true,
|
| 246 |
+
})
|
| 247 |
+
.expect("recovery should publish the pending exit");
|
| 248 |
+
|
| 249 |
+
let mut events = state.subscribe_events();
|
| 250 |
+
assert!(matches!(
|
| 251 |
+
events.recv().await,
|
| 252 |
+
Ok(ExecProcessEvent::Output(_))
|
| 253 |
+
));
|
| 254 |
+
assert_eq!(
|
| 255 |
+
events.recv().await,
|
| 256 |
+
Ok(ExecProcessEvent::Exited {
|
| 257 |
+
seq: 2,
|
| 258 |
+
exit_code: 1,
|
| 259 |
+
sandbox_denied: Some(true),
|
| 260 |
+
})
|
| 261 |
+
);
|
| 262 |
+
}
|
codex-rs/exec-server/src/client_refresh.rs
ADDED
|
@@ -0,0 +1,269 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Explicit connection refresh after a planned executor replacement.
|
| 2 |
+
//!
|
| 3 |
+
//! Ordinary recovery tries to resume the same executor session after a transient
|
| 4 |
+
//! disconnect. Replacement needs a fresh session, without waiting for old recovery
|
| 5 |
+
//! to give up. The caller supplies the ordering: register the replacement first,
|
| 6 |
+
//! then refresh. Executor identity stays inside this connection layer.
|
| 7 |
+
//!
|
| 8 |
+
//! Flow: fresh registry lookup -> reuse or retire session -> connect if needed ->
|
| 9 |
+
//! live status probe. The lazy client and its `Environment` remain the same objects;
|
| 10 |
+
//! only the underlying `ExecServerClient` may change. The public caller contract is on
|
| 11 |
+
//! `Environment::refresh_connection`.
|
| 12 |
+
//!
|
| 13 |
+
//! Two races determine the synchronization here. A client installed during the
|
| 14 |
+
//! lookup makes that lookup stale, so refresh checks again. A connection attempt
|
| 15 |
+
//! cancelled by refresh must never install later. Cancellation and installation
|
| 16 |
+
//! synchronize on the `current_client` lock; acquire `reconnect` first when both are needed.
|
| 17 |
+
//! `refresh_lock` serializes only explicit refreshes; ordinary connection and recovery
|
| 18 |
+
//! work can continue concurrently. Retired sessions cannot publish environment state.
|
| 19 |
+
|
| 20 |
+
use std::sync::Arc;
|
| 21 |
+
use std::sync::Mutex as StdMutex;
|
| 22 |
+
|
| 23 |
+
use futures::future::BoxFuture;
|
| 24 |
+
use tokio::sync::OnceCell;
|
| 25 |
+
use tokio::sync::watch;
|
| 26 |
+
use tokio::time::timeout;
|
| 27 |
+
use tokio_util::sync::CancellationToken;
|
| 28 |
+
|
| 29 |
+
use super::ConnectionResult;
|
| 30 |
+
use super::ConnectionStatus;
|
| 31 |
+
use super::ExecServerClient;
|
| 32 |
+
use super::ExecServerError;
|
| 33 |
+
use super::Inner;
|
| 34 |
+
use super::LazyRemoteExecServerClient;
|
| 35 |
+
use super::fail_all_in_flight_work;
|
| 36 |
+
use crate::EnvironmentConnectionState;
|
| 37 |
+
use crate::NoiseChannelPublicKey;
|
| 38 |
+
use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT;
|
| 39 |
+
use crate::client_api::ExecServerTransportParams;
|
| 40 |
+
use crate::client_api::NoiseRendezvousConnectBundle;
|
| 41 |
+
use crate::client_api::NoiseRendezvousConnectProvider;
|
| 42 |
+
use crate::client_transport::ExecServerReconnectStrategy;
|
| 43 |
+
|
| 44 |
+
/// Shared startup/reconnect result plus cancellation for work superseded by refresh.
|
| 45 |
+
/// The optional transport carries a refresh lookup's bundle into the normal connector.
|
| 46 |
+
#[derive(Default)]
|
| 47 |
+
pub(super) struct ConnectionAttempt {
|
| 48 |
+
pub(super) result: OnceCell<ConnectionResult>,
|
| 49 |
+
pub(super) cancelled: CancellationToken,
|
| 50 |
+
pub(super) transport: Option<ExecServerTransportParams>,
|
| 51 |
+
}
|
| 52 |
+
|
| 53 |
+
// Use the compared bundle intact for the first connection: address, key and authorization
|
| 54 |
+
// belong together. Later lookups, including authorization refresh, use the real provider.
|
| 55 |
+
struct PrefetchedConnectProvider {
|
| 56 |
+
bundle: StdMutex<Option<NoiseRendezvousConnectBundle>>,
|
| 57 |
+
provider: Arc<dyn NoiseRendezvousConnectProvider>,
|
| 58 |
+
}
|
| 59 |
+
|
| 60 |
+
impl NoiseRendezvousConnectProvider for PrefetchedConnectProvider {
|
| 61 |
+
fn connect_bundle(
|
| 62 |
+
&self,
|
| 63 |
+
harness_public_key: NoiseChannelPublicKey,
|
| 64 |
+
) -> BoxFuture<'_, Result<NoiseRendezvousConnectBundle, ExecServerError>> {
|
| 65 |
+
Box::pin(async move {
|
| 66 |
+
let bundle = self
|
| 67 |
+
.bundle
|
| 68 |
+
.lock()
|
| 69 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner)
|
| 70 |
+
.take();
|
| 71 |
+
match bundle {
|
| 72 |
+
Some(bundle) => Ok(bundle),
|
| 73 |
+
None => self.provider.connect_bundle(harness_public_key).await,
|
| 74 |
+
}
|
| 75 |
+
})
|
| 76 |
+
}
|
| 77 |
+
}
|
| 78 |
+
|
| 79 |
+
impl LazyRemoteExecServerClient {
|
| 80 |
+
#[expect(
|
| 81 |
+
clippy::await_holding_invalid_type,
|
| 82 |
+
reason = "serialize explicit refreshes, not ordinary connection or recovery attempts"
|
| 83 |
+
)]
|
| 84 |
+
pub(crate) async fn refresh_connection(&self) -> Result<(), ExecServerError> {
|
| 85 |
+
let _refresh = self.refresh_lock.lock().await;
|
| 86 |
+
let (previous, attempt) = loop {
|
| 87 |
+
let observed = self.cached_client();
|
| 88 |
+
let mut transport = self.transport_params.clone().ok_or_else(|| {
|
| 89 |
+
ExecServerError::Protocol(
|
| 90 |
+
"connection refresh requires a Noise registry".to_string(),
|
| 91 |
+
)
|
| 92 |
+
})?;
|
| 93 |
+
let target = match &mut transport {
|
| 94 |
+
ExecServerTransportParams::Deferred(deferred) => &mut deferred.transport,
|
| 95 |
+
transport => transport,
|
| 96 |
+
};
|
| 97 |
+
let ExecServerTransportParams::NoiseRendezvous { provider, identity } = target else {
|
| 98 |
+
return Err(ExecServerError::Protocol(
|
| 99 |
+
"connection refresh requires a Noise registry".to_string(),
|
| 100 |
+
));
|
| 101 |
+
};
|
| 102 |
+
// This lookup is independent of the old session and its recovery deadline.
|
| 103 |
+
let bundle = timeout(
|
| 104 |
+
DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT,
|
| 105 |
+
provider.connect_bundle(identity.public_key()),
|
| 106 |
+
)
|
| 107 |
+
.await
|
| 108 |
+
.map_err(|_| {
|
| 109 |
+
ExecServerError::EnvironmentRegistryRequest(
|
| 110 |
+
codex_http_client::RouteAwareRequestError::Timeout,
|
| 111 |
+
)
|
| 112 |
+
})??;
|
| 113 |
+
let executor_public_key = bundle.executor_public_key.clone();
|
| 114 |
+
*provider = Arc::new(PrefetchedConnectProvider {
|
| 115 |
+
bundle: StdMutex::new(Some(bundle)),
|
| 116 |
+
provider: Arc::clone(provider),
|
| 117 |
+
});
|
| 118 |
+
|
| 119 |
+
let mut reconnect = self
|
| 120 |
+
.reconnect
|
| 121 |
+
.lock()
|
| 122 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 123 |
+
let current = self
|
| 124 |
+
.current_client
|
| 125 |
+
.lock()
|
| 126 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 127 |
+
// An ordinary connection may have finished during the lookup. Re-read the
|
| 128 |
+
// registry rather than retire a newer client using a superseded response.
|
| 129 |
+
if !match (&observed, &*current) {
|
| 130 |
+
(Some(observed), Some(current)) => Arc::ptr_eq(&observed.inner, ¤t.inner),
|
| 131 |
+
(None, None) => true,
|
| 132 |
+
_ => false,
|
| 133 |
+
} {
|
| 134 |
+
continue;
|
| 135 |
+
}
|
| 136 |
+
// is_disconnected means terminally failed, not temporarily recovering.
|
| 137 |
+
// Preserve same-executor recovery; the final probe fails fast if still recovering.
|
| 138 |
+
if current.as_ref().is_some_and(|client| {
|
| 139 |
+
!client.is_disconnected()
|
| 140 |
+
&& matches!(
|
| 141 |
+
client.inner.reconnect_strategy.as_ref(),
|
| 142 |
+
Some(ExecServerReconnectStrategy::NoiseRendezvous {
|
| 143 |
+
executor_public_key: key, ..
|
| 144 |
+
}) if key == &executor_public_key
|
| 145 |
+
)
|
| 146 |
+
}) {
|
| 147 |
+
break (current.clone(), None);
|
| 148 |
+
}
|
| 149 |
+
// Cancellation and connection installation use the same lock. A late
|
| 150 |
+
// handshake cannot install a client after its attempt has been superseded.
|
| 151 |
+
self.startup.cancelled.cancel();
|
| 152 |
+
if let Some(attempt) = reconnect.as_ref() {
|
| 153 |
+
attempt.cancelled.cancel();
|
| 154 |
+
}
|
| 155 |
+
self.environment_connection_state_tx
|
| 156 |
+
.send_replace(EnvironmentConnectionState::Disconnected);
|
| 157 |
+
let attempt = Arc::new(ConnectionAttempt {
|
| 158 |
+
transport: Some(transport),
|
| 159 |
+
..Default::default()
|
| 160 |
+
});
|
| 161 |
+
*reconnect = Some(Arc::clone(&attempt));
|
| 162 |
+
break (current.clone(), Some(attempt));
|
| 163 |
+
};
|
| 164 |
+
let client = match attempt {
|
| 165 |
+
Some(attempt) => {
|
| 166 |
+
if let Some(previous) = previous {
|
| 167 |
+
previous.inner.retire().await;
|
| 168 |
+
}
|
| 169 |
+
let result = attempt
|
| 170 |
+
.result
|
| 171 |
+
.get_or_init(|| self.connect_once(&attempt))
|
| 172 |
+
.await
|
| 173 |
+
.clone();
|
| 174 |
+
let mut reconnect = self
|
| 175 |
+
.reconnect
|
| 176 |
+
.lock()
|
| 177 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 178 |
+
if reconnect
|
| 179 |
+
.as_ref()
|
| 180 |
+
.is_some_and(|current| Arc::ptr_eq(current, &attempt))
|
| 181 |
+
{
|
| 182 |
+
*reconnect = None;
|
| 183 |
+
}
|
| 184 |
+
result.map_err(ExecServerError::ConnectionAttempt)?
|
| 185 |
+
}
|
| 186 |
+
None => previous.ok_or_else(|| {
|
| 187 |
+
ExecServerError::Protocol("current executor session is missing".to_string())
|
| 188 |
+
})?,
|
| 189 |
+
};
|
| 190 |
+
// Metadata may be cached; readiness requires a live, non-recovering probe.
|
| 191 |
+
client.environment_status().await.map(drop)
|
| 192 |
+
}
|
| 193 |
+
|
| 194 |
+
#[tracing::instrument(name = "codex.exec_server.remote.connect", skip_all)]
|
| 195 |
+
pub(super) fn connect_once<'a>(
|
| 196 |
+
&'a self,
|
| 197 |
+
attempt: &'a ConnectionAttempt,
|
| 198 |
+
) -> BoxFuture<'a, ConnectionResult> {
|
| 199 |
+
// Keep the transport future out of every caller's async layout, including
|
| 200 |
+
// the CLI entry point, which otherwise exceeds rustc's query-depth limit.
|
| 201 |
+
Box::pin(async move {
|
| 202 |
+
let transport = attempt
|
| 203 |
+
.transport
|
| 204 |
+
.as_ref()
|
| 205 |
+
.or(self.transport_params.as_ref())
|
| 206 |
+
.ok_or_else(|| {
|
| 207 |
+
Arc::new(ExecServerError::Protocol(
|
| 208 |
+
"missing transport params for lazy exec-server connection".to_string(),
|
| 209 |
+
))
|
| 210 |
+
})?;
|
| 211 |
+
let client = tokio::select! {
|
| 212 |
+
biased;
|
| 213 |
+
_ = attempt.cancelled.cancelled() => return Err(Arc::new(ExecServerError::Disconnected("connection attempt was superseded".to_string()))),
|
| 214 |
+
result = ExecServerClient::connect_for_transport(transport.clone(), self.http_client_factory.clone()) => result.map_err(Arc::new)?,
|
| 215 |
+
};
|
| 216 |
+
// Cancellation can race with a completed handshake. Recheck before attaching
|
| 217 |
+
// state or installing the client, under the same lock used by refresh.
|
| 218 |
+
{
|
| 219 |
+
let mut current = self
|
| 220 |
+
.current_client
|
| 221 |
+
.lock()
|
| 222 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 223 |
+
if !attempt.cancelled.is_cancelled() {
|
| 224 |
+
client.attach_environment_connection_state(
|
| 225 |
+
self.environment_connection_state_tx.clone(),
|
| 226 |
+
);
|
| 227 |
+
*current = Some(client.clone());
|
| 228 |
+
return Ok(client);
|
| 229 |
+
}
|
| 230 |
+
}
|
| 231 |
+
client.inner.retire().await;
|
| 232 |
+
Err(Arc::new(ExecServerError::Disconnected(
|
| 233 |
+
"connection attempt was superseded".to_string(),
|
| 234 |
+
)))
|
| 235 |
+
})
|
| 236 |
+
}
|
| 237 |
+
}
|
| 238 |
+
|
| 239 |
+
impl Inner {
|
| 240 |
+
async fn retire(self: &Arc<Self>) {
|
| 241 |
+
let message = "exec-server executor was replaced".to_string();
|
| 242 |
+
let rpc_client = {
|
| 243 |
+
let mut connection = self
|
| 244 |
+
.connection
|
| 245 |
+
.lock()
|
| 246 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 247 |
+
// Detach before a later transport completion can publish stale state.
|
| 248 |
+
connection.environment_connection_state_tx =
|
| 249 |
+
watch::channel(EnvironmentConnectionState::Disconnected).0;
|
| 250 |
+
let rpc_client = match &connection.status {
|
| 251 |
+
ConnectionStatus::Connected(client) => Some(Arc::clone(client)),
|
| 252 |
+
ConnectionStatus::Recovering | ConnectionStatus::Failed(_) => None,
|
| 253 |
+
};
|
| 254 |
+
self.retired.cancel();
|
| 255 |
+
connection.set_status(ConnectionStatus::Failed(message.clone()));
|
| 256 |
+
rpc_client
|
| 257 |
+
};
|
| 258 |
+
self.connection_changed.send_replace(());
|
| 259 |
+
// Drain pending RPCs before stream cleanup, which may wait for other work.
|
| 260 |
+
if let Some(rpc_client) = rpc_client {
|
| 261 |
+
rpc_client.close_transport().await;
|
| 262 |
+
}
|
| 263 |
+
fail_all_in_flight_work(self, message).await;
|
| 264 |
+
}
|
| 265 |
+
}
|
| 266 |
+
|
| 267 |
+
#[cfg(test)]
|
| 268 |
+
#[path = "client_refresh_tests.rs"]
|
| 269 |
+
mod tests;
|
codex-rs/exec-server/src/client_refresh_tests.rs
ADDED
|
@@ -0,0 +1,621 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::sync::Arc;
|
| 2 |
+
use std::sync::Mutex;
|
| 3 |
+
use std::time::Duration;
|
| 4 |
+
|
| 5 |
+
use anyhow::Result;
|
| 6 |
+
use futures::future::BoxFuture;
|
| 7 |
+
use pretty_assertions::assert_eq;
|
| 8 |
+
use tokio::net::TcpListener;
|
| 9 |
+
use tokio::sync::Notify;
|
| 10 |
+
use tokio::sync::oneshot;
|
| 11 |
+
use tokio::task::JoinSet;
|
| 12 |
+
use tokio_util::task::AbortOnDropHandle;
|
| 13 |
+
|
| 14 |
+
use super::*;
|
| 15 |
+
use crate::ExecServerRuntimePaths;
|
| 16 |
+
use crate::NoiseChannelIdentity;
|
| 17 |
+
use crate::ProcessId;
|
| 18 |
+
use crate::relay::HarnessKeyValidator;
|
| 19 |
+
use crate::relay::run_multiplexed_environment;
|
| 20 |
+
use crate::server::ConnectionProcessor;
|
| 21 |
+
use codex_http_client::HttpClientFactory;
|
| 22 |
+
use codex_http_client::OutboundProxyPolicy;
|
| 23 |
+
|
| 24 |
+
// These tests exercise real sockets, not keepalive deadlines. The shared unit-test
|
| 25 |
+
// Pong timeout is only 100 ms and can expire during Noise handshakes under load.
|
| 26 |
+
// A blocking task prevents paused time from auto-advancing while socket I/O is
|
| 27 |
+
// pending; dropping the returned sender releases it without polling or sleeping.
|
| 28 |
+
fn freeze_clock() -> std::sync::mpsc::Sender<()> {
|
| 29 |
+
tokio::time::pause();
|
| 30 |
+
let (guard, dropped) = std::sync::mpsc::channel();
|
| 31 |
+
tokio::task::spawn_blocking(move || {
|
| 32 |
+
let _ = dropped.recv();
|
| 33 |
+
});
|
| 34 |
+
guard
|
| 35 |
+
}
|
| 36 |
+
|
| 37 |
+
#[derive(Clone)]
|
| 38 |
+
struct Target {
|
| 39 |
+
url: String,
|
| 40 |
+
identity: NoiseChannelIdentity,
|
| 41 |
+
registration: String,
|
| 42 |
+
}
|
| 43 |
+
|
| 44 |
+
struct Registry {
|
| 45 |
+
target: Mutex<Target>,
|
| 46 |
+
next_lookup: Mutex<Option<oneshot::Receiver<()>>>,
|
| 47 |
+
lookup_started: Notify,
|
| 48 |
+
}
|
| 49 |
+
|
| 50 |
+
impl Registry {
|
| 51 |
+
fn block_next_lookup(&self) -> oneshot::Sender<()> {
|
| 52 |
+
let (tx, rx) = oneshot::channel();
|
| 53 |
+
*self.next_lookup.lock().unwrap() = Some(rx);
|
| 54 |
+
tx
|
| 55 |
+
}
|
| 56 |
+
}
|
| 57 |
+
|
| 58 |
+
impl NoiseRendezvousConnectProvider for Registry {
|
| 59 |
+
fn connect_bundle(
|
| 60 |
+
&self,
|
| 61 |
+
_: NoiseChannelPublicKey,
|
| 62 |
+
) -> BoxFuture<'_, Result<NoiseRendezvousConnectBundle, ExecServerError>> {
|
| 63 |
+
Box::pin(async move {
|
| 64 |
+
let target = self.target.lock().unwrap().clone();
|
| 65 |
+
let block = self.next_lookup.lock().unwrap().take();
|
| 66 |
+
if let Some(block) = block {
|
| 67 |
+
self.lookup_started.notify_one();
|
| 68 |
+
block
|
| 69 |
+
.await
|
| 70 |
+
.map_err(|_| ExecServerError::Protocol("test lookup failed".to_owned()))?;
|
| 71 |
+
}
|
| 72 |
+
Ok(NoiseRendezvousConnectBundle {
|
| 73 |
+
websocket_url: target.url,
|
| 74 |
+
environment_id: "environment".to_owned(),
|
| 75 |
+
executor_registration_id: target.registration,
|
| 76 |
+
executor_public_key: target.identity.public_key(),
|
| 77 |
+
harness_key_authorization: "authorization".to_owned(),
|
| 78 |
+
})
|
| 79 |
+
})
|
| 80 |
+
}
|
| 81 |
+
}
|
| 82 |
+
|
| 83 |
+
#[derive(Clone, Default)]
|
| 84 |
+
struct Validator {
|
| 85 |
+
handshake: Option<Arc<Notify>>,
|
| 86 |
+
started: Arc<Notify>,
|
| 87 |
+
}
|
| 88 |
+
|
| 89 |
+
impl HarnessKeyValidator for Validator {
|
| 90 |
+
async fn validate_harness_key(
|
| 91 |
+
&self,
|
| 92 |
+
_: &NoiseChannelPublicKey,
|
| 93 |
+
_: &str,
|
| 94 |
+
) -> Result<(), ExecServerError> {
|
| 95 |
+
if let Some(handshake) = &self.handshake {
|
| 96 |
+
self.started.notify_one();
|
| 97 |
+
handshake.notified().await;
|
| 98 |
+
}
|
| 99 |
+
Ok(())
|
| 100 |
+
}
|
| 101 |
+
}
|
| 102 |
+
|
| 103 |
+
struct Executor {
|
| 104 |
+
target: Target,
|
| 105 |
+
_server: AbortOnDropHandle<()>,
|
| 106 |
+
}
|
| 107 |
+
|
| 108 |
+
impl Executor {
|
| 109 |
+
async fn start(validator: Validator) -> Result<Self> {
|
| 110 |
+
let listener = TcpListener::bind("127.0.0.1:0").await?;
|
| 111 |
+
let target = Target {
|
| 112 |
+
url: format!("ws://{}", listener.local_addr()?),
|
| 113 |
+
identity: NoiseChannelIdentity::generate()?,
|
| 114 |
+
registration: uuid::Uuid::new_v4().to_string(),
|
| 115 |
+
};
|
| 116 |
+
let executor = target.clone();
|
| 117 |
+
let processor = ConnectionProcessor::new(ExecServerRuntimePaths::new(
|
| 118 |
+
std::env::current_exe()?,
|
| 119 |
+
/*codex_linux_sandbox_exe*/ None,
|
| 120 |
+
)?);
|
| 121 |
+
let server = tokio::spawn(async move {
|
| 122 |
+
let mut connections = JoinSet::new();
|
| 123 |
+
loop {
|
| 124 |
+
tokio::select! {
|
| 125 |
+
socket = listener.accept() => {
|
| 126 |
+
let (socket, _) = socket.unwrap();
|
| 127 |
+
let executor = executor.clone();
|
| 128 |
+
let processor = processor.clone();
|
| 129 |
+
let validator = validator.clone();
|
| 130 |
+
connections.spawn(async move {
|
| 131 |
+
let socket = tokio_tungstenite::accept_async(socket).await.unwrap();
|
| 132 |
+
run_multiplexed_environment(socket, processor, "environment".to_owned(), executor.registration, executor.identity, validator).await;
|
| 133 |
+
});
|
| 134 |
+
}
|
| 135 |
+
_ = connections.join_next(), if !connections.is_empty() => {}
|
| 136 |
+
}
|
| 137 |
+
}
|
| 138 |
+
});
|
| 139 |
+
Ok(Self {
|
| 140 |
+
target,
|
| 141 |
+
_server: AbortOnDropHandle::new(server),
|
| 142 |
+
})
|
| 143 |
+
}
|
| 144 |
+
|
| 145 |
+
fn client(&self) -> Result<(LazyRemoteExecServerClient, Arc<Registry>)> {
|
| 146 |
+
let registry = Arc::new(Registry {
|
| 147 |
+
target: Mutex::new(self.target.clone()),
|
| 148 |
+
next_lookup: Mutex::new(None),
|
| 149 |
+
lookup_started: Notify::new(),
|
| 150 |
+
});
|
| 151 |
+
let client = LazyRemoteExecServerClient::new(
|
| 152 |
+
ExecServerTransportParams::NoiseRendezvous {
|
| 153 |
+
provider: registry.clone(),
|
| 154 |
+
identity: NoiseChannelIdentity::generate()?,
|
| 155 |
+
},
|
| 156 |
+
HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault),
|
| 157 |
+
);
|
| 158 |
+
Ok((client, registry))
|
| 159 |
+
}
|
| 160 |
+
}
|
| 161 |
+
|
| 162 |
+
async fn disconnect(client: &ExecServerClient) {
|
| 163 |
+
let rpc_client = {
|
| 164 |
+
let connection = client.inner.connection.lock().unwrap();
|
| 165 |
+
let ConnectionStatus::Connected(rpc_client) = &connection.status else {
|
| 166 |
+
panic!("expected connected session")
|
| 167 |
+
};
|
| 168 |
+
Arc::clone(rpc_client)
|
| 169 |
+
};
|
| 170 |
+
rpc_client.close_transport().await;
|
| 171 |
+
}
|
| 172 |
+
|
| 173 |
+
#[tokio::test]
|
| 174 |
+
async fn refresh_cancels_old_recovery_and_connects_without_resuming_old_session() -> Result<()> {
|
| 175 |
+
let _clock = freeze_clock();
|
| 176 |
+
let old = Executor::start(Validator::default()).await?;
|
| 177 |
+
let new = Executor::start(Validator::default()).await?;
|
| 178 |
+
let (client, registry) = old.client()?;
|
| 179 |
+
let original = client.get().await?;
|
| 180 |
+
let process = original
|
| 181 |
+
.register_session(&ProcessId::from("old-process"))
|
| 182 |
+
.await?;
|
| 183 |
+
let blocked_recovery = registry.block_next_lookup();
|
| 184 |
+
disconnect(&original).await;
|
| 185 |
+
registry.lookup_started.notified().await;
|
| 186 |
+
*registry.target.lock().unwrap() = new.target.clone();
|
| 187 |
+
let refresh = client.refresh_connection();
|
| 188 |
+
tokio::pin!(refresh);
|
| 189 |
+
let started = tokio::time::Instant::now();
|
| 190 |
+
assert!(futures::poll!(refresh.as_mut()).is_pending());
|
| 191 |
+
assert!(original.inner.retired.is_cancelled());
|
| 192 |
+
assert!(original.is_disconnected());
|
| 193 |
+
assert_eq!(started.elapsed(), Duration::ZERO);
|
| 194 |
+
refresh.await?;
|
| 195 |
+
let replacement = client.get().await?;
|
| 196 |
+
assert_ne!(original.session_id(), replacement.session_id());
|
| 197 |
+
assert_eq!(
|
| 198 |
+
*client.environment_connection_state_tx.borrow(),
|
| 199 |
+
EnvironmentConnectionState::Connected
|
| 200 |
+
);
|
| 201 |
+
assert!(matches!(
|
| 202 |
+
process.write(b"never replay".to_vec()).await,
|
| 203 |
+
Err(ExecServerError::Disconnected(_))
|
| 204 |
+
));
|
| 205 |
+
// Releasing an old registry response cannot reinstall or disconnect the replacement.
|
| 206 |
+
let _ = blocked_recovery.send(());
|
| 207 |
+
tokio::task::yield_now().await;
|
| 208 |
+
assert!(Arc::ptr_eq(&client.get().await?.inner, &replacement.inner));
|
| 209 |
+
replacement.environment_status().await?;
|
| 210 |
+
Ok(())
|
| 211 |
+
}
|
| 212 |
+
|
| 213 |
+
#[tokio::test]
|
| 214 |
+
async fn refresh_retires_a_still_connected_old_executor() -> Result<()> {
|
| 215 |
+
let _clock = freeze_clock();
|
| 216 |
+
let old = Executor::start(Validator::default()).await?;
|
| 217 |
+
let new = Executor::start(Validator::default()).await?;
|
| 218 |
+
let (client, registry) = old.client()?;
|
| 219 |
+
let original = client.get().await?;
|
| 220 |
+
*registry.target.lock().unwrap() = new.target.clone();
|
| 221 |
+
client.refresh_connection().await?;
|
| 222 |
+
assert!(original.is_disconnected());
|
| 223 |
+
assert_ne!(original.session_id(), client.get().await?.session_id());
|
| 224 |
+
Ok(())
|
| 225 |
+
}
|
| 226 |
+
|
| 227 |
+
#[tokio::test]
|
| 228 |
+
async fn refresh_preserves_a_current_session_across_registration_renewal() -> Result<()> {
|
| 229 |
+
let _clock = freeze_clock();
|
| 230 |
+
let executor = Executor::start(Validator::default()).await?;
|
| 231 |
+
let (client, registry) = executor.client()?;
|
| 232 |
+
let original = client.get().await?;
|
| 233 |
+
registry.target.lock().unwrap().registration = "renewed-registration".to_owned();
|
| 234 |
+
let concurrent = client.clone();
|
| 235 |
+
let concurrent = tokio::spawn(async move { concurrent.refresh_connection().await });
|
| 236 |
+
client.refresh_connection().await?;
|
| 237 |
+
concurrent.await??;
|
| 238 |
+
assert!(Arc::ptr_eq(&original.inner, &client.get().await?.inner));
|
| 239 |
+
original.environment_status().await?;
|
| 240 |
+
Ok(())
|
| 241 |
+
}
|
| 242 |
+
|
| 243 |
+
#[tokio::test]
|
| 244 |
+
async fn refresh_cancels_a_stalled_initial_lookup() -> Result<()> {
|
| 245 |
+
let _clock = freeze_clock();
|
| 246 |
+
let old = Executor::start(Validator::default()).await?;
|
| 247 |
+
let new = Executor::start(Validator::default()).await?;
|
| 248 |
+
let (client, registry) = old.client()?;
|
| 249 |
+
let release = registry.block_next_lookup();
|
| 250 |
+
let initial = client.get();
|
| 251 |
+
tokio::pin!(initial);
|
| 252 |
+
assert!(futures::poll!(initial.as_mut()).is_pending());
|
| 253 |
+
*registry.target.lock().unwrap() = new.target.clone();
|
| 254 |
+
client.refresh_connection().await?;
|
| 255 |
+
assert!(initial.await.is_err());
|
| 256 |
+
let _ = release.send(());
|
| 257 |
+
client.get().await?.environment_status().await?;
|
| 258 |
+
Ok(())
|
| 259 |
+
}
|
| 260 |
+
|
| 261 |
+
#[tokio::test]
|
| 262 |
+
async fn refresh_cancels_a_stalled_noise_handshake() -> Result<()> {
|
| 263 |
+
let _clock = freeze_clock();
|
| 264 |
+
let validator = Validator {
|
| 265 |
+
handshake: Some(Arc::new(Notify::new())),
|
| 266 |
+
..Default::default()
|
| 267 |
+
};
|
| 268 |
+
let old = Executor::start(validator.clone()).await?;
|
| 269 |
+
let new = Executor::start(Validator::default()).await?;
|
| 270 |
+
let (client, registry) = old.client()?;
|
| 271 |
+
let connecting = client.clone();
|
| 272 |
+
let initial = tokio::spawn(async move { connecting.get().await });
|
| 273 |
+
validator.started.notified().await;
|
| 274 |
+
*registry.target.lock().unwrap() = new.target.clone();
|
| 275 |
+
client.refresh_connection().await?;
|
| 276 |
+
assert!(initial.await?.is_err());
|
| 277 |
+
validator.handshake.unwrap().notify_one();
|
| 278 |
+
client.get().await?.environment_status().await?;
|
| 279 |
+
assert_eq!(
|
| 280 |
+
*client.environment_connection_state_tx.borrow(),
|
| 281 |
+
EnvironmentConnectionState::Connected
|
| 282 |
+
);
|
| 283 |
+
Ok(())
|
| 284 |
+
}
|
| 285 |
+
|
| 286 |
+
#[tokio::test]
|
| 287 |
+
async fn failed_refresh_lookup_leaves_the_existing_session_usable() -> Result<()> {
|
| 288 |
+
let _clock = freeze_clock();
|
| 289 |
+
let executor = Executor::start(Validator::default()).await?;
|
| 290 |
+
let (client, registry) = executor.client()?;
|
| 291 |
+
let original = client.get().await?;
|
| 292 |
+
drop(registry.block_next_lookup());
|
| 293 |
+
assert!(client.refresh_connection().await.is_err());
|
| 294 |
+
assert!(Arc::ptr_eq(&original.inner, &client.get().await?.inner));
|
| 295 |
+
original.environment_status().await?;
|
| 296 |
+
Ok(())
|
| 297 |
+
}
|
| 298 |
+
|
| 299 |
+
#[tokio::test]
|
| 300 |
+
async fn failed_replacement_connection_keeps_old_handles_retired_and_get_retries() -> Result<()> {
|
| 301 |
+
let _clock = freeze_clock();
|
| 302 |
+
let old = Executor::start(Validator::default()).await?;
|
| 303 |
+
let new = Executor::start(Validator::default()).await?;
|
| 304 |
+
let (client, registry) = old.client()?;
|
| 305 |
+
let original = client.get().await?;
|
| 306 |
+
let process = original
|
| 307 |
+
.register_session(&ProcessId::from("old-process"))
|
| 308 |
+
.await?;
|
| 309 |
+
|
| 310 |
+
// The registry knows the replacement, but its endpoint drops the new connection.
|
| 311 |
+
let unavailable = TcpListener::bind("127.0.0.1:0").await?;
|
| 312 |
+
let target = Target {
|
| 313 |
+
url: format!("ws://{}", unavailable.local_addr()?),
|
| 314 |
+
..new.target.clone()
|
| 315 |
+
};
|
| 316 |
+
let _rejected_connection = AbortOnDropHandle::new(tokio::spawn(async move {
|
| 317 |
+
drop(unavailable.accept().await.unwrap());
|
| 318 |
+
}));
|
| 319 |
+
*registry.target.lock().unwrap() = target;
|
| 320 |
+
assert!(matches!(
|
| 321 |
+
client.refresh_connection().await,
|
| 322 |
+
Err(ExecServerError::ConnectionAttempt(_))
|
| 323 |
+
));
|
| 324 |
+
assert!(original.inner.retired.is_cancelled());
|
| 325 |
+
assert!(matches!(
|
| 326 |
+
original.environment_status().await,
|
| 327 |
+
Err(ExecServerError::Disconnected(_))
|
| 328 |
+
));
|
| 329 |
+
assert!(matches!(
|
| 330 |
+
process.write(b"never replay".to_vec()).await,
|
| 331 |
+
Err(ExecServerError::Disconnected(_))
|
| 332 |
+
));
|
| 333 |
+
assert_eq!(
|
| 334 |
+
*client.environment_connection_state_tx.borrow(),
|
| 335 |
+
EnvironmentConnectionState::Disconnected
|
| 336 |
+
);
|
| 337 |
+
|
| 338 |
+
// A later caller retries against the registry instead of reviving the retired client.
|
| 339 |
+
*registry.target.lock().unwrap() = new.target.clone();
|
| 340 |
+
let replacement = client.get().await?;
|
| 341 |
+
assert_ne!(original.session_id(), replacement.session_id());
|
| 342 |
+
replacement.environment_status().await?;
|
| 343 |
+
assert_eq!(
|
| 344 |
+
*client.environment_connection_state_tx.borrow(),
|
| 345 |
+
EnvironmentConnectionState::Connected
|
| 346 |
+
);
|
| 347 |
+
assert!(matches!(
|
| 348 |
+
process.write(b"still retired".to_vec()).await,
|
| 349 |
+
Err(ExecServerError::Disconnected(_))
|
| 350 |
+
));
|
| 351 |
+
Ok(())
|
| 352 |
+
}
|
| 353 |
+
|
| 354 |
+
#[tokio::test]
|
| 355 |
+
async fn superseded_refresh_lookup_does_not_retire_a_newer_session() -> Result<()> {
|
| 356 |
+
let _clock = freeze_clock();
|
| 357 |
+
let old = Executor::start(Validator::default()).await?;
|
| 358 |
+
let new = Executor::start(Validator::default()).await?;
|
| 359 |
+
let (client, registry) = old.client()?;
|
| 360 |
+
let original = client.get().await?;
|
| 361 |
+
let release = registry.block_next_lookup();
|
| 362 |
+
let refreshing = client.refresh_connection();
|
| 363 |
+
tokio::pin!(refreshing);
|
| 364 |
+
assert!(futures::poll!(refreshing.as_mut()).is_pending());
|
| 365 |
+
*registry.target.lock().unwrap() = new.target.clone();
|
| 366 |
+
original.inner.retire().await;
|
| 367 |
+
let replacement = client.get().await?;
|
| 368 |
+
release.send(()).unwrap();
|
| 369 |
+
refreshing.await?;
|
| 370 |
+
assert!(Arc::ptr_eq(&replacement.inner, &client.get().await?.inner));
|
| 371 |
+
assert!(!replacement.inner.retired.is_cancelled());
|
| 372 |
+
replacement.environment_status().await?;
|
| 373 |
+
Ok(())
|
| 374 |
+
}
|
| 375 |
+
|
| 376 |
+
#[tokio::test]
|
| 377 |
+
async fn environment_refresh_preserves_environment_and_filesystem_handles() -> Result<()> {
|
| 378 |
+
let _clock = freeze_clock();
|
| 379 |
+
let old = Executor::start(Validator::default()).await?;
|
| 380 |
+
let new = Executor::start(Validator::default()).await?;
|
| 381 |
+
let (_, registry) = old.client()?;
|
| 382 |
+
let manager = crate::EnvironmentManager::from_snapshot(
|
| 383 |
+
crate::environment_provider::EnvironmentProviderSnapshot {
|
| 384 |
+
environments: Vec::new(),
|
| 385 |
+
default: crate::environment_provider::EnvironmentDefault::Disabled,
|
| 386 |
+
include_local: false,
|
| 387 |
+
},
|
| 388 |
+
/*local_runtime_paths*/ None,
|
| 389 |
+
HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault),
|
| 390 |
+
)?;
|
| 391 |
+
let environment = manager
|
| 392 |
+
.materialize_pending_noise_environment("environment".to_owned(), registry.clone())?;
|
| 393 |
+
manager.report_environment_provisioning_status(
|
| 394 |
+
"environment".to_owned(),
|
| 395 |
+
Ok(crate::EnvironmentReadyInfo {
|
| 396 |
+
selected_capability_roots: Vec::new(),
|
| 397 |
+
}),
|
| 398 |
+
registry.clone(),
|
| 399 |
+
)?;
|
| 400 |
+
environment.info().await?;
|
| 401 |
+
let filesystem = environment.get_filesystem();
|
| 402 |
+
*registry.target.lock().unwrap() = new.target.clone();
|
| 403 |
+
environment.refresh_connection().await?;
|
| 404 |
+
assert!(Arc::ptr_eq(
|
| 405 |
+
&environment,
|
| 406 |
+
&manager.get_environment("environment").unwrap()
|
| 407 |
+
));
|
| 408 |
+
assert!(Arc::ptr_eq(&filesystem, &environment.get_filesystem()));
|
| 409 |
+
environment.info().await?;
|
| 410 |
+
Ok(())
|
| 411 |
+
}
|
| 412 |
+
|
| 413 |
+
#[tokio::test]
|
| 414 |
+
async fn refresh_does_not_accept_cached_metadata_during_recovery() -> Result<()> {
|
| 415 |
+
let _clock = freeze_clock();
|
| 416 |
+
let executor = Executor::start(Validator::default()).await?;
|
| 417 |
+
let (client, registry) = executor.client()?;
|
| 418 |
+
let original = client.get().await?;
|
| 419 |
+
original.environment_info().await?;
|
| 420 |
+
let blocked_recovery = registry.block_next_lookup();
|
| 421 |
+
disconnect(&original).await;
|
| 422 |
+
registry.lookup_started.notified().await;
|
| 423 |
+
let started = tokio::time::Instant::now();
|
| 424 |
+
assert!(matches!(
|
| 425 |
+
client.refresh_connection().await,
|
| 426 |
+
Err(ExecServerError::Disconnected(_))
|
| 427 |
+
));
|
| 428 |
+
assert_eq!(started.elapsed(), Duration::ZERO);
|
| 429 |
+
assert!(!original.inner.retired.is_cancelled());
|
| 430 |
+
drop(blocked_recovery);
|
| 431 |
+
Ok(())
|
| 432 |
+
}
|
| 433 |
+
|
| 434 |
+
#[tokio::test]
|
| 435 |
+
async fn refresh_before_startup_marks_startup_finished() -> Result<()> {
|
| 436 |
+
let _clock = freeze_clock();
|
| 437 |
+
let executor = Executor::start(Validator::default()).await?;
|
| 438 |
+
let (client, _) = executor.client()?;
|
| 439 |
+
assert!(!client.startup_finished());
|
| 440 |
+
client.refresh_connection().await?;
|
| 441 |
+
assert!(client.startup_finished());
|
| 442 |
+
assert!(matches!(client.readiness_result(), Some(Ok(()))));
|
| 443 |
+
Ok(())
|
| 444 |
+
}
|
| 445 |
+
|
| 446 |
+
struct ControlledRpc {
|
| 447 |
+
client: ExecServerClient,
|
| 448 |
+
requests: tokio::sync::mpsc::Receiver<codex_exec_server_protocol::JSONRPCMessage>,
|
| 449 |
+
responses: tokio::sync::mpsc::Sender<crate::connection::JsonRpcConnectionEvent>,
|
| 450 |
+
}
|
| 451 |
+
|
| 452 |
+
async fn controlled_rpc() -> Result<ControlledRpc> {
|
| 453 |
+
use crate::connection::JsonRpcConnection;
|
| 454 |
+
use crate::connection::JsonRpcConnectionEvent;
|
| 455 |
+
use crate::connection::JsonRpcTransport;
|
| 456 |
+
use codex_exec_server_protocol::JSONRPCMessage;
|
| 457 |
+
use codex_exec_server_protocol::JSONRPCResponse;
|
| 458 |
+
let (outgoing_tx, mut requests) = tokio::sync::mpsc::channel(/*buffer*/ 8);
|
| 459 |
+
let (responses, incoming_rx) = tokio::sync::mpsc::channel(/*buffer*/ 8);
|
| 460 |
+
let connection = JsonRpcConnection {
|
| 461 |
+
outgoing_tx,
|
| 462 |
+
incoming_rx,
|
| 463 |
+
disconnected_rx: tokio::sync::watch::channel(/*init*/ false).1,
|
| 464 |
+
task_handles: Vec::new(),
|
| 465 |
+
transport: JsonRpcTransport::Plain,
|
| 466 |
+
};
|
| 467 |
+
let connecting = ExecServerClient::connect(connection, /*options*/ Default::default());
|
| 468 |
+
tokio::pin!(connecting);
|
| 469 |
+
assert!(futures::poll!(connecting.as_mut()).is_pending());
|
| 470 |
+
let Some(JSONRPCMessage::Request(initialize)) = requests.recv().await else {
|
| 471 |
+
anyhow::bail!("expected initialize request");
|
| 472 |
+
};
|
| 473 |
+
responses
|
| 474 |
+
.send(JsonRpcConnectionEvent::message(JSONRPCMessage::Response(
|
| 475 |
+
JSONRPCResponse {
|
| 476 |
+
id: initialize.id,
|
| 477 |
+
result: serde_json::json!({"sessionId": "controlled-session"}),
|
| 478 |
+
},
|
| 479 |
+
)))
|
| 480 |
+
.await?;
|
| 481 |
+
let client = connecting.await?;
|
| 482 |
+
assert!(matches!(
|
| 483 |
+
requests.recv().await,
|
| 484 |
+
Some(JSONRPCMessage::Notification(_))
|
| 485 |
+
));
|
| 486 |
+
Ok(ControlledRpc {
|
| 487 |
+
client,
|
| 488 |
+
requests,
|
| 489 |
+
responses,
|
| 490 |
+
})
|
| 491 |
+
}
|
| 492 |
+
|
| 493 |
+
#[tokio::test]
|
| 494 |
+
#[expect(
|
| 495 |
+
clippy::await_holding_invalid_type,
|
| 496 |
+
reason = "hold stream cleanup pending to exercise retirement ordering"
|
| 497 |
+
)]
|
| 498 |
+
async fn retirement_rejects_pending_mutation_before_stream_cleanup() -> Result<()> {
|
| 499 |
+
use crate::connection::JsonRpcConnectionEvent;
|
| 500 |
+
use codex_exec_server_protocol::JSONRPCMessage;
|
| 501 |
+
use codex_exec_server_protocol::JSONRPCResponse;
|
| 502 |
+
for response_queued in [false, true] {
|
| 503 |
+
let mut rpc = controlled_rpc().await?;
|
| 504 |
+
let call = rpc.client.fs_remove(crate::protocol::FsRemoveParams {
|
| 505 |
+
path: "file:///retired-file".parse()?,
|
| 506 |
+
recursive: None,
|
| 507 |
+
force: None,
|
| 508 |
+
follow_symlinks: None,
|
| 509 |
+
sandbox: None,
|
| 510 |
+
});
|
| 511 |
+
tokio::pin!(call);
|
| 512 |
+
assert!(futures::poll!(call.as_mut()).is_pending());
|
| 513 |
+
let Some(JSONRPCMessage::Request(request)) = rpc.requests.recv().await else {
|
| 514 |
+
anyhow::bail!("expected filesystem request");
|
| 515 |
+
};
|
| 516 |
+
let response = JSONRPCMessage::Response(JSONRPCResponse {
|
| 517 |
+
id: request.id,
|
| 518 |
+
result: serde_json::json!({}),
|
| 519 |
+
});
|
| 520 |
+
if response_queued {
|
| 521 |
+
rpc.responses
|
| 522 |
+
.send(JsonRpcConnectionEvent::message(response.clone()))
|
| 523 |
+
.await?;
|
| 524 |
+
let transport = rpc.client.rpc_client_without_recovery()?;
|
| 525 |
+
tokio::time::timeout(Duration::from_secs(5), async {
|
| 526 |
+
while transport.pending_request_count().await != 0 {
|
| 527 |
+
tokio::task::yield_now().await;
|
| 528 |
+
}
|
| 529 |
+
})
|
| 530 |
+
.await?;
|
| 531 |
+
}
|
| 532 |
+
// Hold stream cleanup. Cover both a late response and one already queued
|
| 533 |
+
// for the caller when retirement begins; closing the socket alone misses the latter.
|
| 534 |
+
let streams = rpc.client.inner.http_body_streams_write_lock.lock().await;
|
| 535 |
+
let retirement = rpc.client.inner.retire();
|
| 536 |
+
tokio::pin!(retirement);
|
| 537 |
+
assert!(futures::poll!(retirement.as_mut()).is_pending());
|
| 538 |
+
if !response_queued {
|
| 539 |
+
rpc.responses
|
| 540 |
+
.send(JsonRpcConnectionEvent::message(response))
|
| 541 |
+
.await?;
|
| 542 |
+
}
|
| 543 |
+
assert!(matches!(call.await, Err(ExecServerError::Disconnected(_))));
|
| 544 |
+
drop(streams);
|
| 545 |
+
retirement.await;
|
| 546 |
+
}
|
| 547 |
+
Ok(())
|
| 548 |
+
}
|
| 549 |
+
|
| 550 |
+
#[tokio::test]
|
| 551 |
+
#[expect(
|
| 552 |
+
clippy::await_holding_invalid_type,
|
| 553 |
+
reason = "hold stream cleanup pending to exercise retirement ordering"
|
| 554 |
+
)]
|
| 555 |
+
async fn retirement_rejects_pending_process_start_before_stream_cleanup() -> Result<()> {
|
| 556 |
+
use crate::connection::JsonRpcConnectionEvent;
|
| 557 |
+
use codex_exec_server_protocol::JSONRPCMessage;
|
| 558 |
+
use codex_exec_server_protocol::JSONRPCResponse;
|
| 559 |
+
for response_queued in [false, true] {
|
| 560 |
+
let mut rpc = controlled_rpc().await?;
|
| 561 |
+
let process_id = ProcessId::from("retired-process");
|
| 562 |
+
let call = rpc.client.start_process(
|
| 563 |
+
crate::protocol::ExecParams {
|
| 564 |
+
metadata: Default::default(),
|
| 565 |
+
process_id: process_id.clone(),
|
| 566 |
+
argv: vec!["unused".to_owned()],
|
| 567 |
+
cwd: "file:///".parse()?,
|
| 568 |
+
shell_snapshot: None,
|
| 569 |
+
env_policy: None,
|
| 570 |
+
env: Default::default(),
|
| 571 |
+
tty: false,
|
| 572 |
+
pipe_stdin: false,
|
| 573 |
+
arg0: None,
|
| 574 |
+
sandbox: None,
|
| 575 |
+
enforce_managed_network: false,
|
| 576 |
+
managed_network: None,
|
| 577 |
+
network_proxy: None,
|
| 578 |
+
},
|
| 579 |
+
/*network_policy_decider*/ None,
|
| 580 |
+
);
|
| 581 |
+
tokio::pin!(call);
|
| 582 |
+
assert!(futures::poll!(call.as_mut()).is_pending());
|
| 583 |
+
let Some(JSONRPCMessage::Request(request)) = rpc.requests.recv().await else {
|
| 584 |
+
anyhow::bail!("expected process start request");
|
| 585 |
+
};
|
| 586 |
+
let response = JSONRPCMessage::Response(JSONRPCResponse {
|
| 587 |
+
id: request.id,
|
| 588 |
+
result: serde_json::json!({"processId": "retired-process"}),
|
| 589 |
+
});
|
| 590 |
+
if response_queued {
|
| 591 |
+
rpc.responses
|
| 592 |
+
.send(JsonRpcConnectionEvent::message(response.clone()))
|
| 593 |
+
.await?;
|
| 594 |
+
let state = rpc
|
| 595 |
+
.client
|
| 596 |
+
.inner
|
| 597 |
+
.get_session(&process_id)
|
| 598 |
+
.expect("pending process");
|
| 599 |
+
// The start task has a second response channel to the original caller.
|
| 600 |
+
tokio::time::timeout(Duration::from_secs(5), async {
|
| 601 |
+
while !state.recoverable.load(std::sync::atomic::Ordering::Acquire) {
|
| 602 |
+
tokio::task::yield_now().await;
|
| 603 |
+
}
|
| 604 |
+
})
|
| 605 |
+
.await?;
|
| 606 |
+
}
|
| 607 |
+
let streams = rpc.client.inner.http_body_streams_write_lock.lock().await;
|
| 608 |
+
let retirement = rpc.client.inner.retire();
|
| 609 |
+
tokio::pin!(retirement);
|
| 610 |
+
assert!(futures::poll!(retirement.as_mut()).is_pending());
|
| 611 |
+
if !response_queued {
|
| 612 |
+
rpc.responses
|
| 613 |
+
.send(JsonRpcConnectionEvent::message(response))
|
| 614 |
+
.await?;
|
| 615 |
+
}
|
| 616 |
+
assert!(matches!(call.await, Err(ExecServerError::Disconnected(_))));
|
| 617 |
+
drop(streams);
|
| 618 |
+
retirement.await;
|
| 619 |
+
}
|
| 620 |
+
Ok(())
|
| 621 |
+
}
|
codex-rs/exec-server/src/client_telemetry.rs
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Caller-side RPC attempt counts using the client's protocol method names, independent of tracing.
|
| 2 |
+
|
| 3 |
+
use codex_otel::EXEC_SERVER_CLIENT_REQUEST_COUNT_METRIC;
|
| 4 |
+
use codex_otel::MetricsClient;
|
| 5 |
+
|
| 6 |
+
pub(crate) fn record_client_request(metrics: Option<&MetricsClient>, method: &str) {
|
| 7 |
+
let Some(metrics) = metrics else {
|
| 8 |
+
return;
|
| 9 |
+
};
|
| 10 |
+
// Record before local admission so failures and cancelled calls still count
|
| 11 |
+
// as attempts. Notifications and responses never enter these call paths.
|
| 12 |
+
if metrics
|
| 13 |
+
.counter_with_description(
|
| 14 |
+
EXEC_SERVER_CLIENT_REQUEST_COUNT_METRIC,
|
| 15 |
+
"Total number of client-side exec-server RPC attempts, including local failures.",
|
| 16 |
+
/*inc*/ 1,
|
| 17 |
+
&[("method", method)],
|
| 18 |
+
)
|
| 19 |
+
.is_err()
|
| 20 |
+
{
|
| 21 |
+
tracing::warn!("failed to emit exec-server client request counter");
|
| 22 |
+
}
|
| 23 |
+
}
|
codex-rs/exec-server/src/client_transport.rs
ADDED
|
@@ -0,0 +1,814 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::process::Stdio;
|
| 2 |
+
use std::sync::Arc;
|
| 3 |
+
use std::time::Duration;
|
| 4 |
+
|
| 5 |
+
use tokio::io::AsyncBufReadExt;
|
| 6 |
+
use tokio::io::BufReader;
|
| 7 |
+
use tokio::process::Command;
|
| 8 |
+
use tokio::sync::OwnedSemaphorePermit;
|
| 9 |
+
use tokio::time::Instant;
|
| 10 |
+
use tokio::time::sleep;
|
| 11 |
+
use tokio::time::timeout;
|
| 12 |
+
use tokio::time::timeout_at;
|
| 13 |
+
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
|
| 14 |
+
use tracing::debug;
|
| 15 |
+
use tracing::warn;
|
| 16 |
+
|
| 17 |
+
use codex_api::AuthError;
|
| 18 |
+
use codex_api::AuthProvider;
|
| 19 |
+
use codex_http_client::HttpClientFactory;
|
| 20 |
+
use codex_http_client::Request;
|
| 21 |
+
use codex_http_client::RequestCompression;
|
| 22 |
+
use codex_protocol::shell_environment::scrub_non_inheritable_env_vars;
|
| 23 |
+
use codex_utils_rustls_provider::ensure_rustls_crypto_provider;
|
| 24 |
+
use codex_websocket_client::WebSocketConnection;
|
| 25 |
+
use codex_websocket_client::WebSocketConnector;
|
| 26 |
+
use codex_websocket_client::WebSocketTlsMode;
|
| 27 |
+
use http::HeaderMap;
|
| 28 |
+
|
| 29 |
+
use crate::ExecServerClient;
|
| 30 |
+
use crate::ExecServerError;
|
| 31 |
+
use crate::client::NoiseInitializeContext;
|
| 32 |
+
use crate::client::accepted::AcceptedConnectionSource;
|
| 33 |
+
use crate::client::is_retryable_registry_error;
|
| 34 |
+
use crate::client::registry_recovery_retry_delay;
|
| 35 |
+
use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT;
|
| 36 |
+
use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT;
|
| 37 |
+
use crate::client_api::ExecServerClientConnectOptions;
|
| 38 |
+
use crate::client_api::ExecServerTransportParams;
|
| 39 |
+
use crate::client_api::NoiseRendezvousConnectArgs;
|
| 40 |
+
use crate::client_api::NoiseRendezvousConnectBundle;
|
| 41 |
+
use crate::client_api::NoiseRendezvousConnectProvider;
|
| 42 |
+
use crate::client_api::RemoteExecServerConnectArgs;
|
| 43 |
+
use crate::client_api::StdioExecServerCommand;
|
| 44 |
+
use crate::client_api::StdioExecServerConnectArgs;
|
| 45 |
+
use crate::connection::JsonRpcConnection;
|
| 46 |
+
use crate::noise_channel::NoiseChannelIdentity;
|
| 47 |
+
use crate::noise_relay::NoiseHarnessConnectionArgs;
|
| 48 |
+
use crate::noise_relay::noise_harness_connection_from_websocket_with_readiness;
|
| 49 |
+
use crate::noise_relay::noise_relay_websocket_config;
|
| 50 |
+
use crate::relay::harness_connection_from_websocket;
|
| 51 |
+
use crate::trace_context::current_rendezvous_headers;
|
| 52 |
+
|
| 53 |
+
const ENVIRONMENT_CLIENT_NAME: &str = "codex-environment";
|
| 54 |
+
const INITIAL_REGISTRY_MAX_RETRIES: u32 = 4;
|
| 55 |
+
const INITIAL_REGISTRY_REQUEST_TIMEOUT: Duration = Duration::from_secs(6);
|
| 56 |
+
const INITIAL_REGISTRY_OPERATION_TIMEOUT: Duration = Duration::from_secs(14);
|
| 57 |
+
|
| 58 |
+
pub(crate) async fn connect_websocket_request(
|
| 59 |
+
request: http::Request<()>,
|
| 60 |
+
diagnostic_url: String,
|
| 61 |
+
connector: WebSocketConnector,
|
| 62 |
+
connect_timeout: Duration,
|
| 63 |
+
use_loopback_direct: bool,
|
| 64 |
+
) -> Result<WebSocketConnection, ExecServerError> {
|
| 65 |
+
let websocket_config = tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default();
|
| 66 |
+
timeout(connect_timeout, async {
|
| 67 |
+
if use_loopback_direct {
|
| 68 |
+
connector
|
| 69 |
+
.connect_loopback_direct(request, websocket_config)
|
| 70 |
+
.await
|
| 71 |
+
} else {
|
| 72 |
+
connector.connect(request, websocket_config).await
|
| 73 |
+
}
|
| 74 |
+
})
|
| 75 |
+
.await
|
| 76 |
+
.map_err(|_| ExecServerError::WebSocketConnectTimeout {
|
| 77 |
+
url: diagnostic_url.clone(),
|
| 78 |
+
timeout: connect_timeout,
|
| 79 |
+
})?
|
| 80 |
+
.map(|(websocket, _)| websocket)
|
| 81 |
+
.map_err(|source| ExecServerError::WebSocketConnect {
|
| 82 |
+
url: diagnostic_url,
|
| 83 |
+
source,
|
| 84 |
+
})
|
| 85 |
+
}
|
| 86 |
+
|
| 87 |
+
pub(crate) async fn authenticate_websocket_request(
|
| 88 |
+
request: &mut http::Request<()>,
|
| 89 |
+
auth_provider: &dyn AuthProvider,
|
| 90 |
+
) -> Result<(), AuthError> {
|
| 91 |
+
let url = request.uri().to_string();
|
| 92 |
+
let signing_url = if let Some(rest) = url.strip_prefix("wss://") {
|
| 93 |
+
format!("https://{rest}")
|
| 94 |
+
} else if let Some(rest) = url.strip_prefix("ws://") {
|
| 95 |
+
format!("http://{rest}")
|
| 96 |
+
} else {
|
| 97 |
+
url
|
| 98 |
+
};
|
| 99 |
+
let mut auth_request = Request::new(request.method().clone(), signing_url);
|
| 100 |
+
// Intermediaries may rewrite WebSocket and hop-by-hop headers after signing.
|
| 101 |
+
if let Some(host) = request.headers().get(http::header::HOST) {
|
| 102 |
+
auth_request
|
| 103 |
+
.headers
|
| 104 |
+
.insert(http::header::HOST, host.clone());
|
| 105 |
+
}
|
| 106 |
+
let authenticated = auth_provider.apply_auth(auth_request).await?;
|
| 107 |
+
if authenticated.method != *request.method() {
|
| 108 |
+
return Err(AuthError::Build(
|
| 109 |
+
"authentication changed the WebSocket request method".to_string(),
|
| 110 |
+
));
|
| 111 |
+
}
|
| 112 |
+
if authenticated.body.is_some() || authenticated.compression != RequestCompression::None {
|
| 113 |
+
return Err(AuthError::Build(
|
| 114 |
+
"authentication added a body or compression to the WebSocket request".to_string(),
|
| 115 |
+
));
|
| 116 |
+
}
|
| 117 |
+
|
| 118 |
+
let authenticated_websocket_url = websocket_url_from_authenticated_url(&authenticated.url)?;
|
| 119 |
+
let authenticated_uri = authenticated_websocket_url.parse().map_err(|error| {
|
| 120 |
+
AuthError::Build(format!("invalid authenticated WebSocket URL: {error}"))
|
| 121 |
+
})?;
|
| 122 |
+
let original_host = request.headers().get(http::header::HOST).cloned();
|
| 123 |
+
for (name, value) in &authenticated.headers {
|
| 124 |
+
if is_websocket_handshake_header(name) {
|
| 125 |
+
if name == http::header::HOST && original_host.as_ref() == Some(value) {
|
| 126 |
+
continue;
|
| 127 |
+
}
|
| 128 |
+
return Err(AuthError::Build(format!(
|
| 129 |
+
"authentication changed WebSocket handshake header {name}"
|
| 130 |
+
)));
|
| 131 |
+
}
|
| 132 |
+
request.headers_mut().insert(name, value.clone());
|
| 133 |
+
}
|
| 134 |
+
*request.uri_mut() = authenticated_uri;
|
| 135 |
+
Ok(())
|
| 136 |
+
}
|
| 137 |
+
|
| 138 |
+
fn websocket_url_from_authenticated_url(url: &str) -> Result<String, AuthError> {
|
| 139 |
+
let mut url = url::Url::parse(url)
|
| 140 |
+
.map_err(|error| AuthError::Build(format!("invalid authenticated request URL: {error}")))?;
|
| 141 |
+
let websocket_scheme = match url.scheme() {
|
| 142 |
+
"https" => "wss",
|
| 143 |
+
"http" => "ws",
|
| 144 |
+
scheme => {
|
| 145 |
+
return Err(AuthError::Build(format!(
|
| 146 |
+
"authentication returned unsupported WebSocket URL scheme: {scheme}"
|
| 147 |
+
)));
|
| 148 |
+
}
|
| 149 |
+
};
|
| 150 |
+
url.set_scheme(websocket_scheme).map_err(|_| {
|
| 151 |
+
AuthError::Build("failed to convert authenticated URL to WebSocket scheme".to_string())
|
| 152 |
+
})?;
|
| 153 |
+
Ok(url.into())
|
| 154 |
+
}
|
| 155 |
+
|
| 156 |
+
fn is_websocket_handshake_header(name: &http::header::HeaderName) -> bool {
|
| 157 |
+
name == http::header::HOST
|
| 158 |
+
|| name == http::header::CONNECTION
|
| 159 |
+
|| name == http::header::UPGRADE
|
| 160 |
+
|| name == http::header::CONTENT_LENGTH
|
| 161 |
+
|| name == http::header::TRANSFER_ENCODING
|
| 162 |
+
|| name.as_str().starts_with("sec-websocket-")
|
| 163 |
+
}
|
| 164 |
+
|
| 165 |
+
/// Everything the recovery loop needs for one connection attempt.
|
| 166 |
+
///
|
| 167 |
+
/// An attempt may also carry a permit whose lifetime must extend until the
|
| 168 |
+
/// attempt finishes.
|
| 169 |
+
pub(crate) struct ReconnectAttempt {
|
| 170 |
+
connection: JsonRpcConnection,
|
| 171 |
+
options: ExecServerClientConnectOptions,
|
| 172 |
+
attempt_permit: Option<OwnedSemaphorePermit>,
|
| 173 |
+
noise_context: Option<NoiseInitializeContext>,
|
| 174 |
+
}
|
| 175 |
+
|
| 176 |
+
struct OpenNoiseRendezvousConnection {
|
| 177 |
+
connection: JsonRpcConnection,
|
| 178 |
+
options: ExecServerClientConnectOptions,
|
| 179 |
+
handshake_ready: tokio::sync::oneshot::Receiver<()>,
|
| 180 |
+
}
|
| 181 |
+
|
| 182 |
+
struct ReadyNoiseRendezvousConnection {
|
| 183 |
+
connection: JsonRpcConnection,
|
| 184 |
+
options: ExecServerClientConnectOptions,
|
| 185 |
+
noise_context: NoiseInitializeContext,
|
| 186 |
+
}
|
| 187 |
+
|
| 188 |
+
impl ReconnectAttempt {
|
| 189 |
+
pub(crate) fn new(
|
| 190 |
+
connection: JsonRpcConnection,
|
| 191 |
+
options: ExecServerClientConnectOptions,
|
| 192 |
+
) -> Self {
|
| 193 |
+
Self {
|
| 194 |
+
connection,
|
| 195 |
+
options,
|
| 196 |
+
attempt_permit: None,
|
| 197 |
+
noise_context: None,
|
| 198 |
+
}
|
| 199 |
+
}
|
| 200 |
+
|
| 201 |
+
fn with_noise_context(
|
| 202 |
+
connection: JsonRpcConnection,
|
| 203 |
+
options: ExecServerClientConnectOptions,
|
| 204 |
+
noise_context: NoiseInitializeContext,
|
| 205 |
+
) -> Self {
|
| 206 |
+
Self {
|
| 207 |
+
connection,
|
| 208 |
+
options,
|
| 209 |
+
attempt_permit: None,
|
| 210 |
+
noise_context: Some(noise_context),
|
| 211 |
+
}
|
| 212 |
+
}
|
| 213 |
+
|
| 214 |
+
pub(crate) fn with_attempt_permit(
|
| 215 |
+
connection: JsonRpcConnection,
|
| 216 |
+
options: ExecServerClientConnectOptions,
|
| 217 |
+
attempt_permit: OwnedSemaphorePermit,
|
| 218 |
+
) -> Self {
|
| 219 |
+
Self {
|
| 220 |
+
connection,
|
| 221 |
+
options,
|
| 222 |
+
attempt_permit: Some(attempt_permit),
|
| 223 |
+
noise_context: None,
|
| 224 |
+
}
|
| 225 |
+
}
|
| 226 |
+
|
| 227 |
+
pub(crate) fn into_parts(
|
| 228 |
+
self,
|
| 229 |
+
) -> (
|
| 230 |
+
JsonRpcConnection,
|
| 231 |
+
ExecServerClientConnectOptions,
|
| 232 |
+
Option<OwnedSemaphorePermit>,
|
| 233 |
+
Option<NoiseInitializeContext>,
|
| 234 |
+
) {
|
| 235 |
+
(
|
| 236 |
+
self.connection,
|
| 237 |
+
self.options,
|
| 238 |
+
self.attempt_permit,
|
| 239 |
+
self.noise_context,
|
| 240 |
+
)
|
| 241 |
+
}
|
| 242 |
+
}
|
| 243 |
+
|
| 244 |
+
/// Reopens the transport for one logical exec-server client session.
|
| 245 |
+
///
|
| 246 |
+
/// URL connections reuse their configured endpoint. Noise connections retain
|
| 247 |
+
/// the harness identity but fetch a fresh single-use authorization bundle for
|
| 248 |
+
/// every physical connection attempt.
|
| 249 |
+
#[derive(Clone)]
|
| 250 |
+
pub(crate) enum ExecServerReconnectStrategy {
|
| 251 |
+
Accepted(AcceptedConnectionSource),
|
| 252 |
+
WebSocket {
|
| 253 |
+
args: RemoteExecServerConnectArgs,
|
| 254 |
+
http_headers: HeaderMap,
|
| 255 |
+
},
|
| 256 |
+
NoiseRendezvous {
|
| 257 |
+
// The executor that created the session, not the latest recovery lookup.
|
| 258 |
+
executor_public_key: crate::NoiseChannelPublicKey,
|
| 259 |
+
provider: Arc<dyn NoiseRendezvousConnectProvider>,
|
| 260 |
+
identity: NoiseChannelIdentity,
|
| 261 |
+
client_name: String,
|
| 262 |
+
connect_timeout: Duration,
|
| 263 |
+
initialize_timeout: Duration,
|
| 264 |
+
http_client_factory: HttpClientFactory,
|
| 265 |
+
},
|
| 266 |
+
}
|
| 267 |
+
|
| 268 |
+
impl ExecServerReconnectStrategy {
|
| 269 |
+
pub(crate) async fn resume(
|
| 270 |
+
&self,
|
| 271 |
+
session_id: &str,
|
| 272 |
+
) -> Result<ReconnectAttempt, ExecServerError> {
|
| 273 |
+
match self {
|
| 274 |
+
Self::Accepted(source) => source.next_connection(session_id).await,
|
| 275 |
+
Self::WebSocket { args, http_headers } => {
|
| 276 |
+
let mut args = args.clone();
|
| 277 |
+
args.resume_session_id = Some(session_id.to_string());
|
| 278 |
+
let connection =
|
| 279 |
+
ExecServerClient::open_websocket_connection(&args, http_headers).await?;
|
| 280 |
+
Ok(ReconnectAttempt::new(connection, args.into()))
|
| 281 |
+
}
|
| 282 |
+
Self::NoiseRendezvous {
|
| 283 |
+
executor_public_key: _,
|
| 284 |
+
provider,
|
| 285 |
+
identity,
|
| 286 |
+
client_name,
|
| 287 |
+
connect_timeout,
|
| 288 |
+
initialize_timeout,
|
| 289 |
+
http_client_factory,
|
| 290 |
+
} => {
|
| 291 |
+
let bundle = provider.connect_bundle(identity.public_key()).await?;
|
| 292 |
+
let opened = ExecServerClient::open_noise_rendezvous_connection(
|
| 293 |
+
NoiseRendezvousConnectArgs {
|
| 294 |
+
bundle,
|
| 295 |
+
harness_identity: identity.clone(),
|
| 296 |
+
client_name: client_name.clone(),
|
| 297 |
+
connect_timeout: *connect_timeout,
|
| 298 |
+
initialize_timeout: *initialize_timeout,
|
| 299 |
+
resume_session_id: Some(session_id.to_string()),
|
| 300 |
+
http_client_factory: http_client_factory.clone(),
|
| 301 |
+
},
|
| 302 |
+
)
|
| 303 |
+
.await?;
|
| 304 |
+
let ready = ExecServerClient::finish_noise_rendezvous_connection(opened).await?;
|
| 305 |
+
Ok(ReconnectAttempt::with_noise_context(
|
| 306 |
+
ready.connection,
|
| 307 |
+
ready.options,
|
| 308 |
+
ready.noise_context,
|
| 309 |
+
))
|
| 310 |
+
}
|
| 311 |
+
}
|
| 312 |
+
}
|
| 313 |
+
}
|
| 314 |
+
|
| 315 |
+
impl ExecServerClient {
|
| 316 |
+
/// Open the selected transport and run the common JSON-RPC initialization.
|
| 317 |
+
/// Noise connection details are fetched here so reconnects get a fresh URL
|
| 318 |
+
/// and authorization without replacing the harness identity.
|
| 319 |
+
pub(crate) async fn connect_for_transport(
|
| 320 |
+
transport_params: ExecServerTransportParams,
|
| 321 |
+
http_client_factory: HttpClientFactory,
|
| 322 |
+
) -> Result<Self, ExecServerError> {
|
| 323 |
+
let (transport_params, deferred_readiness) = match transport_params {
|
| 324 |
+
ExecServerTransportParams::Deferred(deferred) => {
|
| 325 |
+
(deferred.transport, Some(deferred.readiness))
|
| 326 |
+
}
|
| 327 |
+
transport_params => (transport_params, None),
|
| 328 |
+
};
|
| 329 |
+
|
| 330 |
+
if let Some(mut readiness) = deferred_readiness {
|
| 331 |
+
let provisioning_result = readiness
|
| 332 |
+
.wait_for(Option::is_some)
|
| 333 |
+
.await
|
| 334 |
+
.map_err(|_| {
|
| 335 |
+
ExecServerError::Disconnected(
|
| 336 |
+
"environment unavailable: environment provisioning ended before completion"
|
| 337 |
+
.to_string(),
|
| 338 |
+
)
|
| 339 |
+
})?
|
| 340 |
+
.clone()
|
| 341 |
+
.ok_or_else(|| {
|
| 342 |
+
ExecServerError::Disconnected(
|
| 343 |
+
"environment unavailable: provisioning remained pending after completion"
|
| 344 |
+
.to_string(),
|
| 345 |
+
)
|
| 346 |
+
})?;
|
| 347 |
+
provisioning_result.map_err(ExecServerError::ProvisioningFailed)?;
|
| 348 |
+
}
|
| 349 |
+
|
| 350 |
+
let websocket = match transport_params {
|
| 351 |
+
ExecServerTransportParams::Deferred(_) => {
|
| 352 |
+
return Err(ExecServerError::Protocol(
|
| 353 |
+
"nested deferred exec-server transports are unsupported".to_string(),
|
| 354 |
+
));
|
| 355 |
+
}
|
| 356 |
+
ExecServerTransportParams::WebSocketUrl {
|
| 357 |
+
websocket_url,
|
| 358 |
+
connect_timeout,
|
| 359 |
+
initialize_timeout,
|
| 360 |
+
http_headers,
|
| 361 |
+
} => (
|
| 362 |
+
websocket_url,
|
| 363 |
+
connect_timeout,
|
| 364 |
+
initialize_timeout,
|
| 365 |
+
http_headers,
|
| 366 |
+
),
|
| 367 |
+
ExecServerTransportParams::NoiseRendezvous { provider, identity } => {
|
| 368 |
+
let (ready, executor_public_key) = Self::open_initial_noise_rendezvous_connection(
|
| 369 |
+
&provider,
|
| 370 |
+
&identity,
|
| 371 |
+
http_client_factory.clone(),
|
| 372 |
+
)
|
| 373 |
+
.await?;
|
| 374 |
+
let reconnect_strategy = ExecServerReconnectStrategy::NoiseRendezvous {
|
| 375 |
+
executor_public_key,
|
| 376 |
+
provider,
|
| 377 |
+
identity,
|
| 378 |
+
client_name: ENVIRONMENT_CLIENT_NAME.to_string(),
|
| 379 |
+
connect_timeout: DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT,
|
| 380 |
+
initialize_timeout: DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT,
|
| 381 |
+
http_client_factory,
|
| 382 |
+
};
|
| 383 |
+
return Self::connect_with_recovery_and_noise_context(
|
| 384 |
+
ready.connection,
|
| 385 |
+
ready.options,
|
| 386 |
+
Some(reconnect_strategy),
|
| 387 |
+
ready.noise_context,
|
| 388 |
+
)
|
| 389 |
+
.await;
|
| 390 |
+
}
|
| 391 |
+
ExecServerTransportParams::StdioCommand {
|
| 392 |
+
command,
|
| 393 |
+
initialize_timeout,
|
| 394 |
+
} => {
|
| 395 |
+
return Self::connect_stdio_command(StdioExecServerConnectArgs {
|
| 396 |
+
command,
|
| 397 |
+
client_name: ENVIRONMENT_CLIENT_NAME.to_string(),
|
| 398 |
+
initialize_timeout,
|
| 399 |
+
resume_session_id: None,
|
| 400 |
+
})
|
| 401 |
+
.await;
|
| 402 |
+
}
|
| 403 |
+
};
|
| 404 |
+
let (websocket_url, connect_timeout, initialize_timeout, http_headers) = websocket;
|
| 405 |
+
Self::connect_websocket_with_headers(
|
| 406 |
+
RemoteExecServerConnectArgs {
|
| 407 |
+
websocket_url,
|
| 408 |
+
client_name: ENVIRONMENT_CLIENT_NAME.to_string(),
|
| 409 |
+
connect_timeout,
|
| 410 |
+
initialize_timeout,
|
| 411 |
+
resume_session_id: None,
|
| 412 |
+
http_client_factory,
|
| 413 |
+
},
|
| 414 |
+
http_headers,
|
| 415 |
+
)
|
| 416 |
+
.await
|
| 417 |
+
}
|
| 418 |
+
|
| 419 |
+
#[tracing::instrument(name = "codex.exec_server.remote.noise.connect", skip_all)]
|
| 420 |
+
async fn open_initial_noise_rendezvous_connection(
|
| 421 |
+
provider: &Arc<dyn NoiseRendezvousConnectProvider>,
|
| 422 |
+
identity: &NoiseChannelIdentity,
|
| 423 |
+
http_client_factory: HttpClientFactory,
|
| 424 |
+
) -> Result<(ReadyNoiseRendezvousConnection, crate::NoiseChannelPublicKey), ExecServerError>
|
| 425 |
+
{
|
| 426 |
+
let open_connection = |bundle: NoiseRendezvousConnectBundle| {
|
| 427 |
+
Self::open_noise_rendezvous_connection(NoiseRendezvousConnectArgs {
|
| 428 |
+
bundle,
|
| 429 |
+
harness_identity: identity.clone(),
|
| 430 |
+
client_name: ENVIRONMENT_CLIENT_NAME.to_string(),
|
| 431 |
+
connect_timeout: DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT,
|
| 432 |
+
initialize_timeout: DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT,
|
| 433 |
+
resume_session_id: None,
|
| 434 |
+
http_client_factory: http_client_factory.clone(),
|
| 435 |
+
})
|
| 436 |
+
};
|
| 437 |
+
let mut deadline = Instant::now() + INITIAL_REGISTRY_OPERATION_TIMEOUT;
|
| 438 |
+
let retry_key = uuid::Uuid::new_v4().to_string();
|
| 439 |
+
let mut retries = 0;
|
| 440 |
+
let mut refreshed_unauthorized_bundle = false;
|
| 441 |
+
let connect_bundle = || async {
|
| 442 |
+
timeout(
|
| 443 |
+
INITIAL_REGISTRY_REQUEST_TIMEOUT,
|
| 444 |
+
provider.connect_bundle(identity.public_key()),
|
| 445 |
+
)
|
| 446 |
+
.await
|
| 447 |
+
.unwrap_or_else(|_| {
|
| 448 |
+
Err(ExecServerError::EnvironmentRegistryRequest(
|
| 449 |
+
codex_http_client::RouteAwareRequestError::Timeout,
|
| 450 |
+
))
|
| 451 |
+
})
|
| 452 |
+
};
|
| 453 |
+
let mut result = connect_bundle().await;
|
| 454 |
+
loop {
|
| 455 |
+
let bundle = match result {
|
| 456 |
+
Ok(bundle) => bundle,
|
| 457 |
+
Err(error)
|
| 458 |
+
if is_retryable_registry_error(&error)
|
| 459 |
+
&& retries < INITIAL_REGISTRY_MAX_RETRIES =>
|
| 460 |
+
{
|
| 461 |
+
// Session resumption owns its separate recovery deadline.
|
| 462 |
+
let delay = registry_recovery_retry_delay(&retry_key, retries);
|
| 463 |
+
retries += 1;
|
| 464 |
+
result = match timeout_at(deadline, async {
|
| 465 |
+
sleep(delay).await;
|
| 466 |
+
connect_bundle().await
|
| 467 |
+
})
|
| 468 |
+
.await
|
| 469 |
+
{
|
| 470 |
+
Ok(result) => result,
|
| 471 |
+
Err(_) => return Err(error),
|
| 472 |
+
};
|
| 473 |
+
continue;
|
| 474 |
+
}
|
| 475 |
+
Err(error) => return Err(error),
|
| 476 |
+
};
|
| 477 |
+
let executor_public_key = bundle.executor_public_key.clone();
|
| 478 |
+
match open_connection(bundle).await {
|
| 479 |
+
Err(error)
|
| 480 |
+
if !refreshed_unauthorized_bundle
|
| 481 |
+
&& matches!(
|
| 482 |
+
&error,
|
| 483 |
+
ExecServerError::WebSocketConnect { source, .. }
|
| 484 |
+
if matches!(
|
| 485 |
+
source,
|
| 486 |
+
tokio_tungstenite::tungstenite::Error::Http(response)
|
| 487 |
+
if response.status().as_u16() == 401
|
| 488 |
+
)
|
| 489 |
+
) =>
|
| 490 |
+
{
|
| 491 |
+
refreshed_unauthorized_bundle = true;
|
| 492 |
+
deadline = Instant::now() + INITIAL_REGISTRY_OPERATION_TIMEOUT;
|
| 493 |
+
retries = 0;
|
| 494 |
+
result = connect_bundle().await;
|
| 495 |
+
}
|
| 496 |
+
result => {
|
| 497 |
+
let opened = result?;
|
| 498 |
+
let ready = Self::finish_noise_rendezvous_connection(opened).await?;
|
| 499 |
+
return Ok((ready, executor_public_key));
|
| 500 |
+
}
|
| 501 |
+
}
|
| 502 |
+
}
|
| 503 |
+
}
|
| 504 |
+
|
| 505 |
+
pub async fn connect_websocket(
|
| 506 |
+
args: RemoteExecServerConnectArgs,
|
| 507 |
+
) -> Result<Self, ExecServerError> {
|
| 508 |
+
Self::connect_websocket_with_headers(args, HeaderMap::new()).await
|
| 509 |
+
}
|
| 510 |
+
|
| 511 |
+
async fn connect_websocket_with_headers(
|
| 512 |
+
args: RemoteExecServerConnectArgs,
|
| 513 |
+
http_headers: HeaderMap,
|
| 514 |
+
) -> Result<Self, ExecServerError> {
|
| 515 |
+
let connection = Self::open_websocket_connection(&args, &http_headers).await?;
|
| 516 |
+
let options = args.clone().into();
|
| 517 |
+
Self::connect_with_recovery(
|
| 518 |
+
connection,
|
| 519 |
+
options,
|
| 520 |
+
Some(ExecServerReconnectStrategy::WebSocket { args, http_headers }),
|
| 521 |
+
)
|
| 522 |
+
.await
|
| 523 |
+
}
|
| 524 |
+
|
| 525 |
+
pub(crate) async fn open_websocket_connection(
|
| 526 |
+
args: &RemoteExecServerConnectArgs,
|
| 527 |
+
http_headers: &HeaderMap,
|
| 528 |
+
) -> Result<JsonRpcConnection, ExecServerError> {
|
| 529 |
+
ensure_rustls_crypto_provider();
|
| 530 |
+
let websocket_url = args.websocket_url.clone();
|
| 531 |
+
let connect_timeout = args.connect_timeout;
|
| 532 |
+
let mut request = websocket_url
|
| 533 |
+
.as_str()
|
| 534 |
+
.into_client_request()
|
| 535 |
+
.map_err(|source| ExecServerError::WebSocketConnect {
|
| 536 |
+
url: websocket_url.clone(),
|
| 537 |
+
source,
|
| 538 |
+
})?;
|
| 539 |
+
request.headers_mut().extend(http_headers.clone());
|
| 540 |
+
let connector = WebSocketConnector::new_with_tls_mode(
|
| 541 |
+
&args.http_client_factory,
|
| 542 |
+
WebSocketTlsMode::TungsteniteDefault,
|
| 543 |
+
)
|
| 544 |
+
.map_err(|error| ExecServerError::WebSocketConfiguration(error.to_string()))?;
|
| 545 |
+
let stream = connect_websocket_request(
|
| 546 |
+
request,
|
| 547 |
+
websocket_url.clone(),
|
| 548 |
+
connector,
|
| 549 |
+
connect_timeout,
|
| 550 |
+
!http_headers.is_empty() && websocket_url.starts_with("ws://"),
|
| 551 |
+
)
|
| 552 |
+
.await?;
|
| 553 |
+
|
| 554 |
+
let connection_label = format!("exec-server websocket {websocket_url}");
|
| 555 |
+
let connection = if is_rendezvous_harness_url(&websocket_url) {
|
| 556 |
+
harness_connection_from_websocket(stream, connection_label)
|
| 557 |
+
} else {
|
| 558 |
+
JsonRpcConnection::from_websocket(stream, connection_label)
|
| 559 |
+
};
|
| 560 |
+
Ok(connection)
|
| 561 |
+
}
|
| 562 |
+
|
| 563 |
+
/// Connect to one exec-server through an authenticated rendezvous stream
|
| 564 |
+
/// using a caller-supplied single-use authorization bundle.
|
| 565 |
+
///
|
| 566 |
+
/// The executor key is pinned before JSON-RPC starts; the websocket carries
|
| 567 |
+
/// only ciphertext after that. Environment-managed connections use a
|
| 568 |
+
/// retained [`NoiseRendezvousConnectProvider`] so recovery can fetch a fresh
|
| 569 |
+
/// bundle for each reconnect.
|
| 570 |
+
#[tracing::instrument(
|
| 571 |
+
name = "codex.exec_server.remote.harness.connect",
|
| 572 |
+
skip_all,
|
| 573 |
+
fields(
|
| 574 |
+
otel.kind = "client",
|
| 575 |
+
otel.name = "codex.exec_server.remote.harness.connect",
|
| 576 |
+
)
|
| 577 |
+
)]
|
| 578 |
+
pub async fn connect_noise_rendezvous(
|
| 579 |
+
args: NoiseRendezvousConnectArgs,
|
| 580 |
+
) -> Result<Self, ExecServerError> {
|
| 581 |
+
let opened = Self::open_noise_rendezvous_connection(args).await?;
|
| 582 |
+
let ready = Self::finish_noise_rendezvous_connection(opened).await?;
|
| 583 |
+
Self::connect_with_recovery_and_noise_context(
|
| 584 |
+
ready.connection,
|
| 585 |
+
ready.options,
|
| 586 |
+
/*reconnect_strategy*/ None,
|
| 587 |
+
ready.noise_context,
|
| 588 |
+
)
|
| 589 |
+
.await
|
| 590 |
+
}
|
| 591 |
+
|
| 592 |
+
#[tracing::instrument(
|
| 593 |
+
name = "codex.exec_server.remote.noise.websocket_connect",
|
| 594 |
+
skip_all,
|
| 595 |
+
fields(
|
| 596 |
+
otel.kind = "client",
|
| 597 |
+
otel.name = "codex.exec_server.remote.noise.websocket_connect",
|
| 598 |
+
environment_id = %args.bundle.environment_id,
|
| 599 |
+
executor_registration_id = %args.bundle.executor_registration_id,
|
| 600 |
+
)
|
| 601 |
+
)]
|
| 602 |
+
async fn open_noise_rendezvous_connection(
|
| 603 |
+
args: NoiseRendezvousConnectArgs,
|
| 604 |
+
) -> Result<OpenNoiseRendezvousConnection, ExecServerError> {
|
| 605 |
+
ensure_rustls_crypto_provider();
|
| 606 |
+
// Keep the registry-issued URL, key, and authorization together for this
|
| 607 |
+
// connection attempt.
|
| 608 |
+
let NoiseRendezvousConnectArgs {
|
| 609 |
+
bundle,
|
| 610 |
+
harness_identity,
|
| 611 |
+
client_name,
|
| 612 |
+
connect_timeout,
|
| 613 |
+
initialize_timeout,
|
| 614 |
+
resume_session_id,
|
| 615 |
+
http_client_factory,
|
| 616 |
+
} = args;
|
| 617 |
+
let NoiseRendezvousConnectBundle {
|
| 618 |
+
websocket_url,
|
| 619 |
+
environment_id,
|
| 620 |
+
executor_registration_id,
|
| 621 |
+
executor_public_key,
|
| 622 |
+
harness_key_authorization,
|
| 623 |
+
} = bundle;
|
| 624 |
+
let diagnostic_url = websocket_url
|
| 625 |
+
.split(['?', '#'])
|
| 626 |
+
.next()
|
| 627 |
+
.unwrap_or(websocket_url.as_str())
|
| 628 |
+
.to_string();
|
| 629 |
+
let mut request = websocket_url
|
| 630 |
+
.as_str()
|
| 631 |
+
.into_client_request()
|
| 632 |
+
.map_err(|source| ExecServerError::WebSocketConnect {
|
| 633 |
+
url: diagnostic_url.clone(),
|
| 634 |
+
source,
|
| 635 |
+
})?;
|
| 636 |
+
request.headers_mut().extend(current_rendezvous_headers());
|
| 637 |
+
let (stream, _) = timeout(
|
| 638 |
+
connect_timeout,
|
| 639 |
+
WebSocketConnector::new_with_tls_mode(
|
| 640 |
+
&http_client_factory,
|
| 641 |
+
WebSocketTlsMode::TungsteniteDefault,
|
| 642 |
+
)
|
| 643 |
+
.map_err(|error| ExecServerError::WebSocketConfiguration(error.to_string()))?
|
| 644 |
+
.with_tcp_nodelay()
|
| 645 |
+
.connect(request, noise_relay_websocket_config()),
|
| 646 |
+
)
|
| 647 |
+
.await
|
| 648 |
+
.map_err(|_| ExecServerError::WebSocketConnectTimeout {
|
| 649 |
+
url: diagnostic_url.clone(),
|
| 650 |
+
timeout: connect_timeout,
|
| 651 |
+
})?
|
| 652 |
+
.map_err(|source| ExecServerError::WebSocketConnect {
|
| 653 |
+
url: diagnostic_url.clone(),
|
| 654 |
+
source,
|
| 655 |
+
})?;
|
| 656 |
+
|
| 657 |
+
let connection_label = format!("Noise exec-server rendezvous websocket {diagnostic_url}");
|
| 658 |
+
let connection = noise_harness_connection_from_websocket_with_readiness(
|
| 659 |
+
stream,
|
| 660 |
+
NoiseHarnessConnectionArgs {
|
| 661 |
+
connection_label,
|
| 662 |
+
environment_id,
|
| 663 |
+
executor_registration_id,
|
| 664 |
+
identity: harness_identity,
|
| 665 |
+
responder_public_key: executor_public_key,
|
| 666 |
+
harness_key_authorization,
|
| 667 |
+
},
|
| 668 |
+
);
|
| 669 |
+
Ok(OpenNoiseRendezvousConnection {
|
| 670 |
+
connection: connection.connection,
|
| 671 |
+
options: ExecServerClientConnectOptions {
|
| 672 |
+
client_name,
|
| 673 |
+
initialize_timeout,
|
| 674 |
+
resume_session_id,
|
| 675 |
+
},
|
| 676 |
+
handshake_ready: connection.handshake_ready,
|
| 677 |
+
})
|
| 678 |
+
}
|
| 679 |
+
|
| 680 |
+
#[tracing::instrument(
|
| 681 |
+
name = "codex.exec_server.remote.noise.handshake",
|
| 682 |
+
skip_all,
|
| 683 |
+
parent = initialize_span,
|
| 684 |
+
fields(
|
| 685 |
+
otel.kind = "client",
|
| 686 |
+
otel.name = "codex.exec_server.remote.noise.handshake",
|
| 687 |
+
)
|
| 688 |
+
)]
|
| 689 |
+
async fn wait_for_noise_handshake(
|
| 690 |
+
handshake_ready: &mut tokio::sync::oneshot::Receiver<()>,
|
| 691 |
+
deadline: Instant,
|
| 692 |
+
initialize_timeout: Duration,
|
| 693 |
+
initialize_span: &tracing::Span,
|
| 694 |
+
) -> Result<(), ExecServerError> {
|
| 695 |
+
match timeout_at(deadline, handshake_ready).await {
|
| 696 |
+
Ok(Ok(())) => Ok(()),
|
| 697 |
+
Ok(Err(_)) => Err(ExecServerError::Disconnected(
|
| 698 |
+
"Noise harness handshake failed before connection became ready".to_string(),
|
| 699 |
+
)),
|
| 700 |
+
Err(_) => Err(ExecServerError::InitializeTimedOut {
|
| 701 |
+
timeout: initialize_timeout,
|
| 702 |
+
}),
|
| 703 |
+
}
|
| 704 |
+
}
|
| 705 |
+
|
| 706 |
+
async fn finish_noise_rendezvous_connection(
|
| 707 |
+
mut connection: OpenNoiseRendezvousConnection,
|
| 708 |
+
) -> Result<ReadyNoiseRendezvousConnection, ExecServerError> {
|
| 709 |
+
// Preserve the legacy initialize request span as the post-WebSocket
|
| 710 |
+
// startup parent while making its two child operations visible.
|
| 711 |
+
let initialize_timeout = connection.options.initialize_timeout;
|
| 712 |
+
let noise_context = NoiseInitializeContext {
|
| 713 |
+
span: tracing::info_span!(
|
| 714 |
+
"codex.exec_server.request",
|
| 715 |
+
otel.kind = "client",
|
| 716 |
+
otel.name = "initialize",
|
| 717 |
+
method = "initialize",
|
| 718 |
+
),
|
| 719 |
+
timeout_for_error: initialize_timeout,
|
| 720 |
+
};
|
| 721 |
+
let deadline = Instant::now() + initialize_timeout;
|
| 722 |
+
let readiness = Self::wait_for_noise_handshake(
|
| 723 |
+
&mut connection.handshake_ready,
|
| 724 |
+
deadline,
|
| 725 |
+
initialize_timeout,
|
| 726 |
+
&noise_context.span,
|
| 727 |
+
)
|
| 728 |
+
.await;
|
| 729 |
+
if let Err(error) = readiness {
|
| 730 |
+
// Unlike the normal connect path, the connection has not reached
|
| 731 |
+
// RpcClient yet, so its Drop implementation cannot abort the
|
| 732 |
+
// transport task for us.
|
| 733 |
+
connection.connection.transport.terminate();
|
| 734 |
+
for task in &connection.connection.task_handles {
|
| 735 |
+
task.abort();
|
| 736 |
+
}
|
| 737 |
+
return Err(error);
|
| 738 |
+
}
|
| 739 |
+
let mut options = connection.options;
|
| 740 |
+
options.initialize_timeout = deadline.saturating_duration_since(Instant::now());
|
| 741 |
+
Ok(ReadyNoiseRendezvousConnection {
|
| 742 |
+
connection: connection.connection,
|
| 743 |
+
options,
|
| 744 |
+
noise_context,
|
| 745 |
+
})
|
| 746 |
+
}
|
| 747 |
+
|
| 748 |
+
pub(crate) async fn connect_stdio_command(
|
| 749 |
+
args: StdioExecServerConnectArgs,
|
| 750 |
+
) -> Result<Self, ExecServerError> {
|
| 751 |
+
let mut child = stdio_command_process(&args.command)
|
| 752 |
+
.stdin(Stdio::piped())
|
| 753 |
+
.stdout(Stdio::piped())
|
| 754 |
+
.stderr(Stdio::piped())
|
| 755 |
+
.spawn()
|
| 756 |
+
.map_err(ExecServerError::Spawn)?;
|
| 757 |
+
|
| 758 |
+
let stdin = child.stdin.take().ok_or_else(|| {
|
| 759 |
+
ExecServerError::Protocol("spawned exec-server command has no stdin".to_string())
|
| 760 |
+
})?;
|
| 761 |
+
let stdout = child.stdout.take().ok_or_else(|| {
|
| 762 |
+
ExecServerError::Protocol("spawned exec-server command has no stdout".to_string())
|
| 763 |
+
})?;
|
| 764 |
+
if let Some(stderr) = child.stderr.take() {
|
| 765 |
+
tokio::spawn(async move {
|
| 766 |
+
let mut lines = BufReader::new(stderr).lines();
|
| 767 |
+
loop {
|
| 768 |
+
match lines.next_line().await {
|
| 769 |
+
Ok(Some(line)) => debug!("exec-server stdio stderr: {line}"),
|
| 770 |
+
Ok(None) => break,
|
| 771 |
+
Err(err) => {
|
| 772 |
+
warn!("failed to read exec-server stdio stderr: {err}");
|
| 773 |
+
break;
|
| 774 |
+
}
|
| 775 |
+
}
|
| 776 |
+
}
|
| 777 |
+
});
|
| 778 |
+
}
|
| 779 |
+
|
| 780 |
+
Self::connect(
|
| 781 |
+
JsonRpcConnection::from_stdio(stdout, stdin, "exec-server stdio command".to_string())
|
| 782 |
+
.with_child_process(child),
|
| 783 |
+
args.into(),
|
| 784 |
+
)
|
| 785 |
+
.await
|
| 786 |
+
}
|
| 787 |
+
}
|
| 788 |
+
|
| 789 |
+
fn is_rendezvous_harness_url(websocket_url: &str) -> bool {
|
| 790 |
+
let Some((_path, query)) = websocket_url.split_once('?') else {
|
| 791 |
+
return false;
|
| 792 |
+
};
|
| 793 |
+
query
|
| 794 |
+
.split('&')
|
| 795 |
+
.filter_map(|pair| pair.split_once('='))
|
| 796 |
+
.any(|(key, value)| key == "role" && value == "harness")
|
| 797 |
+
}
|
| 798 |
+
|
| 799 |
+
fn stdio_command_process(stdio_command: &StdioExecServerCommand) -> Command {
|
| 800 |
+
let mut command = Command::new(&stdio_command.program);
|
| 801 |
+
command.args(&stdio_command.args);
|
| 802 |
+
command.envs(&stdio_command.env);
|
| 803 |
+
scrub_non_inheritable_env_vars(command.as_std_mut());
|
| 804 |
+
if let Some(cwd) = &stdio_command.cwd {
|
| 805 |
+
command.current_dir(cwd);
|
| 806 |
+
}
|
| 807 |
+
#[cfg(unix)]
|
| 808 |
+
command.process_group(0);
|
| 809 |
+
command
|
| 810 |
+
}
|
| 811 |
+
|
| 812 |
+
#[cfg(test)]
|
| 813 |
+
#[path = "client_transport_tests.rs"]
|
| 814 |
+
mod tests;
|
codex-rs/exec-server/src/client_transport_tests.rs
ADDED
|
@@ -0,0 +1,581 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::collections::VecDeque;
|
| 2 |
+
use std::future::Future;
|
| 3 |
+
use std::sync::Arc;
|
| 4 |
+
use std::sync::Mutex;
|
| 5 |
+
|
| 6 |
+
use anyhow::Result;
|
| 7 |
+
use codex_exec_server_protocol::JSONRPCMessage;
|
| 8 |
+
use futures::FutureExt;
|
| 9 |
+
use futures::SinkExt;
|
| 10 |
+
use futures::StreamExt;
|
| 11 |
+
use futures::future::BoxFuture;
|
| 12 |
+
use pretty_assertions::assert_eq;
|
| 13 |
+
use tokio::io::AsyncBufReadExt;
|
| 14 |
+
use tokio::io::AsyncReadExt;
|
| 15 |
+
use tokio::io::AsyncWriteExt;
|
| 16 |
+
use tokio::io::BufReader;
|
| 17 |
+
use tokio::io::duplex;
|
| 18 |
+
use tokio::net::TcpListener;
|
| 19 |
+
use tokio_tungstenite::accept_async;
|
| 20 |
+
use tokio_tungstenite::tungstenite::Message;
|
| 21 |
+
|
| 22 |
+
use super::ExecServerClient;
|
| 23 |
+
use super::ExecServerReconnectStrategy;
|
| 24 |
+
use super::INITIAL_REGISTRY_MAX_RETRIES;
|
| 25 |
+
use super::INITIAL_REGISTRY_OPERATION_TIMEOUT;
|
| 26 |
+
use super::INITIAL_REGISTRY_REQUEST_TIMEOUT;
|
| 27 |
+
use crate::ExecServerError;
|
| 28 |
+
use crate::NoiseChannelIdentity;
|
| 29 |
+
use crate::NoiseChannelPublicKey;
|
| 30 |
+
use crate::NoiseRendezvousConnectArgs;
|
| 31 |
+
use crate::NoiseRendezvousConnectBundle;
|
| 32 |
+
use crate::NoiseRendezvousConnectProvider;
|
| 33 |
+
use crate::client::NoiseInitializeContext;
|
| 34 |
+
use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT;
|
| 35 |
+
use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT;
|
| 36 |
+
use crate::client_api::ExecServerClientConnectOptions;
|
| 37 |
+
use crate::connection::JsonRpcConnection;
|
| 38 |
+
use crate::noise_channel::PendingResponderHandshake;
|
| 39 |
+
use crate::noise_channel::noise_channel_prologue;
|
| 40 |
+
use crate::protocol::INITIALIZE_METHOD;
|
| 41 |
+
use crate::relay::RelayFrameBodyKind;
|
| 42 |
+
use crate::relay::decode_relay_message_frame;
|
| 43 |
+
use crate::relay::encode_relay_message_frame;
|
| 44 |
+
use crate::relay_proto::RelayMessageFrame;
|
| 45 |
+
|
| 46 |
+
#[derive(Default)]
|
| 47 |
+
struct SequenceNoiseConnectProvider {
|
| 48 |
+
bundles:
|
| 49 |
+
Mutex<VecDeque<BoxFuture<'static, Result<NoiseRendezvousConnectBundle, ExecServerError>>>>,
|
| 50 |
+
returned_urls: Mutex<Vec<String>>,
|
| 51 |
+
requested_keys: Mutex<Vec<NoiseChannelPublicKey>>,
|
| 52 |
+
}
|
| 53 |
+
|
| 54 |
+
impl SequenceNoiseConnectProvider {
|
| 55 |
+
fn push_response(
|
| 56 |
+
&self,
|
| 57 |
+
response: impl Future<Output = Result<NoiseRendezvousConnectBundle, ExecServerError>>
|
| 58 |
+
+ Send
|
| 59 |
+
+ 'static,
|
| 60 |
+
) {
|
| 61 |
+
self.bundles.lock().unwrap().push_back(response.boxed());
|
| 62 |
+
}
|
| 63 |
+
|
| 64 |
+
fn push_error(&self, error: ExecServerError) {
|
| 65 |
+
self.push_response(futures::future::ready(Err(error)));
|
| 66 |
+
}
|
| 67 |
+
|
| 68 |
+
fn push_pending(&self) {
|
| 69 |
+
self.push_response(futures::future::pending());
|
| 70 |
+
}
|
| 71 |
+
|
| 72 |
+
fn requested_keys(&self) -> Vec<NoiseChannelPublicKey> {
|
| 73 |
+
self.requested_keys.lock().unwrap().clone()
|
| 74 |
+
}
|
| 75 |
+
|
| 76 |
+
fn assert_requested_identity(&self, identity: &NoiseChannelIdentity, requests: usize) {
|
| 77 |
+
assert_eq!(self.requested_keys(), vec![identity.public_key(); requests]);
|
| 78 |
+
}
|
| 79 |
+
|
| 80 |
+
fn returned_urls(&self) -> Vec<String> {
|
| 81 |
+
self.returned_urls
|
| 82 |
+
.lock()
|
| 83 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner)
|
| 84 |
+
.clone()
|
| 85 |
+
}
|
| 86 |
+
|
| 87 |
+
async fn connect(
|
| 88 |
+
self: &Arc<Self>,
|
| 89 |
+
identity: &NoiseChannelIdentity,
|
| 90 |
+
) -> Result<
|
| 91 |
+
(
|
| 92 |
+
super::JsonRpcConnection,
|
| 93 |
+
super::ExecServerClientConnectOptions,
|
| 94 |
+
),
|
| 95 |
+
ExecServerError,
|
| 96 |
+
> {
|
| 97 |
+
let provider: Arc<dyn NoiseRendezvousConnectProvider> = self.clone();
|
| 98 |
+
ExecServerClient::open_initial_noise_rendezvous_connection(
|
| 99 |
+
&provider,
|
| 100 |
+
identity,
|
| 101 |
+
codex_http_client::HttpClientFactory::new(
|
| 102 |
+
codex_http_client::OutboundProxyPolicy::ReqwestDefault,
|
| 103 |
+
),
|
| 104 |
+
)
|
| 105 |
+
.await
|
| 106 |
+
.map(|(ready, _)| (ready.connection, ready.options))
|
| 107 |
+
}
|
| 108 |
+
}
|
| 109 |
+
|
| 110 |
+
impl NoiseRendezvousConnectProvider for SequenceNoiseConnectProvider {
|
| 111 |
+
fn connect_bundle(
|
| 112 |
+
&self,
|
| 113 |
+
harness_public_key: NoiseChannelPublicKey,
|
| 114 |
+
) -> BoxFuture<'_, Result<NoiseRendezvousConnectBundle, ExecServerError>> {
|
| 115 |
+
self.requested_keys.lock().unwrap().push(harness_public_key);
|
| 116 |
+
let response = self
|
| 117 |
+
.bundles
|
| 118 |
+
.lock()
|
| 119 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner)
|
| 120 |
+
.pop_front()
|
| 121 |
+
.expect("test Noise provider exhausted");
|
| 122 |
+
Box::pin(async move {
|
| 123 |
+
let result = response.await;
|
| 124 |
+
if let Ok(bundle) = &result {
|
| 125 |
+
self.returned_urls
|
| 126 |
+
.lock()
|
| 127 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner)
|
| 128 |
+
.push(bundle.websocket_url.clone());
|
| 129 |
+
}
|
| 130 |
+
result
|
| 131 |
+
})
|
| 132 |
+
}
|
| 133 |
+
}
|
| 134 |
+
|
| 135 |
+
fn test_bundle(websocket_url: String) -> Result<NoiseRendezvousConnectBundle> {
|
| 136 |
+
Ok(NoiseRendezvousConnectBundle {
|
| 137 |
+
websocket_url,
|
| 138 |
+
environment_id: "environment".to_string(),
|
| 139 |
+
executor_registration_id: "registration".to_string(),
|
| 140 |
+
executor_public_key: NoiseChannelIdentity::generate()?.public_key(),
|
| 141 |
+
harness_key_authorization: "authorization".to_string(),
|
| 142 |
+
})
|
| 143 |
+
}
|
| 144 |
+
|
| 145 |
+
fn registry_error(status: http::StatusCode, code: &str) -> ExecServerError {
|
| 146 |
+
ExecServerError::EnvironmentRegistryHttp {
|
| 147 |
+
status,
|
| 148 |
+
code: Some(code.to_string()),
|
| 149 |
+
message: "registry unavailable".to_string(),
|
| 150 |
+
}
|
| 151 |
+
}
|
| 152 |
+
|
| 153 |
+
#[tokio::test]
|
| 154 |
+
async fn noise_handshake_uses_initialize_timeout() -> Result<()> {
|
| 155 |
+
let listener = TcpListener::bind("127.0.0.1:0").await?;
|
| 156 |
+
let websocket_url = format!("ws://{}", listener.local_addr()?);
|
| 157 |
+
let server = tokio::spawn(async move {
|
| 158 |
+
let (socket, _) = listener.accept().await?;
|
| 159 |
+
let mut websocket = accept_async(socket).await?;
|
| 160 |
+
// Drain the frames sent before the harness waits for the responder,
|
| 161 |
+
// then verify that a timed-out readiness wait closes the socket.
|
| 162 |
+
assert!(websocket.next().await.is_some());
|
| 163 |
+
assert!(websocket.next().await.is_some());
|
| 164 |
+
let closed =
|
| 165 |
+
tokio::time::timeout(std::time::Duration::from_secs(1), websocket.next()).await?;
|
| 166 |
+
assert!(
|
| 167 |
+
matches!(closed, None | Some(Ok(Message::Close(_))) | Some(Err(_))),
|
| 168 |
+
"timed-out Noise handshake must close its websocket"
|
| 169 |
+
);
|
| 170 |
+
anyhow::Ok(())
|
| 171 |
+
});
|
| 172 |
+
let initialize_timeout = std::time::Duration::from_millis(1);
|
| 173 |
+
let opened = ExecServerClient::open_noise_rendezvous_connection(NoiseRendezvousConnectArgs {
|
| 174 |
+
bundle: NoiseRendezvousConnectBundle {
|
| 175 |
+
websocket_url,
|
| 176 |
+
environment_id: "environment".to_string(),
|
| 177 |
+
executor_registration_id: "registration".to_string(),
|
| 178 |
+
executor_public_key: NoiseChannelIdentity::generate()?.public_key(),
|
| 179 |
+
harness_key_authorization: "authorization".to_string(),
|
| 180 |
+
},
|
| 181 |
+
harness_identity: NoiseChannelIdentity::generate()?,
|
| 182 |
+
client_name: "test".to_string(),
|
| 183 |
+
connect_timeout: DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT,
|
| 184 |
+
initialize_timeout,
|
| 185 |
+
resume_session_id: None,
|
| 186 |
+
http_client_factory: codex_http_client::HttpClientFactory::new(
|
| 187 |
+
codex_http_client::OutboundProxyPolicy::ReqwestDefault,
|
| 188 |
+
),
|
| 189 |
+
})
|
| 190 |
+
.await?;
|
| 191 |
+
|
| 192 |
+
let error = ExecServerClient::finish_noise_rendezvous_connection(opened)
|
| 193 |
+
.await
|
| 194 |
+
.err()
|
| 195 |
+
.expect("stalled Noise handshake must time out");
|
| 196 |
+
assert!(matches!(
|
| 197 |
+
error,
|
| 198 |
+
ExecServerError::InitializeTimedOut { timeout } if timeout == initialize_timeout
|
| 199 |
+
));
|
| 200 |
+
|
| 201 |
+
server.await??;
|
| 202 |
+
Ok(())
|
| 203 |
+
}
|
| 204 |
+
|
| 205 |
+
#[tokio::test]
|
| 206 |
+
async fn deferred_initialize_timeout_reports_configured_budget() {
|
| 207 |
+
let (client_stdin, server_reader) = duplex(1 << 20);
|
| 208 |
+
let (server_writer, client_stdout) = duplex(1 << 20);
|
| 209 |
+
let server = tokio::spawn(async move {
|
| 210 |
+
let _server_writer = server_writer;
|
| 211 |
+
let mut lines = BufReader::new(server_reader).lines();
|
| 212 |
+
let line = lines
|
| 213 |
+
.next_line()
|
| 214 |
+
.await
|
| 215 |
+
.expect("initialize read should succeed")
|
| 216 |
+
.expect("initialize request should arrive");
|
| 217 |
+
let request: JSONRPCMessage =
|
| 218 |
+
serde_json::from_str(&line).expect("initialize request should parse");
|
| 219 |
+
assert!(
|
| 220 |
+
matches!(
|
| 221 |
+
request,
|
| 222 |
+
JSONRPCMessage::Request(ref request) if request.method == INITIALIZE_METHOD
|
| 223 |
+
),
|
| 224 |
+
"expected initialize request, got {request:?}"
|
| 225 |
+
);
|
| 226 |
+
futures::future::pending::<()>().await;
|
| 227 |
+
});
|
| 228 |
+
let configured_timeout = std::time::Duration::from_secs(10);
|
| 229 |
+
let error = ExecServerClient::connect_with_recovery_and_noise_context(
|
| 230 |
+
JsonRpcConnection::from_stdio(
|
| 231 |
+
client_stdout,
|
| 232 |
+
client_stdin,
|
| 233 |
+
"timeout-test-client".to_string(),
|
| 234 |
+
),
|
| 235 |
+
ExecServerClientConnectOptions {
|
| 236 |
+
client_name: "timeout-test-client".to_string(),
|
| 237 |
+
initialize_timeout: std::time::Duration::from_millis(1),
|
| 238 |
+
resume_session_id: None,
|
| 239 |
+
},
|
| 240 |
+
/*reconnect_strategy*/ None,
|
| 241 |
+
NoiseInitializeContext {
|
| 242 |
+
span: tracing::info_span!("codex.exec_server.request"),
|
| 243 |
+
timeout_for_error: configured_timeout,
|
| 244 |
+
},
|
| 245 |
+
)
|
| 246 |
+
.await
|
| 247 |
+
.err()
|
| 248 |
+
.expect("initialize RPC must time out");
|
| 249 |
+
assert!(matches!(
|
| 250 |
+
error,
|
| 251 |
+
ExecServerError::InitializeTimedOut { timeout } if timeout == configured_timeout
|
| 252 |
+
));
|
| 253 |
+
server.abort();
|
| 254 |
+
let _ = server.await;
|
| 255 |
+
}
|
| 256 |
+
|
| 257 |
+
#[tokio::test(start_paused = true)]
|
| 258 |
+
async fn initial_noise_connection_bounds_offline_retries() -> Result<()> {
|
| 259 |
+
let sequence = Arc::new(SequenceNoiseConnectProvider::default());
|
| 260 |
+
for _ in 0..=INITIAL_REGISTRY_MAX_RETRIES {
|
| 261 |
+
sequence.push_error(registry_error(
|
| 262 |
+
http::StatusCode::CONFLICT,
|
| 263 |
+
"environment_offline",
|
| 264 |
+
));
|
| 265 |
+
}
|
| 266 |
+
let identity = NoiseChannelIdentity::generate()?;
|
| 267 |
+
let started = tokio::time::Instant::now();
|
| 268 |
+
let error = sequence
|
| 269 |
+
.connect(&identity)
|
| 270 |
+
.await
|
| 271 |
+
.err()
|
| 272 |
+
.expect("offline retries must end");
|
| 273 |
+
|
| 274 |
+
assert!(crate::client::is_environment_offline_error(&error));
|
| 275 |
+
let requests = sequence.requested_keys().len();
|
| 276 |
+
assert!((4..=INITIAL_REGISTRY_MAX_RETRIES as usize + 1).contains(&requests));
|
| 277 |
+
sequence.assert_requested_identity(&identity, requests);
|
| 278 |
+
assert!(started.elapsed() <= INITIAL_REGISTRY_OPERATION_TIMEOUT);
|
| 279 |
+
Ok(())
|
| 280 |
+
}
|
| 281 |
+
|
| 282 |
+
#[tokio::test(start_paused = true)]
|
| 283 |
+
async fn initial_noise_connection_bounds_a_stalled_retry_request() -> Result<()> {
|
| 284 |
+
let sequence = Arc::new(SequenceNoiseConnectProvider::default());
|
| 285 |
+
sequence.push_error(registry_error(
|
| 286 |
+
http::StatusCode::CONFLICT,
|
| 287 |
+
"environment_offline",
|
| 288 |
+
));
|
| 289 |
+
for _ in 0..INITIAL_REGISTRY_MAX_RETRIES {
|
| 290 |
+
sequence.push_pending();
|
| 291 |
+
}
|
| 292 |
+
let identity = NoiseChannelIdentity::generate()?;
|
| 293 |
+
let started = tokio::time::Instant::now();
|
| 294 |
+
let error = sequence
|
| 295 |
+
.connect(&identity)
|
| 296 |
+
.await
|
| 297 |
+
.err()
|
| 298 |
+
.expect("stalled retry must time out");
|
| 299 |
+
|
| 300 |
+
assert!(matches!(
|
| 301 |
+
error,
|
| 302 |
+
ExecServerError::EnvironmentRegistryRequest(error) if error.is_timeout()
|
| 303 |
+
));
|
| 304 |
+
assert_eq!(started.elapsed(), INITIAL_REGISTRY_OPERATION_TIMEOUT);
|
| 305 |
+
let requests = sequence.requested_keys().len();
|
| 306 |
+
assert!((2..=3).contains(&requests));
|
| 307 |
+
sequence.assert_requested_identity(&identity, requests);
|
| 308 |
+
Ok(())
|
| 309 |
+
}
|
| 310 |
+
|
| 311 |
+
#[tokio::test(start_paused = true)]
|
| 312 |
+
async fn initial_noise_connection_bounds_a_stalled_initial_request() -> Result<()> {
|
| 313 |
+
let sequence = Arc::new(SequenceNoiseConnectProvider::default());
|
| 314 |
+
for _ in 0..=INITIAL_REGISTRY_MAX_RETRIES {
|
| 315 |
+
sequence.push_pending();
|
| 316 |
+
}
|
| 317 |
+
let identity = NoiseChannelIdentity::generate()?;
|
| 318 |
+
let started = tokio::time::Instant::now();
|
| 319 |
+
|
| 320 |
+
let error = sequence
|
| 321 |
+
.connect(&identity)
|
| 322 |
+
.await
|
| 323 |
+
.err()
|
| 324 |
+
.expect("stalled initial request must time out");
|
| 325 |
+
|
| 326 |
+
assert!(matches!(
|
| 327 |
+
error,
|
| 328 |
+
ExecServerError::EnvironmentRegistryRequest(error) if error.is_timeout()
|
| 329 |
+
));
|
| 330 |
+
assert_eq!(started.elapsed(), INITIAL_REGISTRY_OPERATION_TIMEOUT);
|
| 331 |
+
let requests = sequence.requested_keys().len();
|
| 332 |
+
assert!((2..=3).contains(&requests));
|
| 333 |
+
sequence.assert_requested_identity(&identity, requests);
|
| 334 |
+
Ok(())
|
| 335 |
+
}
|
| 336 |
+
|
| 337 |
+
#[tokio::test(start_paused = true)]
|
| 338 |
+
async fn initial_noise_connection_retries_a_stalled_initial_request() -> Result<()> {
|
| 339 |
+
let sequence = Arc::new(SequenceNoiseConnectProvider::default());
|
| 340 |
+
sequence.push_pending();
|
| 341 |
+
sequence.push_error(registry_error(http::StatusCode::FORBIDDEN, "forbidden"));
|
| 342 |
+
let identity = NoiseChannelIdentity::generate()?;
|
| 343 |
+
let started = tokio::time::Instant::now();
|
| 344 |
+
|
| 345 |
+
let error = sequence
|
| 346 |
+
.connect(&identity)
|
| 347 |
+
.await
|
| 348 |
+
.err()
|
| 349 |
+
.expect("terminal response must stop the retry sequence");
|
| 350 |
+
|
| 351 |
+
assert!(matches!(
|
| 352 |
+
error,
|
| 353 |
+
ExecServerError::EnvironmentRegistryHttp {
|
| 354 |
+
status: http::StatusCode::FORBIDDEN,
|
| 355 |
+
..
|
| 356 |
+
}
|
| 357 |
+
));
|
| 358 |
+
assert!(started.elapsed() >= INITIAL_REGISTRY_REQUEST_TIMEOUT);
|
| 359 |
+
assert!(started.elapsed() < INITIAL_REGISTRY_OPERATION_TIMEOUT);
|
| 360 |
+
sequence.assert_requested_identity(&identity, /*requests*/ 2);
|
| 361 |
+
Ok(())
|
| 362 |
+
}
|
| 363 |
+
|
| 364 |
+
#[tokio::test(start_paused = true)]
|
| 365 |
+
async fn initial_noise_connection_retries_transient_registry_statuses() -> Result<()> {
|
| 366 |
+
for status in [
|
| 367 |
+
http::StatusCode::REQUEST_TIMEOUT,
|
| 368 |
+
http::StatusCode::TOO_MANY_REQUESTS,
|
| 369 |
+
http::StatusCode::INTERNAL_SERVER_ERROR,
|
| 370 |
+
http::StatusCode::BAD_GATEWAY,
|
| 371 |
+
http::StatusCode::SERVICE_UNAVAILABLE,
|
| 372 |
+
] {
|
| 373 |
+
let sequence = Arc::new(SequenceNoiseConnectProvider::default());
|
| 374 |
+
sequence.push_error(registry_error(status, "temporarily_unavailable"));
|
| 375 |
+
sequence.push_error(registry_error(http::StatusCode::FORBIDDEN, "forbidden"));
|
| 376 |
+
let identity = NoiseChannelIdentity::generate()?;
|
| 377 |
+
|
| 378 |
+
let error = sequence
|
| 379 |
+
.connect(&identity)
|
| 380 |
+
.await
|
| 381 |
+
.err()
|
| 382 |
+
.expect("terminal response must stop the retry sequence");
|
| 383 |
+
|
| 384 |
+
assert!(matches!(
|
| 385 |
+
error,
|
| 386 |
+
ExecServerError::EnvironmentRegistryHttp {
|
| 387 |
+
status: http::StatusCode::FORBIDDEN,
|
| 388 |
+
..
|
| 389 |
+
}
|
| 390 |
+
));
|
| 391 |
+
sequence.assert_requested_identity(&identity, /*requests*/ 2);
|
| 392 |
+
}
|
| 393 |
+
Ok(())
|
| 394 |
+
}
|
| 395 |
+
|
| 396 |
+
#[tokio::test(start_paused = true)]
|
| 397 |
+
async fn initial_noise_connection_retries_registry_request_timeouts() -> Result<()> {
|
| 398 |
+
let sequence = Arc::new(SequenceNoiseConnectProvider::default());
|
| 399 |
+
sequence.push_error(ExecServerError::EnvironmentRegistryRequest(
|
| 400 |
+
codex_http_client::RouteAwareRequestError::Timeout,
|
| 401 |
+
));
|
| 402 |
+
sequence.push_error(registry_error(http::StatusCode::FORBIDDEN, "forbidden"));
|
| 403 |
+
let identity = NoiseChannelIdentity::generate()?;
|
| 404 |
+
|
| 405 |
+
let error = sequence
|
| 406 |
+
.connect(&identity)
|
| 407 |
+
.await
|
| 408 |
+
.err()
|
| 409 |
+
.expect("terminal response must stop the retry sequence");
|
| 410 |
+
|
| 411 |
+
assert!(matches!(
|
| 412 |
+
error,
|
| 413 |
+
ExecServerError::EnvironmentRegistryHttp {
|
| 414 |
+
status: http::StatusCode::FORBIDDEN,
|
| 415 |
+
..
|
| 416 |
+
}
|
| 417 |
+
));
|
| 418 |
+
sequence.assert_requested_identity(&identity, /*requests*/ 2);
|
| 419 |
+
Ok(())
|
| 420 |
+
}
|
| 421 |
+
|
| 422 |
+
#[tokio::test(start_paused = true)]
|
| 423 |
+
async fn initial_noise_connection_does_not_retry_permanent_registry_errors() -> Result<()> {
|
| 424 |
+
for (status, code) in [
|
| 425 |
+
(http::StatusCode::UNAUTHORIZED, "unauthorized"),
|
| 426 |
+
(http::StatusCode::FORBIDDEN, "forbidden"),
|
| 427 |
+
(http::StatusCode::BAD_REQUEST, "bad_request"),
|
| 428 |
+
(http::StatusCode::NOT_FOUND, "environment_not_found"),
|
| 429 |
+
(http::StatusCode::CONFLICT, "registration_conflict"),
|
| 430 |
+
(http::StatusCode::CONFLICT, "route_unavailable"),
|
| 431 |
+
] {
|
| 432 |
+
// A terminal error must also stop a retry sequence already in progress.
|
| 433 |
+
for initial_offline in [false, true] {
|
| 434 |
+
let sequence = Arc::new(SequenceNoiseConnectProvider::default());
|
| 435 |
+
if initial_offline {
|
| 436 |
+
sequence.push_error(registry_error(
|
| 437 |
+
http::StatusCode::CONFLICT,
|
| 438 |
+
"environment_offline",
|
| 439 |
+
));
|
| 440 |
+
}
|
| 441 |
+
sequence.push_error(registry_error(status, code));
|
| 442 |
+
let identity = NoiseChannelIdentity::generate()?;
|
| 443 |
+
let error = sequence
|
| 444 |
+
.connect(&identity)
|
| 445 |
+
.await
|
| 446 |
+
.err()
|
| 447 |
+
.expect("other errors must propagate");
|
| 448 |
+
assert!(
|
| 449 |
+
matches!(error, ExecServerError::EnvironmentRegistryHttp { status: actual_status, code: Some(actual_code), .. } if actual_status == status && actual_code == code)
|
| 450 |
+
);
|
| 451 |
+
sequence.assert_requested_identity(&identity, 1 + usize::from(initial_offline));
|
| 452 |
+
}
|
| 453 |
+
}
|
| 454 |
+
Ok(())
|
| 455 |
+
}
|
| 456 |
+
|
| 457 |
+
#[tokio::test(start_paused = true)]
|
| 458 |
+
async fn noise_session_resume_leaves_offline_retries_to_recovery() -> Result<()> {
|
| 459 |
+
let sequence = Arc::new(SequenceNoiseConnectProvider::default());
|
| 460 |
+
sequence.push_error(registry_error(
|
| 461 |
+
http::StatusCode::CONFLICT,
|
| 462 |
+
"environment_offline",
|
| 463 |
+
));
|
| 464 |
+
let identity = NoiseChannelIdentity::generate()?;
|
| 465 |
+
let strategy = ExecServerReconnectStrategy::NoiseRendezvous {
|
| 466 |
+
executor_public_key: NoiseChannelIdentity::generate()?.public_key(),
|
| 467 |
+
provider: sequence.clone(),
|
| 468 |
+
identity: identity.clone(),
|
| 469 |
+
client_name: "test".to_string(),
|
| 470 |
+
connect_timeout: DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT,
|
| 471 |
+
initialize_timeout: DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT,
|
| 472 |
+
http_client_factory: codex_http_client::HttpClientFactory::new(
|
| 473 |
+
codex_http_client::OutboundProxyPolicy::ReqwestDefault,
|
| 474 |
+
),
|
| 475 |
+
};
|
| 476 |
+
let started = tokio::time::Instant::now();
|
| 477 |
+
let error = strategy
|
| 478 |
+
.resume("session")
|
| 479 |
+
.await
|
| 480 |
+
.err()
|
| 481 |
+
.expect("resume must return the offline error");
|
| 482 |
+
assert!(crate::client::is_environment_offline_error(&error));
|
| 483 |
+
assert_eq!(started.elapsed(), std::time::Duration::ZERO);
|
| 484 |
+
sequence.assert_requested_identity(&identity, /*requests*/ 1);
|
| 485 |
+
Ok(())
|
| 486 |
+
}
|
| 487 |
+
|
| 488 |
+
#[tokio::test]
|
| 489 |
+
async fn initial_noise_connection_refreshes_bundle_after_exhausting_initial_retries() -> Result<()>
|
| 490 |
+
{
|
| 491 |
+
let unauthorized_listener = TcpListener::bind("127.0.0.1:0").await?;
|
| 492 |
+
let unauthorized_url = format!("ws://{}", unauthorized_listener.local_addr()?);
|
| 493 |
+
let unauthorized_server = tokio::spawn(async move {
|
| 494 |
+
let (mut socket, _) = unauthorized_listener.accept().await?;
|
| 495 |
+
let mut request = [0_u8; 4096];
|
| 496 |
+
let _ = socket.read(&mut request).await?;
|
| 497 |
+
socket
|
| 498 |
+
.write_all(
|
| 499 |
+
b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
|
| 500 |
+
)
|
| 501 |
+
.await?;
|
| 502 |
+
socket.shutdown().await?;
|
| 503 |
+
anyhow::Ok(())
|
| 504 |
+
});
|
| 505 |
+
let accepted_listener = TcpListener::bind("127.0.0.1:0").await?;
|
| 506 |
+
let accepted_url = format!("ws://{}", accepted_listener.local_addr()?);
|
| 507 |
+
let executor_identity = NoiseChannelIdentity::generate()?;
|
| 508 |
+
let executor_public_key = executor_identity.public_key();
|
| 509 |
+
let accepted_server = tokio::spawn(async move {
|
| 510 |
+
let (socket, _) = accepted_listener.accept().await?;
|
| 511 |
+
let mut websocket = accept_async(socket).await?;
|
| 512 |
+
let Message::Binary(resume_payload) = websocket.next().await.unwrap()? else {
|
| 513 |
+
anyhow::bail!("expected Noise relay resume frame");
|
| 514 |
+
};
|
| 515 |
+
let resume = decode_relay_message_frame(resume_payload.as_ref())?;
|
| 516 |
+
assert_eq!(resume.validate()?, RelayFrameBodyKind::Resume);
|
| 517 |
+
let Message::Binary(handshake_payload) = websocket.next().await.unwrap()? else {
|
| 518 |
+
anyhow::bail!("expected Noise relay handshake frame");
|
| 519 |
+
};
|
| 520 |
+
let handshake = decode_relay_message_frame(handshake_payload.as_ref())?;
|
| 521 |
+
let stream_id = handshake.stream_id.clone();
|
| 522 |
+
let prologue = noise_channel_prologue("environment", "registration", &stream_id);
|
| 523 |
+
let pending = PendingResponderHandshake::read_request(
|
| 524 |
+
&executor_identity,
|
| 525 |
+
&prologue,
|
| 526 |
+
&handshake.into_handshake_payload()?,
|
| 527 |
+
)?;
|
| 528 |
+
let (_transport, response) = pending.complete()?;
|
| 529 |
+
websocket
|
| 530 |
+
.send(Message::Binary(
|
| 531 |
+
encode_relay_message_frame(&RelayMessageFrame::handshake(stream_id, response))
|
| 532 |
+
.into(),
|
| 533 |
+
))
|
| 534 |
+
.await?;
|
| 535 |
+
anyhow::Ok(())
|
| 536 |
+
});
|
| 537 |
+
let sequence = Arc::new(SequenceNoiseConnectProvider::default());
|
| 538 |
+
let unauthorized_bundle = test_bundle(unauthorized_url.clone())?;
|
| 539 |
+
let mut accepted_bundle = test_bundle(accepted_url.clone())?;
|
| 540 |
+
accepted_bundle.executor_public_key = executor_public_key;
|
| 541 |
+
sequence.push_response(async {
|
| 542 |
+
tokio::time::pause();
|
| 543 |
+
Err(registry_error(
|
| 544 |
+
http::StatusCode::CONFLICT,
|
| 545 |
+
"environment_offline",
|
| 546 |
+
))
|
| 547 |
+
});
|
| 548 |
+
for _ in 1..INITIAL_REGISTRY_MAX_RETRIES {
|
| 549 |
+
sequence.push_error(registry_error(
|
| 550 |
+
http::StatusCode::CONFLICT,
|
| 551 |
+
"environment_offline",
|
| 552 |
+
));
|
| 553 |
+
}
|
| 554 |
+
sequence.push_response(async move {
|
| 555 |
+
tokio::time::resume();
|
| 556 |
+
Ok(unauthorized_bundle)
|
| 557 |
+
});
|
| 558 |
+
sequence.push_response(async {
|
| 559 |
+
tokio::time::pause();
|
| 560 |
+
Err(registry_error(
|
| 561 |
+
http::StatusCode::CONFLICT,
|
| 562 |
+
"environment_offline",
|
| 563 |
+
))
|
| 564 |
+
});
|
| 565 |
+
sequence.push_response(async move {
|
| 566 |
+
tokio::time::resume();
|
| 567 |
+
Ok(accepted_bundle)
|
| 568 |
+
});
|
| 569 |
+
let identity = NoiseChannelIdentity::generate()?;
|
| 570 |
+
|
| 571 |
+
let _connection = sequence.connect(&identity).await?;
|
| 572 |
+
|
| 573 |
+
assert_eq!(
|
| 574 |
+
sequence.returned_urls(),
|
| 575 |
+
vec![unauthorized_url, accepted_url]
|
| 576 |
+
);
|
| 577 |
+
sequence.assert_requested_identity(&identity, INITIAL_REGISTRY_MAX_RETRIES as usize + 3);
|
| 578 |
+
unauthorized_server.await??;
|
| 579 |
+
accepted_server.await??;
|
| 580 |
+
Ok(())
|
| 581 |
+
}
|
codex-rs/exec-server/src/connection.rs
ADDED
|
@@ -0,0 +1,1042 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#[cfg(windows)]
|
| 2 |
+
use std::process::Stdio;
|
| 3 |
+
use std::sync::Arc;
|
| 4 |
+
use std::sync::atomic::AtomicBool;
|
| 5 |
+
use std::sync::atomic::Ordering;
|
| 6 |
+
use std::time::Duration;
|
| 7 |
+
use std::time::Instant;
|
| 8 |
+
|
| 9 |
+
use axum::extract::ws::Message as AxumWebSocketMessage;
|
| 10 |
+
use axum::extract::ws::WebSocket as AxumWebSocket;
|
| 11 |
+
use codex_exec_server_protocol::JSONRPCMessage;
|
| 12 |
+
use codex_exec_server_protocol::JSONRPCRequest;
|
| 13 |
+
use futures::Sink;
|
| 14 |
+
use futures::SinkExt;
|
| 15 |
+
use futures::Stream;
|
| 16 |
+
use futures::StreamExt;
|
| 17 |
+
use tokio::io::AsyncRead;
|
| 18 |
+
use tokio::io::AsyncWrite;
|
| 19 |
+
use tokio::process::Child;
|
| 20 |
+
use tokio::sync::mpsc;
|
| 21 |
+
use tokio::sync::watch;
|
| 22 |
+
use tokio::time::timeout;
|
| 23 |
+
#[cfg(test)]
|
| 24 |
+
use tokio_tungstenite::WebSocketStream;
|
| 25 |
+
use tokio_tungstenite::tungstenite::Message;
|
| 26 |
+
use tracing::debug;
|
| 27 |
+
use tracing::warn;
|
| 28 |
+
|
| 29 |
+
use tokio::io::AsyncBufReadExt;
|
| 30 |
+
use tokio::io::AsyncReadExt;
|
| 31 |
+
use tokio::io::AsyncWriteExt;
|
| 32 |
+
use tokio::io::BufReader;
|
| 33 |
+
use tokio::io::BufWriter;
|
| 34 |
+
|
| 35 |
+
pub(crate) const CHANNEL_CAPACITY: usize = 128;
|
| 36 |
+
// Match the existing serialized JSON-RPC message ceiling used by Noise and
|
| 37 |
+
// WebSocket transports so stdio has the same per-message bound.
|
| 38 |
+
const MAX_STDIO_JSONRPC_MESSAGE_LEN: usize = 64 * 1024 * 1024;
|
| 39 |
+
const STDIO_TERMINATION_GRACE_PERIOD: Duration = Duration::from_secs(2);
|
| 40 |
+
#[cfg(test)]
|
| 41 |
+
pub(crate) const WEBSOCKET_KEEPALIVE_INTERVAL: Duration = Duration::from_millis(25);
|
| 42 |
+
#[cfg(not(test))]
|
| 43 |
+
pub(crate) const WEBSOCKET_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(30);
|
| 44 |
+
|
| 45 |
+
#[derive(Debug)]
|
| 46 |
+
pub(crate) enum JsonRpcConnectionEvent {
|
| 47 |
+
Message(JSONRPCMessage),
|
| 48 |
+
QueuedRequest {
|
| 49 |
+
request: JSONRPCRequest,
|
| 50 |
+
request_span: tracing::Span,
|
| 51 |
+
queued_at: Instant,
|
| 52 |
+
},
|
| 53 |
+
MalformedMessage {
|
| 54 |
+
reason: String,
|
| 55 |
+
},
|
| 56 |
+
Disconnected {
|
| 57 |
+
reason: Option<String>,
|
| 58 |
+
},
|
| 59 |
+
}
|
| 60 |
+
|
| 61 |
+
impl JsonRpcConnectionEvent {
|
| 62 |
+
pub(crate) fn message(message: JSONRPCMessage) -> Self {
|
| 63 |
+
let JSONRPCMessage::Request(request) = message else {
|
| 64 |
+
return Self::Message(message);
|
| 65 |
+
};
|
| 66 |
+
|
| 67 |
+
let queued_at = Instant::now();
|
| 68 |
+
let request_span = tracing::info_span!(
|
| 69 |
+
"codex.exec_server.request",
|
| 70 |
+
otel.kind = "server",
|
| 71 |
+
otel.name = "unknown",
|
| 72 |
+
method = request.method.as_str(),
|
| 73 |
+
result = tracing::field::Empty,
|
| 74 |
+
);
|
| 75 |
+
if let Some(trace) = &request.trace
|
| 76 |
+
&& !codex_otel::set_parent_from_w3c_trace_context(&request_span, trace)
|
| 77 |
+
{
|
| 78 |
+
warn!(
|
| 79 |
+
method = request.method.as_str(),
|
| 80 |
+
"ignoring invalid inbound exec-server trace carrier"
|
| 81 |
+
);
|
| 82 |
+
}
|
| 83 |
+
|
| 84 |
+
Self::QueuedRequest {
|
| 85 |
+
request,
|
| 86 |
+
request_span,
|
| 87 |
+
queued_at,
|
| 88 |
+
}
|
| 89 |
+
}
|
| 90 |
+
}
|
| 91 |
+
|
| 92 |
+
#[derive(Clone)]
|
| 93 |
+
pub(crate) enum JsonRpcTransport {
|
| 94 |
+
// Plain means no child process; transport bytes may still be encrypted.
|
| 95 |
+
Plain,
|
| 96 |
+
Stdio { transport: StdioTransport },
|
| 97 |
+
}
|
| 98 |
+
|
| 99 |
+
impl JsonRpcTransport {
|
| 100 |
+
fn from_child_process(child_process: Child) -> Self {
|
| 101 |
+
Self::Stdio {
|
| 102 |
+
transport: StdioTransport::spawn(child_process),
|
| 103 |
+
}
|
| 104 |
+
}
|
| 105 |
+
|
| 106 |
+
pub(crate) fn terminate(&self) {
|
| 107 |
+
match self {
|
| 108 |
+
Self::Plain => {}
|
| 109 |
+
Self::Stdio { transport } => transport.terminate(),
|
| 110 |
+
}
|
| 111 |
+
}
|
| 112 |
+
}
|
| 113 |
+
|
| 114 |
+
#[derive(Clone)]
|
| 115 |
+
pub(crate) struct StdioTransport {
|
| 116 |
+
handle: Arc<StdioTransportHandle>,
|
| 117 |
+
}
|
| 118 |
+
|
| 119 |
+
struct StdioTransportHandle {
|
| 120 |
+
terminate_tx: watch::Sender<bool>,
|
| 121 |
+
terminate_requested: AtomicBool,
|
| 122 |
+
}
|
| 123 |
+
|
| 124 |
+
impl StdioTransport {
|
| 125 |
+
fn spawn(child_process: Child) -> Self {
|
| 126 |
+
let (terminate_tx, terminate_rx) = watch::channel(false);
|
| 127 |
+
let handle = Arc::new(StdioTransportHandle {
|
| 128 |
+
terminate_tx,
|
| 129 |
+
terminate_requested: AtomicBool::new(false),
|
| 130 |
+
});
|
| 131 |
+
spawn_stdio_child_supervisor(child_process, terminate_rx);
|
| 132 |
+
Self { handle }
|
| 133 |
+
}
|
| 134 |
+
|
| 135 |
+
fn terminate(&self) {
|
| 136 |
+
self.handle.terminate();
|
| 137 |
+
}
|
| 138 |
+
}
|
| 139 |
+
|
| 140 |
+
impl StdioTransportHandle {
|
| 141 |
+
fn terminate(&self) {
|
| 142 |
+
if !self.terminate_requested.swap(true, Ordering::AcqRel) {
|
| 143 |
+
let _ = self.terminate_tx.send(true);
|
| 144 |
+
}
|
| 145 |
+
}
|
| 146 |
+
}
|
| 147 |
+
|
| 148 |
+
impl Drop for StdioTransportHandle {
|
| 149 |
+
fn drop(&mut self) {
|
| 150 |
+
self.terminate();
|
| 151 |
+
}
|
| 152 |
+
}
|
| 153 |
+
|
| 154 |
+
fn spawn_stdio_child_supervisor(mut child_process: Child, mut terminate_rx: watch::Receiver<bool>) {
|
| 155 |
+
let process_group_id = child_process.id();
|
| 156 |
+
tokio::spawn(async move {
|
| 157 |
+
tokio::select! {
|
| 158 |
+
result = child_process.wait() => {
|
| 159 |
+
log_stdio_child_wait_result(result);
|
| 160 |
+
kill_process_tree(&mut child_process, process_group_id);
|
| 161 |
+
}
|
| 162 |
+
() = wait_for_stdio_termination(&mut terminate_rx) => {
|
| 163 |
+
terminate_stdio_child(&mut child_process, process_group_id).await;
|
| 164 |
+
}
|
| 165 |
+
}
|
| 166 |
+
});
|
| 167 |
+
}
|
| 168 |
+
|
| 169 |
+
async fn wait_for_stdio_termination(terminate_rx: &mut watch::Receiver<bool>) {
|
| 170 |
+
loop {
|
| 171 |
+
if *terminate_rx.borrow() {
|
| 172 |
+
return;
|
| 173 |
+
}
|
| 174 |
+
if terminate_rx.changed().await.is_err() {
|
| 175 |
+
return;
|
| 176 |
+
}
|
| 177 |
+
}
|
| 178 |
+
}
|
| 179 |
+
|
| 180 |
+
async fn terminate_stdio_child(child_process: &mut Child, process_group_id: Option<u32>) {
|
| 181 |
+
terminate_process_tree(child_process, process_group_id);
|
| 182 |
+
match timeout(STDIO_TERMINATION_GRACE_PERIOD, child_process.wait()).await {
|
| 183 |
+
Ok(result) => {
|
| 184 |
+
log_stdio_child_wait_result(result);
|
| 185 |
+
}
|
| 186 |
+
Err(_) => {
|
| 187 |
+
kill_process_tree(child_process, process_group_id);
|
| 188 |
+
log_stdio_child_wait_result(child_process.wait().await);
|
| 189 |
+
}
|
| 190 |
+
}
|
| 191 |
+
}
|
| 192 |
+
|
| 193 |
+
fn terminate_process_tree(child_process: &mut Child, process_group_id: Option<u32>) {
|
| 194 |
+
let Some(process_group_id) = process_group_id else {
|
| 195 |
+
kill_direct_child(child_process, "terminate");
|
| 196 |
+
return;
|
| 197 |
+
};
|
| 198 |
+
|
| 199 |
+
#[cfg(unix)]
|
| 200 |
+
if let Err(err) = codex_utils_pty::process_group::terminate_process_group(process_group_id) {
|
| 201 |
+
warn!("failed to terminate exec-server stdio process group {process_group_id}: {err}");
|
| 202 |
+
kill_direct_child(child_process, "terminate");
|
| 203 |
+
}
|
| 204 |
+
|
| 205 |
+
#[cfg(windows)]
|
| 206 |
+
if !kill_windows_process_tree(process_group_id) {
|
| 207 |
+
kill_direct_child(child_process, "terminate");
|
| 208 |
+
}
|
| 209 |
+
|
| 210 |
+
#[cfg(not(any(unix, windows)))]
|
| 211 |
+
{
|
| 212 |
+
let _ = process_group_id;
|
| 213 |
+
kill_direct_child(child_process, "terminate");
|
| 214 |
+
}
|
| 215 |
+
}
|
| 216 |
+
|
| 217 |
+
fn kill_process_tree(child_process: &mut Child, process_group_id: Option<u32>) {
|
| 218 |
+
let Some(process_group_id) = process_group_id else {
|
| 219 |
+
kill_direct_child(child_process, "kill");
|
| 220 |
+
return;
|
| 221 |
+
};
|
| 222 |
+
|
| 223 |
+
#[cfg(unix)]
|
| 224 |
+
if let Err(err) = codex_utils_pty::process_group::kill_process_group(process_group_id) {
|
| 225 |
+
warn!("failed to kill exec-server stdio process group {process_group_id}: {err}");
|
| 226 |
+
}
|
| 227 |
+
|
| 228 |
+
#[cfg(windows)]
|
| 229 |
+
if !kill_windows_process_tree(process_group_id) {
|
| 230 |
+
kill_direct_child(child_process, "kill");
|
| 231 |
+
}
|
| 232 |
+
|
| 233 |
+
#[cfg(not(any(unix, windows)))]
|
| 234 |
+
{
|
| 235 |
+
let _ = process_group_id;
|
| 236 |
+
kill_direct_child(child_process, "kill");
|
| 237 |
+
}
|
| 238 |
+
}
|
| 239 |
+
|
| 240 |
+
fn kill_direct_child(child_process: &mut Child, action: &str) {
|
| 241 |
+
if let Err(err) = child_process.start_kill() {
|
| 242 |
+
debug!("failed to {action} exec-server stdio child: {err}");
|
| 243 |
+
}
|
| 244 |
+
}
|
| 245 |
+
|
| 246 |
+
#[cfg(windows)]
|
| 247 |
+
fn kill_windows_process_tree(pid: u32) -> bool {
|
| 248 |
+
let pid = pid.to_string();
|
| 249 |
+
match std::process::Command::new("taskkill")
|
| 250 |
+
.args(["/PID", pid.as_str(), "/T", "/F"])
|
| 251 |
+
.stdin(Stdio::null())
|
| 252 |
+
.stdout(Stdio::null())
|
| 253 |
+
.stderr(Stdio::null())
|
| 254 |
+
.status()
|
| 255 |
+
{
|
| 256 |
+
Ok(status) => status.success(),
|
| 257 |
+
Err(err) => {
|
| 258 |
+
warn!("failed to run taskkill for exec-server stdio process tree {pid}: {err}");
|
| 259 |
+
false
|
| 260 |
+
}
|
| 261 |
+
}
|
| 262 |
+
}
|
| 263 |
+
|
| 264 |
+
fn log_stdio_child_wait_result(result: std::io::Result<std::process::ExitStatus>) {
|
| 265 |
+
if let Err(err) = result {
|
| 266 |
+
debug!("failed to wait for exec-server stdio child: {err}");
|
| 267 |
+
}
|
| 268 |
+
}
|
| 269 |
+
|
| 270 |
+
pub(crate) struct JsonRpcConnection {
|
| 271 |
+
pub(crate) outgoing_tx: mpsc::Sender<JSONRPCMessage>,
|
| 272 |
+
pub(crate) incoming_rx: mpsc::Receiver<JsonRpcConnectionEvent>,
|
| 273 |
+
pub(crate) disconnected_rx: watch::Receiver<bool>,
|
| 274 |
+
pub(crate) task_handles: Vec<tokio::task::JoinHandle<()>>,
|
| 275 |
+
pub(crate) transport: JsonRpcTransport,
|
| 276 |
+
}
|
| 277 |
+
|
| 278 |
+
impl JsonRpcConnection {
|
| 279 |
+
pub(crate) fn from_stdio<R, W>(reader: R, writer: W, connection_label: String) -> Self
|
| 280 |
+
where
|
| 281 |
+
R: AsyncRead + Unpin + Send + 'static,
|
| 282 |
+
W: AsyncWrite + Unpin + Send + 'static,
|
| 283 |
+
{
|
| 284 |
+
Self::from_stdio_with_max_message_len(
|
| 285 |
+
reader,
|
| 286 |
+
writer,
|
| 287 |
+
connection_label,
|
| 288 |
+
MAX_STDIO_JSONRPC_MESSAGE_LEN,
|
| 289 |
+
)
|
| 290 |
+
}
|
| 291 |
+
|
| 292 |
+
fn from_stdio_with_max_message_len<R, W>(
|
| 293 |
+
reader: R,
|
| 294 |
+
writer: W,
|
| 295 |
+
connection_label: String,
|
| 296 |
+
max_message_len: usize,
|
| 297 |
+
) -> Self
|
| 298 |
+
where
|
| 299 |
+
R: AsyncRead + Unpin + Send + 'static,
|
| 300 |
+
W: AsyncWrite + Unpin + Send + 'static,
|
| 301 |
+
{
|
| 302 |
+
let (outgoing_tx, mut outgoing_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 303 |
+
let (incoming_tx, incoming_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 304 |
+
let (disconnected_tx, disconnected_rx) = watch::channel(false);
|
| 305 |
+
|
| 306 |
+
let reader_label = connection_label.clone();
|
| 307 |
+
let incoming_tx_for_reader = incoming_tx.clone();
|
| 308 |
+
let disconnected_tx_for_reader = disconnected_tx.clone();
|
| 309 |
+
// Read one byte past the payload limit so an unterminated oversized
|
| 310 |
+
// message fails promptly. A trailing CR gets one more byte of lookahead
|
| 311 |
+
// because it may be the first half of a valid CRLF terminator.
|
| 312 |
+
let read_limit = u64::try_from(max_message_len.saturating_add(1)).unwrap_or(u64::MAX);
|
| 313 |
+
let reader_task = tokio::spawn(async move {
|
| 314 |
+
let mut reader = BufReader::new(reader);
|
| 315 |
+
let mut line = String::new();
|
| 316 |
+
loop {
|
| 317 |
+
line.clear();
|
| 318 |
+
let read_result = (&mut reader).take(read_limit).read_line(&mut line).await;
|
| 319 |
+
match read_result {
|
| 320 |
+
Ok(0) => {
|
| 321 |
+
send_disconnected(
|
| 322 |
+
&incoming_tx_for_reader,
|
| 323 |
+
&disconnected_tx_for_reader,
|
| 324 |
+
/*reason*/ None,
|
| 325 |
+
)
|
| 326 |
+
.await;
|
| 327 |
+
break;
|
| 328 |
+
}
|
| 329 |
+
Ok(_) => {
|
| 330 |
+
if line.ends_with('\n') {
|
| 331 |
+
line.pop();
|
| 332 |
+
if line.ends_with('\r') {
|
| 333 |
+
line.pop();
|
| 334 |
+
}
|
| 335 |
+
} else if line.len() > max_message_len && line.ends_with('\r') {
|
| 336 |
+
match reader.read_u8().await {
|
| 337 |
+
Ok(b'\n') => {
|
| 338 |
+
line.pop();
|
| 339 |
+
}
|
| 340 |
+
Ok(_) => {}
|
| 341 |
+
Err(err) if err.kind() == std::io::ErrorKind::UnexpectedEof => {}
|
| 342 |
+
Err(err) => {
|
| 343 |
+
send_disconnected(
|
| 344 |
+
&incoming_tx_for_reader,
|
| 345 |
+
&disconnected_tx_for_reader,
|
| 346 |
+
Some(format!(
|
| 347 |
+
"failed to read JSON-RPC message from {reader_label}: {err}"
|
| 348 |
+
)),
|
| 349 |
+
)
|
| 350 |
+
.await;
|
| 351 |
+
break;
|
| 352 |
+
}
|
| 353 |
+
}
|
| 354 |
+
}
|
| 355 |
+
if line.len() > max_message_len {
|
| 356 |
+
send_disconnected(
|
| 357 |
+
&incoming_tx_for_reader,
|
| 358 |
+
&disconnected_tx_for_reader,
|
| 359 |
+
Some(format!(
|
| 360 |
+
"JSON-RPC message from {reader_label} exceeds maximum length of {max_message_len} bytes"
|
| 361 |
+
)),
|
| 362 |
+
)
|
| 363 |
+
.await;
|
| 364 |
+
break;
|
| 365 |
+
}
|
| 366 |
+
if line.trim().is_empty() {
|
| 367 |
+
continue;
|
| 368 |
+
}
|
| 369 |
+
match serde_json::from_str::<JSONRPCMessage>(&line) {
|
| 370 |
+
Ok(message) => {
|
| 371 |
+
if incoming_tx_for_reader
|
| 372 |
+
.send(JsonRpcConnectionEvent::message(message))
|
| 373 |
+
.await
|
| 374 |
+
.is_err()
|
| 375 |
+
{
|
| 376 |
+
break;
|
| 377 |
+
}
|
| 378 |
+
}
|
| 379 |
+
Err(err) => {
|
| 380 |
+
send_malformed_message(
|
| 381 |
+
&incoming_tx_for_reader,
|
| 382 |
+
Some(format!(
|
| 383 |
+
"failed to parse JSON-RPC message from {reader_label}: {err}"
|
| 384 |
+
)),
|
| 385 |
+
)
|
| 386 |
+
.await;
|
| 387 |
+
}
|
| 388 |
+
}
|
| 389 |
+
}
|
| 390 |
+
Err(err) => {
|
| 391 |
+
send_disconnected(
|
| 392 |
+
&incoming_tx_for_reader,
|
| 393 |
+
&disconnected_tx_for_reader,
|
| 394 |
+
Some(format!(
|
| 395 |
+
"failed to read JSON-RPC message from {reader_label}: {err}"
|
| 396 |
+
)),
|
| 397 |
+
)
|
| 398 |
+
.await;
|
| 399 |
+
break;
|
| 400 |
+
}
|
| 401 |
+
}
|
| 402 |
+
}
|
| 403 |
+
});
|
| 404 |
+
|
| 405 |
+
let writer_task = tokio::spawn(async move {
|
| 406 |
+
let mut writer = BufWriter::new(writer);
|
| 407 |
+
while let Some(message) = outgoing_rx.recv().await {
|
| 408 |
+
if let Err(err) = write_jsonrpc_line_message(&mut writer, &message).await {
|
| 409 |
+
send_disconnected(
|
| 410 |
+
&incoming_tx,
|
| 411 |
+
&disconnected_tx,
|
| 412 |
+
Some(format!(
|
| 413 |
+
"failed to write JSON-RPC message to {connection_label}: {err}"
|
| 414 |
+
)),
|
| 415 |
+
)
|
| 416 |
+
.await;
|
| 417 |
+
break;
|
| 418 |
+
}
|
| 419 |
+
}
|
| 420 |
+
});
|
| 421 |
+
|
| 422 |
+
Self {
|
| 423 |
+
outgoing_tx,
|
| 424 |
+
incoming_rx,
|
| 425 |
+
disconnected_rx,
|
| 426 |
+
task_handles: vec![reader_task, writer_task],
|
| 427 |
+
transport: JsonRpcTransport::Plain,
|
| 428 |
+
}
|
| 429 |
+
}
|
| 430 |
+
|
| 431 |
+
pub(crate) fn from_websocket<T, E>(stream: T, connection_label: String) -> Self
|
| 432 |
+
where
|
| 433 |
+
T: Sink<Message, Error = E> + Stream<Item = Result<Message, E>> + Unpin + Send + 'static,
|
| 434 |
+
E: std::fmt::Display + Send + 'static,
|
| 435 |
+
{
|
| 436 |
+
Self::from_websocket_stream(stream, connection_label, /*ping_interval*/ None)
|
| 437 |
+
}
|
| 438 |
+
|
| 439 |
+
pub(crate) fn from_axum_websocket(stream: AxumWebSocket, connection_label: String) -> Self {
|
| 440 |
+
Self::from_websocket_stream(stream, connection_label, Some(WEBSOCKET_KEEPALIVE_INTERVAL))
|
| 441 |
+
}
|
| 442 |
+
|
| 443 |
+
fn from_websocket_stream<T, M, E>(
|
| 444 |
+
mut websocket: T,
|
| 445 |
+
connection_label: String,
|
| 446 |
+
ping_interval: Option<Duration>,
|
| 447 |
+
) -> Self
|
| 448 |
+
where
|
| 449 |
+
T: Sink<M, Error = E> + Stream<Item = Result<M, E>> + Unpin + Send + 'static,
|
| 450 |
+
M: JsonRpcWebSocketMessage,
|
| 451 |
+
E: std::fmt::Display + Send + 'static,
|
| 452 |
+
{
|
| 453 |
+
let (outgoing_tx, mut outgoing_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 454 |
+
let (incoming_tx, incoming_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 455 |
+
let (disconnected_tx, disconnected_rx) = watch::channel(false);
|
| 456 |
+
|
| 457 |
+
let websocket_task = tokio::spawn(async move {
|
| 458 |
+
let mut ping_interval = ping_interval.map(|ping_interval| {
|
| 459 |
+
let mut interval = tokio::time::interval_at(
|
| 460 |
+
tokio::time::Instant::now() + ping_interval,
|
| 461 |
+
ping_interval,
|
| 462 |
+
);
|
| 463 |
+
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
|
| 464 |
+
interval
|
| 465 |
+
});
|
| 466 |
+
|
| 467 |
+
loop {
|
| 468 |
+
tokio::select! {
|
| 469 |
+
maybe_message = outgoing_rx.recv() => {
|
| 470 |
+
let Some(message) = maybe_message else {
|
| 471 |
+
break;
|
| 472 |
+
};
|
| 473 |
+
if let Err(reason) = send_websocket_jsonrpc_message(
|
| 474 |
+
&mut websocket,
|
| 475 |
+
&connection_label,
|
| 476 |
+
&message,
|
| 477 |
+
)
|
| 478 |
+
.await
|
| 479 |
+
{
|
| 480 |
+
send_disconnected(&incoming_tx, &disconnected_tx, Some(reason)).await;
|
| 481 |
+
break;
|
| 482 |
+
}
|
| 483 |
+
}
|
| 484 |
+
_ = async {
|
| 485 |
+
match ping_interval.as_mut() {
|
| 486 |
+
Some(interval) => interval.tick().await,
|
| 487 |
+
None => std::future::pending().await,
|
| 488 |
+
}
|
| 489 |
+
} => {
|
| 490 |
+
if let Err(err) = websocket.send(M::ping()).await {
|
| 491 |
+
send_disconnected(
|
| 492 |
+
&incoming_tx,
|
| 493 |
+
&disconnected_tx,
|
| 494 |
+
Some(format!(
|
| 495 |
+
"failed to write websocket ping to {connection_label}: {err}"
|
| 496 |
+
)),
|
| 497 |
+
)
|
| 498 |
+
.await;
|
| 499 |
+
break;
|
| 500 |
+
}
|
| 501 |
+
}
|
| 502 |
+
incoming_message = websocket.next() => {
|
| 503 |
+
match incoming_message {
|
| 504 |
+
Some(Ok(message)) => match message.parse_jsonrpc_frame() {
|
| 505 |
+
Ok(JsonRpcWebSocketFrame::Message(message)) => {
|
| 506 |
+
if incoming_tx
|
| 507 |
+
.send(JsonRpcConnectionEvent::message(message))
|
| 508 |
+
.await
|
| 509 |
+
.is_err()
|
| 510 |
+
{
|
| 511 |
+
break;
|
| 512 |
+
}
|
| 513 |
+
}
|
| 514 |
+
Ok(JsonRpcWebSocketFrame::Close) => {
|
| 515 |
+
send_disconnected(
|
| 516 |
+
&incoming_tx,
|
| 517 |
+
&disconnected_tx,
|
| 518 |
+
/*reason*/ None,
|
| 519 |
+
)
|
| 520 |
+
.await;
|
| 521 |
+
break;
|
| 522 |
+
}
|
| 523 |
+
Ok(JsonRpcWebSocketFrame::Ignore) => {}
|
| 524 |
+
Err(err) => {
|
| 525 |
+
send_malformed_message(
|
| 526 |
+
&incoming_tx,
|
| 527 |
+
Some(format!(
|
| 528 |
+
"failed to parse websocket JSON-RPC message from {connection_label}: {err}"
|
| 529 |
+
)),
|
| 530 |
+
)
|
| 531 |
+
.await;
|
| 532 |
+
}
|
| 533 |
+
},
|
| 534 |
+
Some(Err(err)) => {
|
| 535 |
+
send_disconnected(
|
| 536 |
+
&incoming_tx,
|
| 537 |
+
&disconnected_tx,
|
| 538 |
+
Some(format!(
|
| 539 |
+
"failed to read websocket JSON-RPC message from {connection_label}: {err}"
|
| 540 |
+
)),
|
| 541 |
+
)
|
| 542 |
+
.await;
|
| 543 |
+
break;
|
| 544 |
+
}
|
| 545 |
+
None => {
|
| 546 |
+
send_disconnected(
|
| 547 |
+
&incoming_tx,
|
| 548 |
+
&disconnected_tx,
|
| 549 |
+
/*reason*/ None,
|
| 550 |
+
)
|
| 551 |
+
.await;
|
| 552 |
+
break;
|
| 553 |
+
}
|
| 554 |
+
}
|
| 555 |
+
}
|
| 556 |
+
}
|
| 557 |
+
}
|
| 558 |
+
});
|
| 559 |
+
|
| 560 |
+
Self {
|
| 561 |
+
outgoing_tx,
|
| 562 |
+
incoming_rx,
|
| 563 |
+
disconnected_rx,
|
| 564 |
+
task_handles: vec![websocket_task],
|
| 565 |
+
transport: JsonRpcTransport::Plain,
|
| 566 |
+
}
|
| 567 |
+
}
|
| 568 |
+
|
| 569 |
+
pub(crate) fn with_child_process(mut self, child_process: Child) -> Self {
|
| 570 |
+
self.transport = JsonRpcTransport::from_child_process(child_process);
|
| 571 |
+
self
|
| 572 |
+
}
|
| 573 |
+
}
|
| 574 |
+
|
| 575 |
+
enum JsonRpcWebSocketFrame {
|
| 576 |
+
Message(JSONRPCMessage),
|
| 577 |
+
Close,
|
| 578 |
+
Ignore,
|
| 579 |
+
}
|
| 580 |
+
|
| 581 |
+
trait JsonRpcWebSocketMessage: Send + 'static {
|
| 582 |
+
fn parse_jsonrpc_frame(self) -> Result<JsonRpcWebSocketFrame, serde_json::Error>;
|
| 583 |
+
fn from_text(text: String) -> Self;
|
| 584 |
+
fn ping() -> Self;
|
| 585 |
+
}
|
| 586 |
+
|
| 587 |
+
impl JsonRpcWebSocketMessage for Message {
|
| 588 |
+
fn parse_jsonrpc_frame(self) -> Result<JsonRpcWebSocketFrame, serde_json::Error> {
|
| 589 |
+
match self {
|
| 590 |
+
Message::Text(text) => {
|
| 591 |
+
serde_json::from_str(text.as_ref()).map(JsonRpcWebSocketFrame::Message)
|
| 592 |
+
}
|
| 593 |
+
Message::Binary(bytes) => {
|
| 594 |
+
serde_json::from_slice(bytes.as_ref()).map(JsonRpcWebSocketFrame::Message)
|
| 595 |
+
}
|
| 596 |
+
Message::Close(_) => Ok(JsonRpcWebSocketFrame::Close),
|
| 597 |
+
Message::Ping(_) | Message::Pong(_) | Message::Frame(_) => {
|
| 598 |
+
Ok(JsonRpcWebSocketFrame::Ignore)
|
| 599 |
+
}
|
| 600 |
+
}
|
| 601 |
+
}
|
| 602 |
+
|
| 603 |
+
fn from_text(text: String) -> Self {
|
| 604 |
+
Self::Text(text.into())
|
| 605 |
+
}
|
| 606 |
+
|
| 607 |
+
fn ping() -> Self {
|
| 608 |
+
Self::Ping(Vec::new().into())
|
| 609 |
+
}
|
| 610 |
+
}
|
| 611 |
+
|
| 612 |
+
impl JsonRpcWebSocketMessage for AxumWebSocketMessage {
|
| 613 |
+
fn parse_jsonrpc_frame(self) -> Result<JsonRpcWebSocketFrame, serde_json::Error> {
|
| 614 |
+
match self {
|
| 615 |
+
AxumWebSocketMessage::Text(text) => {
|
| 616 |
+
serde_json::from_str(text.as_ref()).map(JsonRpcWebSocketFrame::Message)
|
| 617 |
+
}
|
| 618 |
+
AxumWebSocketMessage::Binary(bytes) => {
|
| 619 |
+
serde_json::from_slice(bytes.as_ref()).map(JsonRpcWebSocketFrame::Message)
|
| 620 |
+
}
|
| 621 |
+
AxumWebSocketMessage::Close(_) => Ok(JsonRpcWebSocketFrame::Close),
|
| 622 |
+
AxumWebSocketMessage::Ping(_) | AxumWebSocketMessage::Pong(_) => {
|
| 623 |
+
Ok(JsonRpcWebSocketFrame::Ignore)
|
| 624 |
+
}
|
| 625 |
+
}
|
| 626 |
+
}
|
| 627 |
+
|
| 628 |
+
fn from_text(text: String) -> Self {
|
| 629 |
+
Self::Text(text.into())
|
| 630 |
+
}
|
| 631 |
+
|
| 632 |
+
fn ping() -> Self {
|
| 633 |
+
Self::Ping(Vec::new().into())
|
| 634 |
+
}
|
| 635 |
+
}
|
| 636 |
+
|
| 637 |
+
async fn send_disconnected(
|
| 638 |
+
incoming_tx: &mpsc::Sender<JsonRpcConnectionEvent>,
|
| 639 |
+
disconnected_tx: &watch::Sender<bool>,
|
| 640 |
+
reason: Option<String>,
|
| 641 |
+
) {
|
| 642 |
+
let _ = disconnected_tx.send(true);
|
| 643 |
+
let _ = incoming_tx
|
| 644 |
+
.send(JsonRpcConnectionEvent::Disconnected { reason })
|
| 645 |
+
.await;
|
| 646 |
+
}
|
| 647 |
+
|
| 648 |
+
async fn send_malformed_message(
|
| 649 |
+
incoming_tx: &mpsc::Sender<JsonRpcConnectionEvent>,
|
| 650 |
+
reason: Option<String>,
|
| 651 |
+
) {
|
| 652 |
+
let _ = incoming_tx
|
| 653 |
+
.send(JsonRpcConnectionEvent::MalformedMessage {
|
| 654 |
+
reason: reason.unwrap_or_else(|| "malformed JSON-RPC message".to_string()),
|
| 655 |
+
})
|
| 656 |
+
.await;
|
| 657 |
+
}
|
| 658 |
+
|
| 659 |
+
async fn write_jsonrpc_line_message<W>(
|
| 660 |
+
writer: &mut BufWriter<W>,
|
| 661 |
+
message: &JSONRPCMessage,
|
| 662 |
+
) -> std::io::Result<()>
|
| 663 |
+
where
|
| 664 |
+
W: AsyncWrite + Unpin,
|
| 665 |
+
{
|
| 666 |
+
let encoded =
|
| 667 |
+
serialize_jsonrpc_message(message).map_err(|err| std::io::Error::other(err.to_string()))?;
|
| 668 |
+
writer.write_all(encoded.as_bytes()).await?;
|
| 669 |
+
writer.write_all(b"\n").await?;
|
| 670 |
+
writer.flush().await
|
| 671 |
+
}
|
| 672 |
+
|
| 673 |
+
async fn send_websocket_jsonrpc_message<W, M, E>(
|
| 674 |
+
websocket_writer: &mut W,
|
| 675 |
+
connection_label: &str,
|
| 676 |
+
message: &JSONRPCMessage,
|
| 677 |
+
) -> Result<(), String>
|
| 678 |
+
where
|
| 679 |
+
W: Sink<M, Error = E> + Unpin,
|
| 680 |
+
M: JsonRpcWebSocketMessage,
|
| 681 |
+
E: std::fmt::Display,
|
| 682 |
+
{
|
| 683 |
+
match serialize_jsonrpc_message(message) {
|
| 684 |
+
Ok(encoded) => websocket_writer
|
| 685 |
+
.send(M::from_text(encoded))
|
| 686 |
+
.await
|
| 687 |
+
.map_err(|err| {
|
| 688 |
+
format!("failed to write websocket JSON-RPC message to {connection_label}: {err}")
|
| 689 |
+
}),
|
| 690 |
+
Err(err) => Err(format!(
|
| 691 |
+
"failed to serialize JSON-RPC message for {connection_label}: {err}"
|
| 692 |
+
)),
|
| 693 |
+
}
|
| 694 |
+
}
|
| 695 |
+
|
| 696 |
+
fn serialize_jsonrpc_message(message: &JSONRPCMessage) -> Result<String, serde_json::Error> {
|
| 697 |
+
serde_json::to_string(message)
|
| 698 |
+
}
|
| 699 |
+
|
| 700 |
+
#[cfg(test)]
|
| 701 |
+
mod tests {
|
| 702 |
+
use std::pin::Pin;
|
| 703 |
+
use std::sync::Arc;
|
| 704 |
+
use std::sync::atomic::AtomicBool;
|
| 705 |
+
use std::sync::atomic::Ordering;
|
| 706 |
+
use std::task::Context;
|
| 707 |
+
use std::task::Poll;
|
| 708 |
+
|
| 709 |
+
use codex_exec_server_protocol::JSONRPCRequest;
|
| 710 |
+
use codex_exec_server_protocol::RequestId;
|
| 711 |
+
use futures::channel::mpsc as futures_mpsc;
|
| 712 |
+
use futures::task::AtomicWaker;
|
| 713 |
+
use pretty_assertions::assert_eq;
|
| 714 |
+
use tokio::net::TcpListener;
|
| 715 |
+
use tokio::time::timeout;
|
| 716 |
+
use tokio_tungstenite::accept_async;
|
| 717 |
+
use tokio_tungstenite::connect_async;
|
| 718 |
+
|
| 719 |
+
use super::*;
|
| 720 |
+
|
| 721 |
+
#[tokio::test]
|
| 722 |
+
async fn stdio_connection_accepts_message_at_size_limit() -> anyhow::Result<()> {
|
| 723 |
+
let message = test_jsonrpc_message();
|
| 724 |
+
let encoded = serde_json::to_string(&message)?;
|
| 725 |
+
let max_message_len = encoded.len();
|
| 726 |
+
|
| 727 |
+
for line_ending in [b"\n".as_slice(), b"\r\n".as_slice()] {
|
| 728 |
+
let (reader, mut peer) =
|
| 729 |
+
tokio::io::duplex(max_message_len.saturating_add(line_ending.len()));
|
| 730 |
+
let mut connection = JsonRpcConnection::from_stdio_with_max_message_len(
|
| 731 |
+
reader,
|
| 732 |
+
tokio::io::sink(),
|
| 733 |
+
"test stdio peer".to_string(),
|
| 734 |
+
max_message_len,
|
| 735 |
+
);
|
| 736 |
+
|
| 737 |
+
peer.write_all(encoded.as_bytes()).await?;
|
| 738 |
+
peer.write_all(line_ending).await?;
|
| 739 |
+
let event = timeout(Duration::from_secs(1), connection.incoming_rx.recv())
|
| 740 |
+
.await?
|
| 741 |
+
.expect("stdio connection should report the message");
|
| 742 |
+
match event {
|
| 743 |
+
JsonRpcConnectionEvent::QueuedRequest { request, .. } => {
|
| 744 |
+
assert_eq!(JSONRPCMessage::Request(request), message)
|
| 745 |
+
}
|
| 746 |
+
event => anyhow::bail!("expected JSON-RPC message, got {event:?}"),
|
| 747 |
+
}
|
| 748 |
+
|
| 749 |
+
drop(peer);
|
| 750 |
+
drop(connection);
|
| 751 |
+
}
|
| 752 |
+
|
| 753 |
+
Ok(())
|
| 754 |
+
}
|
| 755 |
+
|
| 756 |
+
#[tokio::test]
|
| 757 |
+
async fn stdio_connection_rejects_overlong_unterminated_message() -> anyhow::Result<()> {
|
| 758 |
+
let max_message_len: usize = 32;
|
| 759 |
+
let (reader, mut peer) = tokio::io::duplex(max_message_len.saturating_add(1));
|
| 760 |
+
let mut connection = JsonRpcConnection::from_stdio_with_max_message_len(
|
| 761 |
+
reader,
|
| 762 |
+
tokio::io::sink(),
|
| 763 |
+
"hostile stdio peer".to_string(),
|
| 764 |
+
max_message_len,
|
| 765 |
+
);
|
| 766 |
+
let overlong_message = vec![b'x'; max_message_len + 1];
|
| 767 |
+
|
| 768 |
+
peer.write_all(&overlong_message).await?;
|
| 769 |
+
let event = timeout(Duration::from_secs(1), connection.incoming_rx.recv())
|
| 770 |
+
.await?
|
| 771 |
+
.expect("stdio connection should report the framing violation");
|
| 772 |
+
match event {
|
| 773 |
+
JsonRpcConnectionEvent::Disconnected { reason } => assert_eq!(
|
| 774 |
+
reason,
|
| 775 |
+
Some(
|
| 776 |
+
"JSON-RPC message from hostile stdio peer exceeds maximum length of 32 bytes"
|
| 777 |
+
.to_string()
|
| 778 |
+
)
|
| 779 |
+
),
|
| 780 |
+
event => anyhow::bail!("expected stdio disconnect, got {event:?}"),
|
| 781 |
+
}
|
| 782 |
+
|
| 783 |
+
drop(peer);
|
| 784 |
+
drop(connection);
|
| 785 |
+
Ok(())
|
| 786 |
+
}
|
| 787 |
+
|
| 788 |
+
#[tokio::test]
|
| 789 |
+
async fn websocket_connection_sends_configured_ping() -> anyhow::Result<()> {
|
| 790 |
+
let (client_websocket, mut server_websocket) = websocket_pair().await?;
|
| 791 |
+
let connection = JsonRpcConnection::from_websocket_stream(
|
| 792 |
+
client_websocket,
|
| 793 |
+
"test".into(),
|
| 794 |
+
Some(WEBSOCKET_KEEPALIVE_INTERVAL),
|
| 795 |
+
);
|
| 796 |
+
|
| 797 |
+
let message = timeout(Duration::from_secs(1), server_websocket.next())
|
| 798 |
+
.await?
|
| 799 |
+
.expect("websocket should stay open")?;
|
| 800 |
+
assert!(matches!(message, Message::Ping(_)));
|
| 801 |
+
|
| 802 |
+
drop(connection);
|
| 803 |
+
Ok(())
|
| 804 |
+
}
|
| 805 |
+
|
| 806 |
+
#[tokio::test]
|
| 807 |
+
async fn websocket_connection_ignores_server_pong() -> anyhow::Result<()> {
|
| 808 |
+
let (client_websocket, mut server_websocket) = websocket_pair().await?;
|
| 809 |
+
let mut connection = JsonRpcConnection::from_websocket(client_websocket, "test".into());
|
| 810 |
+
|
| 811 |
+
server_websocket
|
| 812 |
+
.send(Message::Pong(b"check".to_vec().into()))
|
| 813 |
+
.await?;
|
| 814 |
+
assert!(
|
| 815 |
+
timeout(Duration::from_millis(50), connection.incoming_rx.recv())
|
| 816 |
+
.await
|
| 817 |
+
.is_err()
|
| 818 |
+
);
|
| 819 |
+
|
| 820 |
+
drop(connection);
|
| 821 |
+
Ok(())
|
| 822 |
+
}
|
| 823 |
+
|
| 824 |
+
#[tokio::test]
|
| 825 |
+
async fn websocket_connection_reports_server_close() -> anyhow::Result<()> {
|
| 826 |
+
let (client_websocket, mut server_websocket) = websocket_pair().await?;
|
| 827 |
+
let mut connection = JsonRpcConnection::from_websocket(client_websocket, "test".into());
|
| 828 |
+
|
| 829 |
+
server_websocket.close(None).await?;
|
| 830 |
+
assert!(matches!(
|
| 831 |
+
timeout(Duration::from_secs(1), connection.incoming_rx.recv()).await?,
|
| 832 |
+
Some(JsonRpcConnectionEvent::Disconnected { reason: None })
|
| 833 |
+
));
|
| 834 |
+
|
| 835 |
+
drop(connection);
|
| 836 |
+
Ok(())
|
| 837 |
+
}
|
| 838 |
+
|
| 839 |
+
#[tokio::test]
|
| 840 |
+
async fn websocket_connection_accepts_binary_jsonrpc_message() -> anyhow::Result<()> {
|
| 841 |
+
let (client_websocket, mut server_websocket) = websocket_pair().await?;
|
| 842 |
+
let mut connection = JsonRpcConnection::from_websocket(client_websocket, "test".into());
|
| 843 |
+
let message = JSONRPCMessage::Request(JSONRPCRequest {
|
| 844 |
+
id: RequestId::Integer(1),
|
| 845 |
+
method: "test".to_string(),
|
| 846 |
+
params: None,
|
| 847 |
+
trace: None,
|
| 848 |
+
});
|
| 849 |
+
|
| 850 |
+
server_websocket
|
| 851 |
+
.send(Message::Binary(serde_json::to_vec(&message)?.into()))
|
| 852 |
+
.await?;
|
| 853 |
+
let Some(JsonRpcConnectionEvent::QueuedRequest { request, .. }) =
|
| 854 |
+
timeout(Duration::from_secs(1), connection.incoming_rx.recv()).await?
|
| 855 |
+
else {
|
| 856 |
+
anyhow::bail!("expected a queued JSON-RPC request");
|
| 857 |
+
};
|
| 858 |
+
assert_eq!(JSONRPCMessage::Request(request), message);
|
| 859 |
+
|
| 860 |
+
drop(connection);
|
| 861 |
+
Ok(())
|
| 862 |
+
}
|
| 863 |
+
|
| 864 |
+
#[tokio::test]
|
| 865 |
+
async fn websocket_connection_keeps_outbound_message_while_send_is_backpressured()
|
| 866 |
+
-> anyhow::Result<()> {
|
| 867 |
+
let (websocket, control, mut outbound_rx) =
|
| 868 |
+
ControlledWebSocket::new(/*write_ready*/ false);
|
| 869 |
+
let mut connection = JsonRpcConnection::from_websocket_stream(
|
| 870 |
+
websocket,
|
| 871 |
+
"test".into(),
|
| 872 |
+
/*ping_interval*/ None,
|
| 873 |
+
);
|
| 874 |
+
let message = test_jsonrpc_message();
|
| 875 |
+
|
| 876 |
+
connection.outgoing_tx.send(message.clone()).await?;
|
| 877 |
+
control.wait_for_blocked_write().await?;
|
| 878 |
+
control.send_inbound(Message::Pong(b"check".to_vec().into()))?;
|
| 879 |
+
assert!(
|
| 880 |
+
timeout(Duration::from_millis(50), connection.incoming_rx.recv())
|
| 881 |
+
.await
|
| 882 |
+
.is_err()
|
| 883 |
+
);
|
| 884 |
+
|
| 885 |
+
control.set_write_ready();
|
| 886 |
+
assert!(matches!(
|
| 887 |
+
timeout(Duration::from_secs(1), outbound_rx.next()).await?,
|
| 888 |
+
Some(Message::Text(text)) if serde_json::from_str::<JSONRPCMessage>(&text)? == message
|
| 889 |
+
));
|
| 890 |
+
drop(connection);
|
| 891 |
+
Ok(())
|
| 892 |
+
}
|
| 893 |
+
|
| 894 |
+
async fn websocket_pair() -> anyhow::Result<(
|
| 895 |
+
WebSocketStream<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>>,
|
| 896 |
+
WebSocketStream<tokio::net::TcpStream>,
|
| 897 |
+
)> {
|
| 898 |
+
let listener = TcpListener::bind("127.0.0.1:0").await?;
|
| 899 |
+
let websocket_url = format!("ws://{}", listener.local_addr()?);
|
| 900 |
+
let server_task = tokio::spawn(async move {
|
| 901 |
+
let (stream, _) = listener.accept().await?;
|
| 902 |
+
accept_async(stream).await.map_err(anyhow::Error::from)
|
| 903 |
+
});
|
| 904 |
+
let (client_websocket, _) = connect_async(websocket_url).await?;
|
| 905 |
+
let server_websocket = server_task.await??;
|
| 906 |
+
Ok((client_websocket, server_websocket))
|
| 907 |
+
}
|
| 908 |
+
|
| 909 |
+
fn test_jsonrpc_message() -> JSONRPCMessage {
|
| 910 |
+
JSONRPCMessage::Request(JSONRPCRequest {
|
| 911 |
+
id: RequestId::Integer(1),
|
| 912 |
+
method: "test".to_string(),
|
| 913 |
+
params: None,
|
| 914 |
+
trace: None,
|
| 915 |
+
})
|
| 916 |
+
}
|
| 917 |
+
|
| 918 |
+
struct ControlledWebSocket {
|
| 919 |
+
inbound_rx: futures_mpsc::UnboundedReceiver<Result<Message, std::convert::Infallible>>,
|
| 920 |
+
outbound_tx: futures_mpsc::UnboundedSender<Message>,
|
| 921 |
+
write_ready: Arc<AtomicBool>,
|
| 922 |
+
write_blocked: Arc<AtomicBool>,
|
| 923 |
+
write_blocked_waker: Arc<AtomicWaker>,
|
| 924 |
+
write_waker: Arc<AtomicWaker>,
|
| 925 |
+
}
|
| 926 |
+
|
| 927 |
+
struct ControlledWebSocketHandle {
|
| 928 |
+
inbound_tx: futures_mpsc::UnboundedSender<Result<Message, std::convert::Infallible>>,
|
| 929 |
+
write_ready: Arc<AtomicBool>,
|
| 930 |
+
write_blocked: Arc<AtomicBool>,
|
| 931 |
+
write_blocked_waker: Arc<AtomicWaker>,
|
| 932 |
+
write_waker: Arc<AtomicWaker>,
|
| 933 |
+
}
|
| 934 |
+
|
| 935 |
+
impl ControlledWebSocket {
|
| 936 |
+
fn new(
|
| 937 |
+
write_ready: bool,
|
| 938 |
+
) -> (
|
| 939 |
+
Self,
|
| 940 |
+
ControlledWebSocketHandle,
|
| 941 |
+
futures_mpsc::UnboundedReceiver<Message>,
|
| 942 |
+
) {
|
| 943 |
+
let (inbound_tx, inbound_rx) = futures_mpsc::unbounded();
|
| 944 |
+
let (outbound_tx, outbound_rx) = futures_mpsc::unbounded();
|
| 945 |
+
let write_ready = Arc::new(AtomicBool::new(write_ready));
|
| 946 |
+
let write_blocked = Arc::new(AtomicBool::new(false));
|
| 947 |
+
let write_blocked_waker = Arc::new(AtomicWaker::new());
|
| 948 |
+
let write_waker = Arc::new(AtomicWaker::new());
|
| 949 |
+
(
|
| 950 |
+
Self {
|
| 951 |
+
inbound_rx,
|
| 952 |
+
outbound_tx,
|
| 953 |
+
write_ready: Arc::clone(&write_ready),
|
| 954 |
+
write_blocked: Arc::clone(&write_blocked),
|
| 955 |
+
write_blocked_waker: Arc::clone(&write_blocked_waker),
|
| 956 |
+
write_waker: Arc::clone(&write_waker),
|
| 957 |
+
},
|
| 958 |
+
ControlledWebSocketHandle {
|
| 959 |
+
inbound_tx,
|
| 960 |
+
write_ready,
|
| 961 |
+
write_blocked,
|
| 962 |
+
write_blocked_waker,
|
| 963 |
+
write_waker,
|
| 964 |
+
},
|
| 965 |
+
outbound_rx,
|
| 966 |
+
)
|
| 967 |
+
}
|
| 968 |
+
}
|
| 969 |
+
|
| 970 |
+
impl ControlledWebSocketHandle {
|
| 971 |
+
fn send_inbound(&self, message: Message) -> anyhow::Result<()> {
|
| 972 |
+
self.inbound_tx
|
| 973 |
+
.unbounded_send(Ok(message))
|
| 974 |
+
.map_err(anyhow::Error::from)
|
| 975 |
+
}
|
| 976 |
+
|
| 977 |
+
fn set_write_ready(&self) {
|
| 978 |
+
self.write_ready.store(true, Ordering::Release);
|
| 979 |
+
self.write_waker.wake();
|
| 980 |
+
}
|
| 981 |
+
|
| 982 |
+
async fn wait_for_blocked_write(&self) -> anyhow::Result<()> {
|
| 983 |
+
timeout(
|
| 984 |
+
Duration::from_secs(1),
|
| 985 |
+
futures::future::poll_fn(|cx| {
|
| 986 |
+
if self.write_blocked.load(Ordering::Acquire) {
|
| 987 |
+
Poll::Ready(())
|
| 988 |
+
} else {
|
| 989 |
+
self.write_blocked_waker.register(cx.waker());
|
| 990 |
+
Poll::Pending
|
| 991 |
+
}
|
| 992 |
+
}),
|
| 993 |
+
)
|
| 994 |
+
.await?;
|
| 995 |
+
Ok(())
|
| 996 |
+
}
|
| 997 |
+
}
|
| 998 |
+
|
| 999 |
+
impl Sink<Message> for ControlledWebSocket {
|
| 1000 |
+
type Error = std::convert::Infallible;
|
| 1001 |
+
|
| 1002 |
+
fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
| 1003 |
+
if self.write_ready.load(Ordering::Acquire) {
|
| 1004 |
+
Poll::Ready(Ok(()))
|
| 1005 |
+
} else {
|
| 1006 |
+
self.write_blocked.store(true, Ordering::Release);
|
| 1007 |
+
self.write_blocked_waker.wake();
|
| 1008 |
+
self.write_waker.register(cx.waker());
|
| 1009 |
+
Poll::Pending
|
| 1010 |
+
}
|
| 1011 |
+
}
|
| 1012 |
+
|
| 1013 |
+
fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
|
| 1014 |
+
self.outbound_tx
|
| 1015 |
+
.unbounded_send(item)
|
| 1016 |
+
.expect("test outbound receiver should stay open");
|
| 1017 |
+
Ok(())
|
| 1018 |
+
}
|
| 1019 |
+
|
| 1020 |
+
fn poll_flush(
|
| 1021 |
+
self: Pin<&mut Self>,
|
| 1022 |
+
_cx: &mut Context<'_>,
|
| 1023 |
+
) -> Poll<Result<(), Self::Error>> {
|
| 1024 |
+
Poll::Ready(Ok(()))
|
| 1025 |
+
}
|
| 1026 |
+
|
| 1027 |
+
fn poll_close(
|
| 1028 |
+
self: Pin<&mut Self>,
|
| 1029 |
+
_cx: &mut Context<'_>,
|
| 1030 |
+
) -> Poll<Result<(), Self::Error>> {
|
| 1031 |
+
Poll::Ready(Ok(()))
|
| 1032 |
+
}
|
| 1033 |
+
}
|
| 1034 |
+
|
| 1035 |
+
impl Stream for ControlledWebSocket {
|
| 1036 |
+
type Item = Result<Message, std::convert::Infallible>;
|
| 1037 |
+
|
| 1038 |
+
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
| 1039 |
+
Pin::new(&mut self.inbound_rx).poll_next(cx)
|
| 1040 |
+
}
|
| 1041 |
+
}
|
| 1042 |
+
}
|