diff --git a/codex-rs/config/src/cloud_requirements.rs b/codex-rs/config/src/cloud_requirements.rs new file mode 100644 index 0000000000..3487cc326a --- /dev/null +++ b/codex-rs/config/src/cloud_requirements.rs @@ -0,0 +1,62 @@ +use crate::config_loader::ConfigRequirementsToml; +use futures::future::BoxFuture; +use futures::future::FutureExt; +use futures::future::Shared; +use std::fmt; +use std::future::Future; + +#[derive(Clone)] +pub struct CloudRequirementsLoader { + // TODO(gt): This should return a Result once we can fail-closed. + fut: Shared>>, +} + +impl CloudRequirementsLoader { + pub fn new(fut: F) -> Self + where + F: Future> + Send + 'static, + { + Self { + fut: fut.boxed().shared(), + } + } + + pub async fn get(&self) -> Option { + self.fut.clone().await + } +} + +impl fmt::Debug for CloudRequirementsLoader { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("CloudRequirementsLoader").finish() + } +} + +impl Default for CloudRequirementsLoader { + fn default() -> Self { + Self::new(async { None }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use pretty_assertions::assert_eq; + use std::sync::Arc; + use std::sync::atomic::AtomicUsize; + use std::sync::atomic::Ordering; + + #[tokio::test] + async fn shared_future_runs_once() { + let counter = Arc::new(AtomicUsize::new(0)); + let counter_clone = Arc::clone(&counter); + let loader = CloudRequirementsLoader::new(async move { + counter_clone.fetch_add(1, Ordering::SeqCst); + Some(ConfigRequirementsToml::default()) + }); + + let (first, second) = tokio::join!(loader.get(), loader.get()); + assert_eq!(first, second); + assert_eq!(counter.load(Ordering::SeqCst), 1); + } +} diff --git a/codex-rs/config/src/config_loader.rs b/codex-rs/config/src/config_loader.rs new file mode 100644 index 0000000000..b9819afd66 --- /dev/null +++ b/codex-rs/config/src/config_loader.rs @@ -0,0 +1,36 @@ +#[path = "cloud_requirements.rs"] +pub mod cloud_requirements; +#[path = "config_requirements.rs"] +pub mod config_requirements; +#[path = "fingerprint.rs"] +pub mod fingerprint; +#[path = "merge.rs"] +pub mod merge; +#[path = "overrides.rs"] +pub mod overrides; +#[path = "requirements_exec_policy.rs"] +pub mod requirements_exec_policy; +#[path = "state.rs"] +pub mod state; + +pub use cloud_requirements::CloudRequirementsLoader; +pub use config_requirements::ConfigRequirements; +pub use config_requirements::ConfigRequirementsToml; +pub use config_requirements::ConfigRequirementsWithSources; +pub use config_requirements::ConstrainedWithSource; +pub use config_requirements::McpServerIdentity; +pub use config_requirements::McpServerRequirement; +pub use config_requirements::NetworkConstraints; +pub use config_requirements::NetworkRequirementsToml; +pub use config_requirements::RequirementSource; +pub use config_requirements::ResidencyRequirement; +pub use config_requirements::SandboxModeRequirement; +pub use config_requirements::Sourced; +pub use config_requirements::WebSearchModeRequirement; +pub use fingerprint::version_for_toml; +pub use merge::merge_toml_values; +pub use overrides::build_cli_overrides_layer; +pub use state::ConfigLayerEntry; +pub use state::ConfigLayerStack; +pub use state::ConfigLayerStackOrdering; +pub use state::LoaderOverrides; diff --git a/codex-rs/config/src/config_requirements.rs b/codex-rs/config/src/config_requirements.rs new file mode 100644 index 0000000000..8632023d48 --- /dev/null +++ b/codex-rs/config/src/config_requirements.rs @@ -0,0 +1,1177 @@ +use codex_protocol::config_types::SandboxMode; +use codex_protocol::config_types::WebSearchMode; +use codex_protocol::protocol::AskForApproval; +use codex_protocol::protocol::SandboxPolicy; +use codex_utils_absolute_path::AbsolutePathBuf; +use serde::Deserialize; +use serde::Serialize; +use std::collections::BTreeMap; +use std::fmt; + +use super::requirements_exec_policy::RequirementsExecPolicy; +use super::requirements_exec_policy::RequirementsExecPolicyToml; +use crate::Constrained; +use crate::ConstraintError; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum RequirementSource { + Unknown, + MdmManagedPreferences { domain: String, key: String }, + CloudRequirements, + SystemRequirementsToml { file: AbsolutePathBuf }, + LegacyManagedConfigTomlFromFile { file: AbsolutePathBuf }, + LegacyManagedConfigTomlFromMdm, +} + +impl fmt::Display for RequirementSource { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + RequirementSource::Unknown => write!(f, ""), + RequirementSource::MdmManagedPreferences { domain, key } => { + write!(f, "MDM {domain}:{key}") + } + RequirementSource::CloudRequirements => { + write!(f, "cloud requirements") + } + RequirementSource::SystemRequirementsToml { file } => { + write!(f, "{}", file.as_path().display()) + } + RequirementSource::LegacyManagedConfigTomlFromFile { file } => { + write!(f, "{}", file.as_path().display()) + } + RequirementSource::LegacyManagedConfigTomlFromMdm => { + write!(f, "MDM managed_config.toml (legacy)") + } + } + } +} + +#[derive(Debug, Clone, PartialEq)] +pub struct ConstrainedWithSource { + pub value: Constrained, + pub source: Option, +} + +impl ConstrainedWithSource { + pub fn new(value: Constrained, source: Option) -> Self { + Self { value, source } + } +} + +impl std::ops::Deref for ConstrainedWithSource { + type Target = Constrained; + + fn deref(&self) -> &Self::Target { + &self.value + } +} + +impl std::ops::DerefMut for ConstrainedWithSource { + fn deref_mut(&mut self) -> &mut Self::Target { + &mut self.value + } +} + +/// Normalized version of [`ConfigRequirementsToml`] after deserialization and +/// normalization. +#[derive(Debug, Clone, PartialEq)] +pub struct ConfigRequirements { + pub approval_policy: ConstrainedWithSource, + pub sandbox_policy: ConstrainedWithSource, + pub web_search_mode: ConstrainedWithSource, + pub mcp_servers: Option>>, + pub exec_policy: Option>, + pub enforce_residency: ConstrainedWithSource>, + /// Managed network constraints derived from requirements. + pub network: Option>, +} + +impl Default for ConfigRequirements { + fn default() -> Self { + Self { + approval_policy: ConstrainedWithSource::new( + Constrained::allow_any_from_default(), + None, + ), + sandbox_policy: ConstrainedWithSource::new( + Constrained::allow_any(SandboxPolicy::ReadOnly), + None, + ), + web_search_mode: ConstrainedWithSource::new( + Constrained::allow_any(WebSearchMode::Cached), + None, + ), + mcp_servers: None, + exec_policy: None, + enforce_residency: ConstrainedWithSource::new(Constrained::allow_any(None), None), + network: None, + } + } +} + +impl ConfigRequirements { + pub fn exec_policy_source(&self) -> Option<&RequirementSource> { + self.exec_policy.as_ref().map(|policy| &policy.source) + } +} + +#[derive(Deserialize, Debug, Clone, PartialEq, Eq)] +#[serde(untagged)] +pub enum McpServerIdentity { + Command { command: String }, + Url { url: String }, +} + +#[derive(Deserialize, Debug, Clone, PartialEq, Eq)] +pub struct McpServerRequirement { + pub identity: McpServerIdentity, +} + +#[derive(Deserialize, Debug, Clone, Default, PartialEq, Eq)] +pub struct NetworkRequirementsToml { + pub enabled: Option, + pub http_port: Option, + pub socks_port: Option, + pub allow_upstream_proxy: Option, + pub dangerously_allow_non_loopback_proxy: Option, + pub dangerously_allow_non_loopback_admin: Option, + pub allowed_domains: Option>, + pub denied_domains: Option>, + pub allow_unix_sockets: Option>, + pub allow_local_binding: Option, +} + +/// Normalized network constraints derived from requirements TOML. +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct NetworkConstraints { + pub enabled: Option, + pub http_port: Option, + pub socks_port: Option, + pub allow_upstream_proxy: Option, + pub dangerously_allow_non_loopback_proxy: Option, + pub dangerously_allow_non_loopback_admin: Option, + pub allowed_domains: Option>, + pub denied_domains: Option>, + pub allow_unix_sockets: Option>, + pub allow_local_binding: Option, +} + +impl From for NetworkConstraints { + fn from(value: NetworkRequirementsToml) -> Self { + let NetworkRequirementsToml { + enabled, + http_port, + socks_port, + allow_upstream_proxy, + dangerously_allow_non_loopback_proxy, + dangerously_allow_non_loopback_admin, + allowed_domains, + denied_domains, + allow_unix_sockets, + allow_local_binding, + } = value; + Self { + enabled, + http_port, + socks_port, + allow_upstream_proxy, + dangerously_allow_non_loopback_proxy, + dangerously_allow_non_loopback_admin, + allowed_domains, + denied_domains, + allow_unix_sockets, + allow_local_binding, + } + } +} + +#[derive(Deserialize, Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[serde(rename_all = "lowercase")] +pub enum WebSearchModeRequirement { + Disabled, + Cached, + Live, +} + +impl From for WebSearchModeRequirement { + fn from(mode: WebSearchMode) -> Self { + match mode { + WebSearchMode::Disabled => WebSearchModeRequirement::Disabled, + WebSearchMode::Cached => WebSearchModeRequirement::Cached, + WebSearchMode::Live => WebSearchModeRequirement::Live, + } + } +} + +impl From for WebSearchMode { + fn from(mode: WebSearchModeRequirement) -> Self { + match mode { + WebSearchModeRequirement::Disabled => WebSearchMode::Disabled, + WebSearchModeRequirement::Cached => WebSearchMode::Cached, + WebSearchModeRequirement::Live => WebSearchMode::Live, + } + } +} + +impl fmt::Display for WebSearchModeRequirement { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + WebSearchModeRequirement::Disabled => write!(f, "disabled"), + WebSearchModeRequirement::Cached => write!(f, "cached"), + WebSearchModeRequirement::Live => write!(f, "live"), + } + } +} + +/// Base config deserialized from system `requirements.toml` or MDM. +#[derive(Deserialize, Debug, Clone, Default, PartialEq)] +pub struct ConfigRequirementsToml { + pub allowed_approval_policies: Option>, + pub allowed_sandbox_modes: Option>, + pub allowed_web_search_modes: Option>, + pub mcp_servers: Option>, + pub rules: Option, + pub enforce_residency: Option, + #[serde(rename = "experimental_network")] + pub network: Option, +} + +/// Value paired with the requirement source it came from, for better error +/// messages. +#[derive(Debug, Clone, PartialEq)] +pub struct Sourced { + pub value: T, + pub source: RequirementSource, +} + +impl Sourced { + pub fn new(value: T, source: RequirementSource) -> Self { + Self { value, source } + } +} + +impl std::ops::Deref for Sourced { + type Target = T; + + fn deref(&self) -> &Self::Target { + &self.value + } +} + +#[derive(Debug, Clone, Default, PartialEq)] +pub struct ConfigRequirementsWithSources { + pub allowed_approval_policies: Option>>, + pub allowed_sandbox_modes: Option>>, + pub allowed_web_search_modes: Option>>, + pub mcp_servers: Option>>, + pub rules: Option>, + pub enforce_residency: Option>, + pub network: Option>, +} + +impl ConfigRequirementsWithSources { + pub fn merge_unset_fields(&mut self, source: RequirementSource, other: ConfigRequirementsToml) { + // For every field in `other` that is `Some`, if the corresponding field + // in `self` is `None`, copy the value from `other` into `self`. + macro_rules! fill_missing_take { + ($base:expr, $other:expr, $source:expr, { $($field:ident),+ $(,)? }) => { + // Destructure without `..` so adding fields to `ConfigRequirementsToml` + // forces this merge logic to be updated. + let ConfigRequirementsToml { $($field: _,)+ } = &$other; + + $( + if $base.$field.is_none() + && let Some(value) = $other.$field.take() + { + $base.$field = Some(Sourced::new(value, $source.clone())); + } + )+ + }; + } + + let mut other = other; + fill_missing_take!( + self, + other, + source, + { + allowed_approval_policies, + allowed_sandbox_modes, + allowed_web_search_modes, + mcp_servers, + rules, + enforce_residency, + network, + } + ); + } + + pub fn into_toml(self) -> ConfigRequirementsToml { + let ConfigRequirementsWithSources { + allowed_approval_policies, + allowed_sandbox_modes, + allowed_web_search_modes, + mcp_servers, + rules, + enforce_residency, + network, + } = self; + ConfigRequirementsToml { + allowed_approval_policies: allowed_approval_policies.map(|sourced| sourced.value), + allowed_sandbox_modes: allowed_sandbox_modes.map(|sourced| sourced.value), + allowed_web_search_modes: allowed_web_search_modes.map(|sourced| sourced.value), + mcp_servers: mcp_servers.map(|sourced| sourced.value), + rules: rules.map(|sourced| sourced.value), + enforce_residency: enforce_residency.map(|sourced| sourced.value), + network: network.map(|sourced| sourced.value), + } + } +} + +/// Currently, `external-sandbox` is not supported in config.toml, but it is +/// supported through programmatic use. +#[derive(Deserialize, Debug, Clone, Copy, PartialEq)] +pub enum SandboxModeRequirement { + #[serde(rename = "read-only")] + ReadOnly, + + #[serde(rename = "workspace-write")] + WorkspaceWrite, + + #[serde(rename = "danger-full-access")] + DangerFullAccess, + + #[serde(rename = "external-sandbox")] + ExternalSandbox, +} + +impl From for SandboxModeRequirement { + fn from(mode: SandboxMode) -> Self { + match mode { + SandboxMode::ReadOnly => SandboxModeRequirement::ReadOnly, + SandboxMode::WorkspaceWrite => SandboxModeRequirement::WorkspaceWrite, + SandboxMode::DangerFullAccess => SandboxModeRequirement::DangerFullAccess, + } + } +} + +#[derive(Deserialize, Debug, Clone, Copy, PartialEq, Eq)] +#[serde(rename_all = "lowercase")] +pub enum ResidencyRequirement { + Us, +} + +impl ConfigRequirementsToml { + pub fn is_empty(&self) -> bool { + self.allowed_approval_policies.is_none() + && self.allowed_sandbox_modes.is_none() + && self.allowed_web_search_modes.is_none() + && self.mcp_servers.is_none() + && self.rules.is_none() + && self.enforce_residency.is_none() + && self.network.is_none() + } +} + +impl TryFrom for ConfigRequirements { + type Error = ConstraintError; + + fn try_from(toml: ConfigRequirementsWithSources) -> Result { + let ConfigRequirementsWithSources { + allowed_approval_policies, + allowed_sandbox_modes, + allowed_web_search_modes, + mcp_servers, + rules, + enforce_residency, + network, + } = toml; + + let approval_policy = match allowed_approval_policies { + Some(Sourced { + value: policies, + source: requirement_source, + }) => { + let Some(initial_value) = policies.first().copied() else { + return Err(ConstraintError::empty_field("allowed_approval_policies")); + }; + + let requirement_source_for_error = requirement_source.clone(); + let constrained = Constrained::new(initial_value, move |candidate| { + if policies.contains(candidate) { + Ok(()) + } else { + Err(ConstraintError::InvalidValue { + field_name: "approval_policy", + candidate: format!("{candidate:?}"), + allowed: format!("{policies:?}"), + requirement_source: requirement_source_for_error.clone(), + }) + } + })?; + ConstrainedWithSource::new(constrained, Some(requirement_source)) + } + None => ConstrainedWithSource::new(Constrained::allow_any_from_default(), None), + }; + + // TODO(gt): `ConfigRequirementsToml` should let the author specify the + // default `SandboxPolicy`? Should do this for `AskForApproval` too? + // + // Currently, we force ReadOnly as the default policy because two of + // the other variants (WorkspaceWrite, ExternalSandbox) require + // additional parameters. Ultimately, we should expand the config + // format to allow specifying those parameters. + let default_sandbox_policy = SandboxPolicy::ReadOnly; + let sandbox_policy = match allowed_sandbox_modes { + Some(Sourced { + value: modes, + source: requirement_source, + }) => { + if !modes.contains(&SandboxModeRequirement::ReadOnly) { + return Err(ConstraintError::InvalidValue { + field_name: "allowed_sandbox_modes", + candidate: format!("{modes:?}"), + allowed: "must include 'read-only' to allow any SandboxPolicy".to_string(), + requirement_source, + }); + }; + + let requirement_source_for_error = requirement_source.clone(); + let constrained = Constrained::new(default_sandbox_policy, move |candidate| { + let mode = match candidate { + SandboxPolicy::ReadOnly => SandboxModeRequirement::ReadOnly, + SandboxPolicy::WorkspaceWrite { .. } => { + SandboxModeRequirement::WorkspaceWrite + } + SandboxPolicy::DangerFullAccess => SandboxModeRequirement::DangerFullAccess, + SandboxPolicy::ExternalSandbox { .. } => { + SandboxModeRequirement::ExternalSandbox + } + }; + if modes.contains(&mode) { + Ok(()) + } else { + Err(ConstraintError::InvalidValue { + field_name: "sandbox_mode", + candidate: format!("{mode:?}"), + allowed: format!("{modes:?}"), + requirement_source: requirement_source_for_error.clone(), + }) + } + })?; + ConstrainedWithSource::new(constrained, Some(requirement_source)) + } + None => { + ConstrainedWithSource::new(Constrained::allow_any(default_sandbox_policy), None) + } + }; + let exec_policy = match rules { + Some(Sourced { value, source }) => { + let policy = value.to_requirements_policy().map_err(|err| { + ConstraintError::ExecPolicyParse { + requirement_source: source.clone(), + reason: err.to_string(), + } + })?; + Some(Sourced::new(policy, source)) + } + None => None, + }; + let web_search_mode = match allowed_web_search_modes { + Some(Sourced { + value: modes, + source: requirement_source, + }) => { + let mut accepted = modes.into_iter().collect::>(); + accepted.insert(WebSearchModeRequirement::Disabled); + let allowed_for_error = format!( + "{:?}", + accepted + .iter() + .copied() + .map(WebSearchMode::from) + .collect::>() + ); + + let initial_value = if accepted.contains(&WebSearchModeRequirement::Cached) { + WebSearchMode::Cached + } else if accepted.contains(&WebSearchModeRequirement::Live) { + WebSearchMode::Live + } else { + WebSearchMode::Disabled + }; + let requirement_source_for_error = requirement_source.clone(); + let constrained = Constrained::new(initial_value, move |candidate| { + if accepted.contains(&(*candidate).into()) { + Ok(()) + } else { + Err(ConstraintError::InvalidValue { + field_name: "web_search_mode", + candidate: format!("{candidate:?}"), + allowed: allowed_for_error.clone(), + requirement_source: requirement_source_for_error.clone(), + }) + } + })?; + ConstrainedWithSource::new(constrained, Some(requirement_source)) + } + None => ConstrainedWithSource::new(Constrained::allow_any(WebSearchMode::Cached), None), + }; + + let enforce_residency = match enforce_residency { + Some(Sourced { + value: residency, + source: requirement_source, + }) => { + let required = Some(residency); + let requirement_source_for_error = requirement_source.clone(); + let constrained = Constrained::new(required, move |candidate| { + if candidate == &required { + Ok(()) + } else { + Err(ConstraintError::InvalidValue { + field_name: "enforce_residency", + candidate: format!("{candidate:?}"), + allowed: format!("{required:?}"), + requirement_source: requirement_source_for_error.clone(), + }) + } + })?; + ConstrainedWithSource::new(constrained, Some(requirement_source)) + } + None => ConstrainedWithSource::new(Constrained::allow_any(None), None), + }; + let network = network.map(|sourced_network| { + let Sourced { value, source } = sourced_network; + Sourced::new(NetworkConstraints::from(value), source) + }); + Ok(ConfigRequirements { + approval_policy, + sandbox_policy, + web_search_mode, + mcp_servers, + exec_policy, + enforce_residency, + network, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use anyhow::Result; + use codex_execpolicy::Decision; + use codex_execpolicy::Evaluation; + use codex_execpolicy::RuleMatch; + use codex_protocol::protocol::NetworkAccess; + use codex_utils_absolute_path::AbsolutePathBuf; + use pretty_assertions::assert_eq; + use toml::from_str; + + fn tokens(cmd: &[&str]) -> Vec { + cmd.iter().map(std::string::ToString::to_string).collect() + } + + fn system_requirements_toml_file_for_test() -> Result { + Ok(AbsolutePathBuf::try_from( + std::env::temp_dir().join("requirements.toml"), + )?) + } + + fn with_unknown_source(toml: ConfigRequirementsToml) -> ConfigRequirementsWithSources { + let ConfigRequirementsToml { + allowed_approval_policies, + allowed_sandbox_modes, + allowed_web_search_modes, + mcp_servers, + rules, + enforce_residency, + network, + } = toml; + ConfigRequirementsWithSources { + allowed_approval_policies: allowed_approval_policies + .map(|value| Sourced::new(value, RequirementSource::Unknown)), + allowed_sandbox_modes: allowed_sandbox_modes + .map(|value| Sourced::new(value, RequirementSource::Unknown)), + allowed_web_search_modes: allowed_web_search_modes + .map(|value| Sourced::new(value, RequirementSource::Unknown)), + mcp_servers: mcp_servers.map(|value| Sourced::new(value, RequirementSource::Unknown)), + rules: rules.map(|value| Sourced::new(value, RequirementSource::Unknown)), + enforce_residency: enforce_residency + .map(|value| Sourced::new(value, RequirementSource::Unknown)), + network: network.map(|value| Sourced::new(value, RequirementSource::Unknown)), + } + } + + #[test] + fn merge_unset_fields_copies_every_field_and_sets_sources() { + let mut target = ConfigRequirementsWithSources::default(); + let source = RequirementSource::LegacyManagedConfigTomlFromMdm; + + let allowed_approval_policies = vec![AskForApproval::UnlessTrusted, AskForApproval::Never]; + let allowed_sandbox_modes = vec![ + SandboxModeRequirement::WorkspaceWrite, + SandboxModeRequirement::DangerFullAccess, + ]; + let allowed_web_search_modes = vec![ + WebSearchModeRequirement::Cached, + WebSearchModeRequirement::Live, + ]; + let enforce_residency = ResidencyRequirement::Us; + let enforce_source = source.clone(); + + // Intentionally constructed without `..Default::default()` so adding a new field to + // `ConfigRequirementsToml` forces this test to be updated. + let other = ConfigRequirementsToml { + allowed_approval_policies: Some(allowed_approval_policies.clone()), + allowed_sandbox_modes: Some(allowed_sandbox_modes.clone()), + allowed_web_search_modes: Some(allowed_web_search_modes.clone()), + mcp_servers: None, + rules: None, + enforce_residency: Some(enforce_residency), + network: None, + }; + + target.merge_unset_fields(source.clone(), other); + + assert_eq!( + target, + ConfigRequirementsWithSources { + allowed_approval_policies: Some(Sourced::new( + allowed_approval_policies, + source.clone() + )), + allowed_sandbox_modes: Some(Sourced::new(allowed_sandbox_modes, source)), + allowed_web_search_modes: Some(Sourced::new( + allowed_web_search_modes, + enforce_source.clone(), + )), + mcp_servers: None, + rules: None, + enforce_residency: Some(Sourced::new(enforce_residency, enforce_source)), + network: None, + } + ); + } + + #[test] + fn merge_unset_fields_fills_missing_values() -> Result<()> { + let source: ConfigRequirementsToml = from_str( + r#" + allowed_approval_policies = ["on-request"] + "#, + )?; + + let source_location = RequirementSource::MdmManagedPreferences { + domain: "com.codex".to_string(), + key: "allowed_approval_policies".to_string(), + }; + + let mut empty_target = ConfigRequirementsWithSources::default(); + empty_target.merge_unset_fields(source_location.clone(), source); + assert_eq!( + empty_target, + ConfigRequirementsWithSources { + allowed_approval_policies: Some(Sourced::new( + vec![AskForApproval::OnRequest], + source_location, + )), + allowed_sandbox_modes: None, + allowed_web_search_modes: None, + mcp_servers: None, + rules: None, + enforce_residency: None, + network: None, + } + ); + Ok(()) + } + + #[test] + fn merge_unset_fields_does_not_overwrite_existing_values() -> Result<()> { + let existing_source = RequirementSource::LegacyManagedConfigTomlFromMdm; + let mut populated_target = ConfigRequirementsWithSources::default(); + let populated_requirements: ConfigRequirementsToml = from_str( + r#" + allowed_approval_policies = ["never"] + "#, + )?; + populated_target.merge_unset_fields(existing_source.clone(), populated_requirements); + + let source: ConfigRequirementsToml = from_str( + r#" + allowed_approval_policies = ["on-request"] + "#, + )?; + let source_location = RequirementSource::MdmManagedPreferences { + domain: "com.codex".to_string(), + key: "allowed_approval_policies".to_string(), + }; + populated_target.merge_unset_fields(source_location, source); + + assert_eq!( + populated_target, + ConfigRequirementsWithSources { + allowed_approval_policies: Some(Sourced::new( + vec![AskForApproval::Never], + existing_source, + )), + allowed_sandbox_modes: None, + allowed_web_search_modes: None, + mcp_servers: None, + rules: None, + enforce_residency: None, + network: None, + } + ); + Ok(()) + } + + #[test] + fn constraint_error_includes_requirement_source() -> Result<()> { + let source: ConfigRequirementsToml = from_str( + r#" + allowed_approval_policies = ["on-request"] + allowed_sandbox_modes = ["read-only"] + "#, + )?; + + let requirements_toml_file = system_requirements_toml_file_for_test()?; + let source_location = RequirementSource::SystemRequirementsToml { + file: requirements_toml_file, + }; + + let mut target = ConfigRequirementsWithSources::default(); + target.merge_unset_fields(source_location.clone(), source); + let requirements = ConfigRequirements::try_from(target)?; + + assert_eq!( + requirements.approval_policy.can_set(&AskForApproval::Never), + Err(ConstraintError::InvalidValue { + field_name: "approval_policy", + candidate: "Never".into(), + allowed: "[OnRequest]".into(), + requirement_source: source_location.clone(), + }) + ); + assert_eq!( + requirements + .sandbox_policy + .can_set(&SandboxPolicy::DangerFullAccess), + Err(ConstraintError::InvalidValue { + field_name: "sandbox_mode", + candidate: "DangerFullAccess".into(), + allowed: "[ReadOnly]".into(), + requirement_source: source_location, + }) + ); + + Ok(()) + } + + #[test] + fn constraint_error_includes_cloud_requirements_source() -> Result<()> { + let source: ConfigRequirementsToml = from_str( + r#" + allowed_approval_policies = ["on-request"] + "#, + )?; + + let source_location = RequirementSource::CloudRequirements; + + let mut target = ConfigRequirementsWithSources::default(); + target.merge_unset_fields(source_location.clone(), source); + let requirements = ConfigRequirements::try_from(target)?; + + assert_eq!( + requirements.approval_policy.can_set(&AskForApproval::Never), + Err(ConstraintError::InvalidValue { + field_name: "approval_policy", + candidate: "Never".into(), + allowed: "[OnRequest]".into(), + requirement_source: source_location, + }) + ); + + Ok(()) + } + + #[test] + fn constrained_fields_store_requirement_source() -> Result<()> { + let source: ConfigRequirementsToml = from_str( + r#" + allowed_approval_policies = ["on-request"] + allowed_sandbox_modes = ["read-only"] + allowed_web_search_modes = ["cached"] + enforce_residency = "us" + "#, + )?; + + let source_location = RequirementSource::CloudRequirements; + let mut target = ConfigRequirementsWithSources::default(); + target.merge_unset_fields(source_location.clone(), source); + let requirements = ConfigRequirements::try_from(target)?; + + assert_eq!( + requirements.approval_policy.source, + Some(source_location.clone()) + ); + assert_eq!( + requirements.sandbox_policy.source, + Some(source_location.clone()) + ); + assert_eq!( + requirements.web_search_mode.source, + Some(source_location.clone()) + ); + assert_eq!(requirements.enforce_residency.source, Some(source_location)); + + Ok(()) + } + + #[test] + fn deserialize_allowed_approval_policies() -> Result<()> { + let toml_str = r#" + allowed_approval_policies = ["untrusted", "on-request"] + "#; + let config: ConfigRequirementsToml = from_str(toml_str)?; + let requirements: ConfigRequirements = with_unknown_source(config).try_into()?; + + assert_eq!( + requirements.approval_policy.value(), + AskForApproval::UnlessTrusted, + "currently, there is no way to specify the default value for approval policy in the toml, so it picks the first allowed value" + ); + assert!( + requirements + .approval_policy + .can_set(&AskForApproval::UnlessTrusted) + .is_ok() + ); + assert_eq!( + requirements + .approval_policy + .can_set(&AskForApproval::OnFailure), + Err(ConstraintError::InvalidValue { + field_name: "approval_policy", + candidate: "OnFailure".into(), + allowed: "[UnlessTrusted, OnRequest]".into(), + requirement_source: RequirementSource::Unknown, + }) + ); + assert!( + requirements + .approval_policy + .can_set(&AskForApproval::OnRequest) + .is_ok() + ); + assert_eq!( + requirements.approval_policy.can_set(&AskForApproval::Never), + Err(ConstraintError::InvalidValue { + field_name: "approval_policy", + candidate: "Never".into(), + allowed: "[UnlessTrusted, OnRequest]".into(), + requirement_source: RequirementSource::Unknown, + }) + ); + assert!( + requirements + .sandbox_policy + .can_set(&SandboxPolicy::ReadOnly) + .is_ok() + ); + + Ok(()) + } + + #[test] + fn deserialize_allowed_sandbox_modes() -> Result<()> { + let toml_str = r#" + allowed_sandbox_modes = ["read-only", "workspace-write"] + "#; + let config: ConfigRequirementsToml = from_str(toml_str)?; + let requirements: ConfigRequirements = with_unknown_source(config).try_into()?; + + let root = if cfg!(windows) { "C:\\repo" } else { "/repo" }; + assert!( + requirements + .sandbox_policy + .can_set(&SandboxPolicy::ReadOnly) + .is_ok() + ); + assert!( + requirements + .sandbox_policy + .can_set(&SandboxPolicy::WorkspaceWrite { + writable_roots: vec![AbsolutePathBuf::from_absolute_path(root)?], + network_access: false, + exclude_tmpdir_env_var: false, + exclude_slash_tmp: false, + }) + .is_ok() + ); + assert_eq!( + requirements + .sandbox_policy + .can_set(&SandboxPolicy::DangerFullAccess), + Err(ConstraintError::InvalidValue { + field_name: "sandbox_mode", + candidate: "DangerFullAccess".into(), + allowed: "[ReadOnly, WorkspaceWrite]".into(), + requirement_source: RequirementSource::Unknown, + }) + ); + assert_eq!( + requirements + .sandbox_policy + .can_set(&SandboxPolicy::ExternalSandbox { + network_access: NetworkAccess::Restricted, + }), + Err(ConstraintError::InvalidValue { + field_name: "sandbox_mode", + candidate: "ExternalSandbox".into(), + allowed: "[ReadOnly, WorkspaceWrite]".into(), + requirement_source: RequirementSource::Unknown, + }) + ); + + Ok(()) + } + + #[test] + fn deserialize_allowed_web_search_modes() -> Result<()> { + let toml_str = r#" + allowed_web_search_modes = ["cached"] + "#; + let config: ConfigRequirementsToml = from_str(toml_str)?; + let requirements: ConfigRequirements = with_unknown_source(config).try_into()?; + + assert_eq!(requirements.web_search_mode.value(), WebSearchMode::Cached); + assert!( + requirements + .web_search_mode + .can_set(&WebSearchMode::Disabled) + .is_ok() + ); + assert_eq!( + requirements.web_search_mode.can_set(&WebSearchMode::Live), + Err(ConstraintError::InvalidValue { + field_name: "web_search_mode", + candidate: "Live".into(), + allowed: "[Disabled, Cached]".into(), + requirement_source: RequirementSource::Unknown, + }) + ); + assert!( + requirements + .web_search_mode + .can_set(&WebSearchMode::Cached) + .is_ok() + ); + + Ok(()) + } + + #[test] + fn allowed_web_search_modes_allows_disabled() -> Result<()> { + let toml_str = r#" + allowed_web_search_modes = ["disabled"] + "#; + let config: ConfigRequirementsToml = from_str(toml_str)?; + let requirements: ConfigRequirements = with_unknown_source(config).try_into()?; + + assert_eq!( + requirements.web_search_mode.value(), + WebSearchMode::Disabled + ); + assert!( + requirements + .web_search_mode + .can_set(&WebSearchMode::Disabled) + .is_ok() + ); + assert_eq!( + requirements.web_search_mode.can_set(&WebSearchMode::Cached), + Err(ConstraintError::InvalidValue { + field_name: "web_search_mode", + candidate: "Cached".into(), + allowed: "[Disabled]".into(), + requirement_source: RequirementSource::Unknown, + }) + ); + Ok(()) + } + + #[test] + fn allowed_web_search_modes_empty_restricts_to_disabled() -> Result<()> { + let toml_str = r#" + allowed_web_search_modes = [] + "#; + let config: ConfigRequirementsToml = from_str(toml_str)?; + let requirements: ConfigRequirements = with_unknown_source(config).try_into()?; + + assert_eq!( + requirements.web_search_mode.value(), + WebSearchMode::Disabled + ); + assert!( + requirements + .web_search_mode + .can_set(&WebSearchMode::Disabled) + .is_ok() + ); + assert_eq!( + requirements.web_search_mode.can_set(&WebSearchMode::Cached), + Err(ConstraintError::InvalidValue { + field_name: "web_search_mode", + candidate: "Cached".into(), + allowed: "[Disabled]".into(), + requirement_source: RequirementSource::Unknown, + }) + ); + Ok(()) + } + + #[test] + fn network_requirements_are_preserved_as_constraints_with_source() -> Result<()> { + let toml_str = r#" + [experimental_network] + enabled = true + allow_upstream_proxy = false + allowed_domains = ["api.example.com", "*.openai.com"] + denied_domains = ["blocked.example.com"] + allow_unix_sockets = ["/tmp/example.sock"] + allow_local_binding = false + "#; + + let source = RequirementSource::CloudRequirements; + let mut requirements_with_sources = ConfigRequirementsWithSources::default(); + requirements_with_sources.merge_unset_fields(source.clone(), from_str(toml_str)?); + + let requirements = ConfigRequirements::try_from(requirements_with_sources)?; + let sourced_network = requirements + .network + .expect("network requirements should be preserved as constraints"); + + assert_eq!(sourced_network.source, source); + assert_eq!(sourced_network.value.enabled, Some(true)); + assert_eq!(sourced_network.value.allow_upstream_proxy, Some(false)); + assert_eq!( + sourced_network.value.allowed_domains.as_ref(), + Some(&vec![ + "api.example.com".to_string(), + "*.openai.com".to_string() + ]) + ); + assert_eq!( + sourced_network.value.denied_domains.as_ref(), + Some(&vec!["blocked.example.com".to_string()]) + ); + assert_eq!( + sourced_network.value.allow_unix_sockets.as_ref(), + Some(&vec!["/tmp/example.sock".to_string()]) + ); + assert_eq!(sourced_network.value.allow_local_binding, Some(false)); + + Ok(()) + } + + #[test] + fn deserialize_mcp_server_requirements() -> Result<()> { + let toml_str = r#" + [mcp_servers.docs.identity] + command = "codex-mcp" + + [mcp_servers.remote.identity] + url = "https://example.com/mcp" + "#; + let requirements: ConfigRequirements = + with_unknown_source(from_str(toml_str)?).try_into()?; + + assert_eq!( + requirements.mcp_servers, + Some(Sourced::new( + BTreeMap::from([ + ( + "docs".to_string(), + McpServerRequirement { + identity: McpServerIdentity::Command { + command: "codex-mcp".to_string(), + }, + }, + ), + ( + "remote".to_string(), + McpServerRequirement { + identity: McpServerIdentity::Url { + url: "https://example.com/mcp".to_string(), + }, + }, + ), + ]), + RequirementSource::Unknown, + )) + ); + Ok(()) + } + + #[test] + fn deserialize_exec_policy_requirements() -> Result<()> { + let toml_str = r#" + [rules] + prefix_rules = [ + { pattern = [{ token = "rm" }], decision = "forbidden" }, + ] + "#; + let config: ConfigRequirementsToml = from_str(toml_str)?; + let requirements: ConfigRequirements = with_unknown_source(config).try_into()?; + let policy = requirements.exec_policy.expect("exec policy").value; + + assert_eq!( + policy.as_ref().check(&tokens(&["rm", "-rf"]), &|_| { + panic!("rule should match so heuristic should not be called"); + }), + Evaluation { + decision: Decision::Forbidden, + matched_rules: vec![RuleMatch::PrefixRuleMatch { + matched_prefix: tokens(&["rm"]), + decision: Decision::Forbidden, + justification: None, + }], + } + ); + + Ok(()) + } + + #[test] + fn exec_policy_error_includes_requirement_source() -> Result<()> { + let toml_str = r#" + [rules] + prefix_rules = [ + { pattern = [{ token = "rm" }] }, + ] + "#; + let config: ConfigRequirementsToml = from_str(toml_str)?; + let requirements_toml_file = system_requirements_toml_file_for_test()?; + let source_location = RequirementSource::SystemRequirementsToml { + file: requirements_toml_file, + }; + + let mut requirements_with_sources = ConfigRequirementsWithSources::default(); + requirements_with_sources.merge_unset_fields(source_location.clone(), config); + let err = ConfigRequirements::try_from(requirements_with_sources) + .expect_err("invalid exec policy"); + + assert_eq!( + err, + ConstraintError::ExecPolicyParse { + requirement_source: source_location, + reason: "rules prefix_rule at index 0 is missing a decision".to_string(), + } + ); + + Ok(()) + } +} diff --git a/codex-rs/config/src/fingerprint.rs b/codex-rs/config/src/fingerprint.rs new file mode 100644 index 0000000000..d8e0263389 --- /dev/null +++ b/codex-rs/config/src/fingerprint.rs @@ -0,0 +1,67 @@ +use codex_app_server_protocol::ConfigLayerMetadata; +use serde_json::Value as JsonValue; +use sha2::Digest; +use sha2::Sha256; +use std::collections::HashMap; +use toml::Value as TomlValue; + +pub(super) fn record_origins( + value: &TomlValue, + meta: &ConfigLayerMetadata, + path: &mut Vec, + origins: &mut HashMap, +) { + match value { + TomlValue::Table(table) => { + for (key, val) in table { + path.push(key.clone()); + record_origins(val, meta, path, origins); + path.pop(); + } + } + TomlValue::Array(items) => { + for (idx, item) in (0_i32..).zip(items.iter()) { + path.push(idx.to_string()); + record_origins(item, meta, path, origins); + path.pop(); + } + } + _ => { + if !path.is_empty() { + origins.insert(path.join("."), meta.clone()); + } + } + } +} + +pub fn version_for_toml(value: &TomlValue) -> String { + let json = serde_json::to_value(value).unwrap_or(JsonValue::Null); + let canonical = canonical_json(&json); + let serialized = serde_json::to_vec(&canonical).unwrap_or_default(); + let mut hasher = Sha256::new(); + hasher.update(serialized); + let hash = hasher.finalize(); + let hex = hash + .iter() + .map(|byte| format!("{byte:02x}")) + .collect::(); + format!("sha256:{hex}") +} + +fn canonical_json(value: &JsonValue) -> JsonValue { + match value { + JsonValue::Object(map) => { + let mut sorted = serde_json::Map::new(); + let mut keys = map.keys().cloned().collect::>(); + keys.sort(); + for key in keys { + if let Some(val) = map.get(&key) { + sorted.insert(key, canonical_json(val)); + } + } + JsonValue::Object(sorted) + } + JsonValue::Array(items) => JsonValue::Array(items.iter().map(canonical_json).collect()), + other => other.clone(), + } +} diff --git a/codex-rs/config/src/merge.rs b/codex-rs/config/src/merge.rs new file mode 100644 index 0000000000..eacbbcb788 --- /dev/null +++ b/codex-rs/config/src/merge.rs @@ -0,0 +1,18 @@ +use toml::Value as TomlValue; + +/// Merge config `overlay` into `base`, giving `overlay` precedence. +pub fn merge_toml_values(base: &mut TomlValue, overlay: &TomlValue) { + if let TomlValue::Table(overlay_table) = overlay + && let TomlValue::Table(base_table) = base + { + for (key, value) in overlay_table { + if let Some(existing) = base_table.get_mut(key) { + merge_toml_values(existing, value); + } else { + base_table.insert(key.clone(), value.clone()); + } + } + } else { + *base = overlay.clone(); + } +} diff --git a/codex-rs/config/src/overrides.rs b/codex-rs/config/src/overrides.rs new file mode 100644 index 0000000000..a27caedb9a --- /dev/null +++ b/codex-rs/config/src/overrides.rs @@ -0,0 +1,55 @@ +use toml::Value as TomlValue; + +pub(crate) fn default_empty_table() -> TomlValue { + TomlValue::Table(Default::default()) +} + +pub fn build_cli_overrides_layer(cli_overrides: &[(String, TomlValue)]) -> TomlValue { + let mut root = default_empty_table(); + for (path, value) in cli_overrides { + apply_toml_override(&mut root, path, value.clone()); + } + root +} + +/// Apply a single dotted-path override onto a TOML value. +fn apply_toml_override(root: &mut TomlValue, path: &str, value: TomlValue) { + use toml::value::Table; + + let mut current = root; + let mut segments_iter = path.split('.').peekable(); + + while let Some(segment) = segments_iter.next() { + let is_last = segments_iter.peek().is_none(); + + if is_last { + match current { + TomlValue::Table(table) => { + table.insert(segment.to_string(), value); + } + _ => { + let mut table = Table::new(); + table.insert(segment.to_string(), value); + *current = TomlValue::Table(table); + } + } + return; + } + + match current { + TomlValue::Table(table) => { + current = table + .entry(segment.to_string()) + .or_insert_with(|| TomlValue::Table(Table::new())); + } + _ => { + *current = TomlValue::Table(Table::new()); + if let TomlValue::Table(tbl) = current { + current = tbl + .entry(segment.to_string()) + .or_insert_with(|| TomlValue::Table(Table::new())); + } + } + } + } +} diff --git a/codex-rs/config/src/requirements_exec_policy.rs b/codex-rs/config/src/requirements_exec_policy.rs new file mode 100644 index 0000000000..64d60f8814 --- /dev/null +++ b/codex-rs/config/src/requirements_exec_policy.rs @@ -0,0 +1,236 @@ +use codex_execpolicy::Decision; +use codex_execpolicy::Policy; +use codex_execpolicy::rule::PatternToken; +use codex_execpolicy::rule::PrefixPattern; +use codex_execpolicy::rule::PrefixRule; +use codex_execpolicy::rule::RuleRef; +use multimap::MultiMap; +use serde::Deserialize; +use std::sync::Arc; +use thiserror::Error; + +#[derive(Debug, Clone)] +pub struct RequirementsExecPolicy { + policy: Policy, +} + +impl RequirementsExecPolicy { + pub fn new(policy: Policy) -> Self { + Self { policy } + } +} + +impl PartialEq for RequirementsExecPolicy { + fn eq(&self, other: &Self) -> bool { + policy_fingerprint(&self.policy) == policy_fingerprint(&other.policy) + } +} + +impl Eq for RequirementsExecPolicy {} + +impl AsRef for RequirementsExecPolicy { + fn as_ref(&self) -> &Policy { + &self.policy + } +} + +fn policy_fingerprint(policy: &Policy) -> Vec { + let mut entries = Vec::new(); + for (program, rules) in policy.rules().iter_all() { + for rule in rules { + entries.push(format!("{program}:{rule:?}")); + } + } + entries.sort(); + entries +} + +/// TOML representation of `[rules]` within `requirements.toml`. +#[derive(Debug, Clone, PartialEq, Eq, Deserialize)] +pub struct RequirementsExecPolicyToml { + pub prefix_rules: Vec, +} + +/// A TOML representation of the `prefix_rule(...)` Starlark builtin. +/// +/// This mirrors the builtin defined in `execpolicy/src/parser.rs`. +#[derive(Debug, Clone, PartialEq, Eq, Deserialize)] +pub struct RequirementsExecPolicyPrefixRuleToml { + pub pattern: Vec, + pub decision: Option, + pub justification: Option, +} + +/// TOML-friendly representation of a pattern token. +/// +/// Starlark supports either a string token or a list of alternative tokens at +/// each position, but TOML arrays cannot mix strings and arrays. Using an +/// array of tables sidesteps that restriction. +#[derive(Debug, Clone, PartialEq, Eq, Deserialize)] +pub struct RequirementsExecPolicyPatternTokenToml { + pub token: Option, + pub any_of: Option>, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum RequirementsExecPolicyDecisionToml { + Allow, + Prompt, + Forbidden, +} + +impl RequirementsExecPolicyDecisionToml { + fn as_decision(self) -> Decision { + match self { + Self::Allow => Decision::Allow, + Self::Prompt => Decision::Prompt, + Self::Forbidden => Decision::Forbidden, + } + } +} + +#[derive(Debug, Error)] +pub enum RequirementsExecPolicyParseError { + #[error("rules prefix_rules cannot be empty")] + EmptyPrefixRules, + + #[error("rules prefix_rule at index {rule_index} has an empty pattern")] + EmptyPattern { rule_index: usize }, + + #[error( + "rules prefix_rule at index {rule_index} has an invalid pattern token at index {token_index}: {reason}" + )] + InvalidPatternToken { + rule_index: usize, + token_index: usize, + reason: String, + }, + + #[error("rules prefix_rule at index {rule_index} has an empty justification")] + EmptyJustification { rule_index: usize }, + + #[error("rules prefix_rule at index {rule_index} is missing a decision")] + MissingDecision { rule_index: usize }, + + #[error( + "rules prefix_rule at index {rule_index} has decision 'allow', which is not permitted in requirements.toml: Codex merges these rules with other config and uses the most restrictive result (use 'prompt' or 'forbidden')" + )] + AllowDecisionNotAllowed { rule_index: usize }, +} + +impl RequirementsExecPolicyToml { + /// Convert requirements TOML rules into the internal `.rules` + /// representation used by `codex-execpolicy`. + pub fn to_policy(&self) -> Result { + if self.prefix_rules.is_empty() { + return Err(RequirementsExecPolicyParseError::EmptyPrefixRules); + } + + let mut rules_by_program: MultiMap = MultiMap::new(); + + for (rule_index, rule) in self.prefix_rules.iter().enumerate() { + if let Some(justification) = &rule.justification + && justification.trim().is_empty() + { + return Err(RequirementsExecPolicyParseError::EmptyJustification { rule_index }); + } + + if rule.pattern.is_empty() { + return Err(RequirementsExecPolicyParseError::EmptyPattern { rule_index }); + } + + let pattern_tokens = rule + .pattern + .iter() + .enumerate() + .map(|(token_index, token)| parse_pattern_token(token, rule_index, token_index)) + .collect::, _>>()?; + + let decision = match rule.decision { + Some(RequirementsExecPolicyDecisionToml::Allow) => { + return Err(RequirementsExecPolicyParseError::AllowDecisionNotAllowed { + rule_index, + }); + } + Some(decision) => decision.as_decision(), + None => { + return Err(RequirementsExecPolicyParseError::MissingDecision { rule_index }); + } + }; + let justification = rule.justification.clone(); + + let (first_token, remaining_tokens) = pattern_tokens + .split_first() + .ok_or(RequirementsExecPolicyParseError::EmptyPattern { rule_index })?; + + let rest: Arc<[PatternToken]> = remaining_tokens.to_vec().into(); + + for head in first_token.alternatives() { + let rule: RuleRef = Arc::new(PrefixRule { + pattern: PrefixPattern { + first: Arc::from(head.as_str()), + rest: rest.clone(), + }, + decision, + justification: justification.clone(), + }); + rules_by_program.insert(head.clone(), rule); + } + } + + Ok(Policy::new(rules_by_program)) + } + + pub(crate) fn to_requirements_policy( + &self, + ) -> Result { + self.to_policy().map(RequirementsExecPolicy::new) + } +} + +fn parse_pattern_token( + token: &RequirementsExecPolicyPatternTokenToml, + rule_index: usize, + token_index: usize, +) -> Result { + match (&token.token, &token.any_of) { + (Some(single), None) => { + if single.trim().is_empty() { + return Err(RequirementsExecPolicyParseError::InvalidPatternToken { + rule_index, + token_index, + reason: "token cannot be empty".to_string(), + }); + } + Ok(PatternToken::Single(single.clone())) + } + (None, Some(alternatives)) => { + if alternatives.is_empty() { + return Err(RequirementsExecPolicyParseError::InvalidPatternToken { + rule_index, + token_index, + reason: "any_of cannot be empty".to_string(), + }); + } + if alternatives.iter().any(|alt| alt.trim().is_empty()) { + return Err(RequirementsExecPolicyParseError::InvalidPatternToken { + rule_index, + token_index, + reason: "any_of cannot include empty tokens".to_string(), + }); + } + Ok(PatternToken::Alts(alternatives.clone())) + } + (Some(_), Some(_)) => Err(RequirementsExecPolicyParseError::InvalidPatternToken { + rule_index, + token_index, + reason: "set either token or any_of, not both".to_string(), + }), + (None, None) => Err(RequirementsExecPolicyParseError::InvalidPatternToken { + rule_index, + token_index, + reason: "set either token or any_of".to_string(), + }), + } +} diff --git a/codex-rs/config/src/state.rs b/codex-rs/config/src/state.rs new file mode 100644 index 0000000000..30afecc7ca --- /dev/null +++ b/codex-rs/config/src/state.rs @@ -0,0 +1,311 @@ +use crate::config_loader::ConfigRequirements; +use crate::config_loader::ConfigRequirementsToml; + +use super::fingerprint::record_origins; +use super::fingerprint::version_for_toml; +use super::merge::merge_toml_values; +use codex_app_server_protocol::ConfigLayer; +use codex_app_server_protocol::ConfigLayerMetadata; +use codex_app_server_protocol::ConfigLayerSource; +use codex_utils_absolute_path::AbsolutePathBuf; +use serde_json::Value as JsonValue; +use std::collections::HashMap; +use std::path::PathBuf; +use toml::Value as TomlValue; + +/// LoaderOverrides overrides managed configuration inputs (primarily for tests). +#[derive(Debug, Default, Clone)] +pub struct LoaderOverrides { + pub managed_config_path: Option, + //TODO(gt): Add a macos_ prefix to this field and remove the target_os check. + #[cfg(target_os = "macos")] + pub managed_preferences_base64: Option, + pub macos_managed_config_requirements_base64: Option, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct ConfigLayerEntry { + pub name: ConfigLayerSource, + pub config: TomlValue, + pub raw_toml: Option, + pub version: String, + pub disabled_reason: Option, +} + +impl ConfigLayerEntry { + pub fn new(name: ConfigLayerSource, config: TomlValue) -> Self { + let version = version_for_toml(&config); + Self { + name, + config, + raw_toml: None, + version, + disabled_reason: None, + } + } + + pub fn new_with_raw_toml(name: ConfigLayerSource, config: TomlValue, raw_toml: String) -> Self { + let version = version_for_toml(&config); + Self { + name, + config, + raw_toml: Some(raw_toml), + version, + disabled_reason: None, + } + } + + pub fn new_disabled( + name: ConfigLayerSource, + config: TomlValue, + disabled_reason: impl Into, + ) -> Self { + let version = version_for_toml(&config); + Self { + name, + config, + raw_toml: None, + version, + disabled_reason: Some(disabled_reason.into()), + } + } + + pub fn is_disabled(&self) -> bool { + self.disabled_reason.is_some() + } + + pub fn raw_toml(&self) -> Option<&str> { + self.raw_toml.as_deref() + } + + pub fn metadata(&self) -> ConfigLayerMetadata { + ConfigLayerMetadata { + name: self.name.clone(), + version: self.version.clone(), + } + } + + pub fn as_layer(&self) -> ConfigLayer { + ConfigLayer { + name: self.name.clone(), + version: self.version.clone(), + config: serde_json::to_value(&self.config).unwrap_or(JsonValue::Null), + disabled_reason: self.disabled_reason.clone(), + } + } + + // Get the `.codex/` folder associated with this config layer, if any. + pub fn config_folder(&self) -> Option { + match &self.name { + ConfigLayerSource::Mdm { .. } => None, + ConfigLayerSource::System { file } => file.parent(), + ConfigLayerSource::User { file } => file.parent(), + ConfigLayerSource::Project { dot_codex_folder } => Some(dot_codex_folder.clone()), + ConfigLayerSource::SessionFlags => None, + ConfigLayerSource::LegacyManagedConfigTomlFromFile { .. } => None, + ConfigLayerSource::LegacyManagedConfigTomlFromMdm => None, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ConfigLayerStackOrdering { + LowestPrecedenceFirst, + HighestPrecedenceFirst, +} + +#[derive(Debug, Clone, Default, PartialEq)] +pub struct ConfigLayerStack { + /// Layers are listed from lowest precedence (base) to highest (top), so + /// later entries in the Vec override earlier ones. + layers: Vec, + + /// Index into [layers] of the user config layer, if any. + user_layer_index: Option, + + /// Constraints that must be enforced when deriving a [Config] from the + /// layers. + requirements: ConfigRequirements, + + /// Raw requirements data as loaded from requirements.toml/MDM/legacy + /// sources. This preserves the original allow-lists so they can be + /// surfaced via APIs. + requirements_toml: ConfigRequirementsToml, +} + +impl ConfigLayerStack { + pub fn new( + layers: Vec, + requirements: ConfigRequirements, + requirements_toml: ConfigRequirementsToml, + ) -> std::io::Result { + let user_layer_index = verify_layer_ordering(&layers)?; + Ok(Self { + layers, + user_layer_index, + requirements, + requirements_toml, + }) + } + + /// Returns the user config layer, if any. + pub fn get_user_layer(&self) -> Option<&ConfigLayerEntry> { + self.user_layer_index + .and_then(|index| self.layers.get(index)) + } + + pub fn requirements(&self) -> &ConfigRequirements { + &self.requirements + } + + pub fn requirements_toml(&self) -> &ConfigRequirementsToml { + &self.requirements_toml + } + + /// Creates a new [ConfigLayerStack] using the specified values to inject a + /// "user layer" into the stack. If such a layer already exists, it is + /// replaced; otherwise, it is inserted into the stack at the appropriate + /// position based on precedence rules. + pub fn with_user_config(&self, config_toml: &AbsolutePathBuf, user_config: TomlValue) -> Self { + let user_layer = ConfigLayerEntry::new( + ConfigLayerSource::User { + file: config_toml.clone(), + }, + user_config, + ); + + let mut layers = self.layers.clone(); + match self.user_layer_index { + Some(index) => { + layers[index] = user_layer; + Self { + layers, + user_layer_index: self.user_layer_index, + requirements: self.requirements.clone(), + requirements_toml: self.requirements_toml.clone(), + } + } + None => { + let user_layer_index = match layers + .iter() + .position(|layer| layer.name.precedence() > user_layer.name.precedence()) + { + Some(index) => { + layers.insert(index, user_layer); + index + } + None => { + layers.push(user_layer); + layers.len() - 1 + } + }; + Self { + layers, + user_layer_index: Some(user_layer_index), + requirements: self.requirements.clone(), + requirements_toml: self.requirements_toml.clone(), + } + } + } + } + + pub fn effective_config(&self) -> TomlValue { + let mut merged = TomlValue::Table(toml::map::Map::new()); + for layer in self.get_layers(ConfigLayerStackOrdering::LowestPrecedenceFirst, false) { + merge_toml_values(&mut merged, &layer.config); + } + merged + } + + pub fn origins(&self) -> HashMap { + let mut origins = HashMap::new(); + let mut path = Vec::new(); + + for layer in self.get_layers(ConfigLayerStackOrdering::LowestPrecedenceFirst, false) { + record_origins(&layer.config, &layer.metadata(), &mut path, &mut origins); + } + + origins + } + + /// Returns the highest-precedence to lowest-precedence layers, so + /// `ConfigLayerSource::SessionFlags` would be first, if present. + pub fn layers_high_to_low(&self) -> Vec<&ConfigLayerEntry> { + self.get_layers(ConfigLayerStackOrdering::HighestPrecedenceFirst, false) + } + + /// Returns the highest-precedence to lowest-precedence layers, so + /// `ConfigLayerSource::SessionFlags` would be first, if present. + pub fn get_layers( + &self, + ordering: ConfigLayerStackOrdering, + include_disabled: bool, + ) -> Vec<&ConfigLayerEntry> { + let mut layers: Vec<&ConfigLayerEntry> = self + .layers + .iter() + .filter(|layer| include_disabled || !layer.is_disabled()) + .collect(); + if ordering == ConfigLayerStackOrdering::HighestPrecedenceFirst { + layers.reverse(); + } + layers + } +} + +/// Ensures precedence ordering of config layers is correct. Returns the index +/// of the user config layer, if any (at most one should exist). +fn verify_layer_ordering(layers: &[ConfigLayerEntry]) -> std::io::Result> { + if !layers.iter().map(|layer| &layer.name).is_sorted() { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "config layers are not in correct precedence order", + )); + } + + // The previous check ensured `layers` is sorted by precedence, so now we + // further verify that: + // 1. There is at most one user config layer. + // 2. Project layers are ordered from root to cwd. + let mut user_layer_index: Option = None; + let mut previous_project_dot_codex_folder: Option<&AbsolutePathBuf> = None; + for (index, layer) in layers.iter().enumerate() { + if matches!(layer.name, ConfigLayerSource::User { .. }) { + if user_layer_index.is_some() { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "multiple user config layers found", + )); + } + user_layer_index = Some(index); + } + + if let ConfigLayerSource::Project { + dot_codex_folder: current_project_dot_codex_folder, + } = &layer.name + { + if let Some(previous) = previous_project_dot_codex_folder { + let Some(parent) = previous.as_path().parent() else { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "project layer has no parent directory", + )); + }; + if previous == current_project_dot_codex_folder + || !current_project_dot_codex_folder + .as_path() + .ancestors() + .any(|ancestor| ancestor == parent) + { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "project layers are not ordered from root to cwd", + )); + } + } + previous_project_dot_codex_folder = Some(current_project_dot_codex_folder); + } + } + + Ok(user_layer_index) +}