diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index bc3610127..f34161c69 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -3553,6 +3553,7 @@ dependencies = [ "rama-tcp", "rama-tls-rustls", "rama-unix", + "rand 0.9.3", "rustls-native-certs", "schannel", "security-framework 3.5.1", diff --git a/codex-rs/network-proxy/Cargo.toml b/codex-rs/network-proxy/Cargo.toml index cc6a8de71..d1d2a72ce 100644 --- a/codex-rs/network-proxy/Cargo.toml +++ b/codex-rs/network-proxy/Cargo.toml @@ -21,6 +21,7 @@ codex-utils-absolute-path = { workspace = true } codex-utils-home-dir = { workspace = true } codex-utils-rustls-provider = { workspace = true } globset = { workspace = true } +rand = { workspace = true } serde = { workspace = true, features = ["derive"] } serde_json = { workspace = true } thiserror = { workspace = true } diff --git a/codex-rs/network-proxy/src/config.rs b/codex-rs/network-proxy/src/config.rs index d9cef46b1..30cdfa7b2 100644 --- a/codex-rs/network-proxy/src/config.rs +++ b/codex-rs/network-proxy/src/config.rs @@ -21,6 +21,13 @@ pub struct NetworkProxyConfig { pub network: NetworkProxySettings, } +impl NetworkProxyConfig { + pub fn set_credential_broker_enabled(&mut self, enabled: bool) { + self.network.credential_broker = enabled; + self.network.mitm |= enabled; + } +} + /// Variant order encodes effective precedence for duplicate patterns: /// `None < Allow < Deny`, so deny wins over allow when entries conflict. #[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, PartialOrd, Ord)] @@ -142,6 +149,10 @@ pub struct NetworkProxySettings { #[serde(default)] pub mitm: bool, #[serde(default)] + pub credential_broker: bool, + #[serde(default)] + pub dangerously_allow_plaintext_credential_injection: bool, + #[serde(default)] pub mitm_hooks: Vec, } @@ -161,6 +172,8 @@ impl Default for NetworkProxySettings { unix_sockets: None, allow_local_binding: false, mitm: false, + credential_broker: false, + dangerously_allow_plaintext_credential_injection: false, mitm_hooks: Vec::new(), } } @@ -593,6 +606,8 @@ mod tests { unix_sockets: None, allow_local_binding: false, mitm: false, + credential_broker: false, + dangerously_allow_plaintext_credential_injection: false, mitm_hooks: Vec::new(), } ); @@ -658,6 +673,8 @@ mod tests { "unix_sockets": null, "allow_local_binding": false, "mitm": false, + "credential_broker": false, + "dangerously_allow_plaintext_credential_injection": false, "mitm_hooks": [], } }) diff --git a/codex-rs/network-proxy/src/credential_broker.rs b/codex-rs/network-proxy/src/credential_broker.rs new file mode 100644 index 000000000..379c45e4c --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker.rs @@ -0,0 +1,270 @@ +mod providers; + +use crate::policy::normalize_host; +use rama_http::HeaderMap; +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::RwLock; + +pub const CREDENTIAL_BROKER_ACTIVE_ENV_KEY: &str = "CODEX_NETWORK_PROXY_CREDENTIAL_BROKER_ACTIVE"; +pub(crate) const BROKERED_CREDENTIALS_ENV_KEY: &str = "CODEX_NETWORK_PROXY_BROKERED_CREDENTIALS"; + +#[derive(Clone)] +pub(crate) struct CredentialBroker { + state: Arc>, +} + +#[derive(Default)] +struct CredentialBrokerState { + enabled: bool, + credentials: Vec, +} + +struct CredentialRecord { + env_var: String, + provider: &'static providers::CredentialProvider, + host_binding: providers::CredentialHostBinding, + real_value: String, + dummy_value: String, +} + +impl CredentialBroker { + pub(crate) fn new(enabled: bool) -> Self { + Self { + state: Arc::new(RwLock::new(CredentialBrokerState { + enabled, + ..CredentialBrokerState::default() + })), + } + } + + pub(crate) fn enabled(&self) -> bool { + self.read_state().enabled + } + + pub(crate) fn virtualize_child_env(&self, env: &mut HashMap) { + let mut state = self.write_state(); + if !state.enabled { + env.remove(CREDENTIAL_BROKER_ACTIVE_ENV_KEY); + env.remove(BROKERED_CREDENTIALS_ENV_KEY); + return; + } + env.insert( + CREDENTIAL_BROKER_ACTIVE_ENV_KEY.to_string(), + "1".to_string(), + ); + + for provider in providers::credential_providers() { + for source in provider.sources() { + if let Some(host_binding) = (source.host_binding)(env) { + for env_var in source.env_vars { + virtualize_env_var( + env, + &mut state, + env_var, + provider, + host_binding.clone(), + ); + } + } + } + } + update_brokered_credentials_marker(&state, env); + } + + pub(crate) fn host_requires_mitm(&self, host: &str) -> bool { + let normalized_host = normalize_host(host); + let state = self.read_state(); + state.enabled + && state + .credentials + .iter() + .any(|credential| credential.matches_host(&normalized_host)) + } + + pub(crate) fn inject_request_headers(&self, host: &str, headers: &mut HeaderMap) { + let normalized_host = normalize_host(host); + let state = self.read_state(); + if !state.enabled { + return; + } + + let matching_credentials = state + .credentials + .iter() + .filter(|credential| credential.matches_host(&normalized_host)) + .collect::>(); + let Some(credential) = select_credential(headers, &matching_credentials) else { + return; + }; + let Some(header_value) = credential + .provider + .request_header_value(&credential.real_value) + else { + return; + }; + credential + .provider + .insert_request_header(headers, header_value); + } + + fn read_state(&self) -> std::sync::RwLockReadGuard<'_, CredentialBrokerState> { + self.state + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + } + + fn write_state(&self) -> std::sync::RwLockWriteGuard<'_, CredentialBrokerState> { + self.state + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) + } +} + +fn virtualize_env_var( + env: &mut HashMap, + state: &mut CredentialBrokerState, + env_var: &str, + provider: &'static providers::CredentialProvider, + host_binding: providers::CredentialHostBinding, +) { + let Some(real_value) = brokerable_credential_value(env, state, env_var, provider) else { + return; + }; + + let dummy_value = state.register(env_var, provider, host_binding, real_value); + env.insert(env_var.to_string(), dummy_value); +} + +fn brokerable_credential_value<'a>( + env: &'a HashMap, + state: &CredentialBrokerState, + env_var: &str, + provider: &providers::CredentialProvider, +) -> Option<&'a str> { + let real_value = env.get(env_var)?.trim(); + (!real_value.is_empty() + && !state.is_dummy_value(real_value) + && provider.request_header_value(real_value).is_some()) + .then_some(real_value) +} + +impl CredentialBrokerState { + fn register( + &mut self, + env_var: &str, + provider: &'static providers::CredentialProvider, + host_binding: providers::CredentialHostBinding, + real_value: &str, + ) -> String { + if let Some(existing) = self.credentials.iter().find(|credential| { + credential.env_var == env_var + && std::ptr::eq(credential.provider, provider) + && credential.host_binding == host_binding + && credential.real_value == real_value + }) { + return existing.dummy_value.clone(); + } + + let dummy_value = loop { + let candidate = provider.dummy_value(real_value); + if candidate != real_value && !self.is_dummy_value(&candidate) { + break candidate; + } + }; + self.credentials.push(CredentialRecord { + env_var: env_var.to_string(), + provider, + host_binding, + real_value: real_value.to_string(), + dummy_value: dummy_value.clone(), + }); + dummy_value + } + + fn is_dummy_value(&self, value: &str) -> bool { + self.credentials + .iter() + .any(|credential| credential.dummy_value == value) + } +} + +impl CredentialRecord { + fn matches_host(&self, host: &str) -> bool { + self.host_binding.matches_host(host) + } +} + +fn select_credential<'a>( + headers: &HeaderMap, + matching_credentials: &[&'a CredentialRecord], +) -> Option<&'a CredentialRecord> { + let dummy_matches = matching_credentials + .iter() + .copied() + .filter(|credential| { + credential + .provider + .request_header(headers) + .and_then(|value| value.to_str().ok()) + .is_some_and(|value| value.contains(&credential.dummy_value)) + }) + .collect::>(); + match dummy_matches.as_slice() { + [credential] => Some(*credential), + [] | [_, _, ..] => None, + } +} + +fn update_brokered_credentials_marker( + state: &CredentialBrokerState, + env: &mut HashMap, +) { + let brokered = providers::credential_broker_env_keys() + .filter_map(|key| { + let value = env.get(key)?; + state.is_dummy_value(value).then_some((key, value.as_str())) + }) + .collect::>(); + match serde_json::to_string(&brokered) { + Ok(marker) => { + env.insert(BROKERED_CREDENTIALS_ENV_KEY.to_string(), marker); + } + Err(_) => { + env.remove(BROKERED_CREDENTIALS_ENV_KEY); + } + } +} + +/// Returns supported environment keys whose current values still match the child-scoped dummy +/// values recorded by the credential broker. +/// +/// The broker marker is treated as untrusted: malformed metadata, unsupported keys, and values +/// replaced by the user are ignored. The environment is not mutated; callers own the decision to +/// remove the returned keys. +pub fn brokered_credential_dummy_env_keys(env: &HashMap) -> Vec { + env.get(BROKERED_CREDENTIALS_ENV_KEY) + .and_then(|marker| serde_json::from_str::>(marker).ok()) + .unwrap_or_default() + .into_iter() + .filter_map(|(key, dummy_value)| { + (providers::credential_broker_env_keys().any(|candidate| candidate == key.as_str()) + && env.get(&key) == Some(&dummy_value)) + .then_some(key) + }) + .collect() +} + +/// Returns supported credential keys only for an environment with an active broker. +pub fn brokered_credential_env_keys( + env: &HashMap, +) -> impl Iterator { + let active = env + .get(CREDENTIAL_BROKER_ACTIVE_ENV_KEY) + .is_some_and(|value| value == "1"); + providers::credential_broker_env_keys().filter(move |_| active) +} + +#[cfg(test)] +#[path = "credential_broker_tests.rs"] +mod tests; diff --git a/codex-rs/network-proxy/src/credential_broker/providers.rs b/codex-rs/network-proxy/src/credential_broker/providers.rs new file mode 100644 index 000000000..94c4b6e37 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/providers.rs @@ -0,0 +1,105 @@ +mod github; +mod openai; + +use rama_http::HeaderMap; +use rama_http::HeaderValue; +use rand::Rng as _; +use std::collections::HashMap; + +const DUMMY_ALPHANUMERIC: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"; + +type RequestHeader = for<'a> fn(&'a HeaderMap) -> Option<&'a HeaderValue>; + +/// Describes how one credential family is recognized and injected. +/// +/// Providers must be declared as `static` values because the broker uses their addresses as stable +/// identities when deduplicating credential records. +pub(super) struct CredentialProvider { + context_env_vars: &'static [&'static str], + sources: &'static [CredentialSource], + dummy_value: fn(&str) -> String, + request_header: RequestHeader, + request_header_value: fn(&str) -> Option, + insert_request_header: fn(&mut HeaderMap, HeaderValue), +} + +#[derive(Clone, PartialEq, Eq)] +pub(super) enum CredentialHostBinding { + ExactHost(String), + HostPattern { + exact_hosts: &'static [&'static str], + suffixes: &'static [&'static str], + }, +} + +pub(super) struct CredentialSource { + pub(super) env_vars: &'static [&'static str], + pub(super) host_binding: fn(&HashMap) -> Option, +} + +const CREDENTIAL_PROVIDERS: &[&CredentialProvider] = &[&github::PROVIDER, &openai::PROVIDER]; + +impl CredentialProvider { + pub(super) fn sources(&self) -> &[CredentialSource] { + self.sources + } + + pub(super) fn dummy_value(&self, real_value: &str) -> String { + (self.dummy_value)(real_value) + } + + pub(super) fn request_header<'a>(&self, headers: &'a HeaderMap) -> Option<&'a HeaderValue> { + (self.request_header)(headers) + } + + pub(super) fn request_header_value(&self, value: &str) -> Option { + (self.request_header_value)(value) + } + + pub(super) fn insert_request_header(&self, headers: &mut HeaderMap, value: HeaderValue) { + (self.insert_request_header)(headers, value); + } +} + +impl CredentialHostBinding { + pub(super) fn matches_host(&self, host: &str) -> bool { + match self { + Self::ExactHost(expected_host) => host == expected_host, + Self::HostPattern { + exact_hosts, + suffixes, + } => { + exact_hosts.contains(&host) || suffixes.iter().any(|suffix| host.ends_with(suffix)) + } + } + } +} + +pub(super) fn credential_broker_env_keys() -> impl Iterator { + credential_providers() + .flat_map(|provider| provider.context_env_vars.iter().copied()) + .chain( + credential_providers() + .flat_map(CredentialProvider::sources) + .flat_map(|source| source.env_vars.iter().copied()), + ) +} + +pub(super) fn credential_providers() -> impl Iterator { + CREDENTIAL_PROVIDERS.iter().copied() +} + +fn shaped_dummy_value(real_value: &str, prefix: &str, minimum_len: usize) -> String { + let target_len = real_value.len().max(minimum_len).max(prefix.len() + 16); + let mut rng = rand::rng(); + let mut dummy = String::with_capacity(target_len); + dummy.push_str(prefix); + for index in prefix.len()..target_len { + let character = match real_value.as_bytes().get(index).copied() { + Some(template) if !template.is_ascii_alphanumeric() => template, + _ => DUMMY_ALPHANUMERIC[rng.random_range(0..DUMMY_ALPHANUMERIC.len())], + }; + dummy.push(char::from(character)); + } + dummy +} diff --git a/codex-rs/network-proxy/src/credential_broker/providers/github.rs b/codex-rs/network-proxy/src/credential_broker/providers/github.rs new file mode 100644 index 000000000..ced353ee7 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/providers/github.rs @@ -0,0 +1,91 @@ +use super::CredentialHostBinding; +use super::CredentialProvider; +use super::CredentialSource; +use super::shaped_dummy_value; +use crate::policy::normalize_host; +use rama_http::HeaderMap; +use rama_http::HeaderValue; +use rama_http::header::AUTHORIZATION; +use std::collections::HashMap; + +const GH_HOST_ENV_VAR: &str = "GH_HOST"; +const GITHUB_TOKEN_PREFIXES: &[&str] = &["github_pat_", "ghp_", "gho_", "ghu_", "ghs_", "ghr_"]; +const GITHUB_TOKEN_MIN_LEN: usize = 40; +const GITHUB_CLOUD_TOKEN_ENV_VARS: &[&str] = &["GH_TOKEN", "GITHUB_TOKEN"]; +const GITHUB_ENTERPRISE_TOKEN_ENV_VARS: &[&str] = + &["GH_ENTERPRISE_TOKEN", "GITHUB_ENTERPRISE_TOKEN"]; +const GITHUB_CLOUD_HOSTS: &[&str] = &["api.github.com", "github.com"]; +const GITHUB_CLOUD_HOST_SUFFIXES: &[&str] = &[".ghe.com"]; + +pub(super) static PROVIDER: CredentialProvider = CredentialProvider { + context_env_vars: &[GH_HOST_ENV_VAR], + sources: &[ + CredentialSource { + env_vars: GITHUB_CLOUD_TOKEN_ENV_VARS, + host_binding: github_cloud_binding, + }, + CredentialSource { + env_vars: GITHUB_ENTERPRISE_TOKEN_ENV_VARS, + host_binding: github_enterprise_binding, + }, + ], + dummy_value, + request_header, + request_header_value, + insert_request_header, +}; + +fn dummy_value(real_value: &str) -> String { + shaped_dummy_value( + real_value, + github_token_prefix(real_value), + GITHUB_TOKEN_MIN_LEN, + ) +} + +fn request_header(headers: &HeaderMap) -> Option<&HeaderValue> { + headers.get(AUTHORIZATION) +} + +fn request_header_value(value: &str) -> Option { + HeaderValue::from_str(&format!("Bearer {value}")).ok() +} + +fn insert_request_header(headers: &mut HeaderMap, value: HeaderValue) { + headers.insert(AUTHORIZATION, value); +} + +fn github_cloud_binding(_: &HashMap) -> Option { + Some(CredentialHostBinding::HostPattern { + exact_hosts: GITHUB_CLOUD_HOSTS, + suffixes: GITHUB_CLOUD_HOST_SUFFIXES, + }) +} + +fn github_enterprise_binding(env: &HashMap) -> Option { + github_host_hint(env) + .filter(|host| !github_cloud_host(host)) + .map(CredentialHostBinding::ExactHost) +} + +fn github_cloud_host(host: &str) -> bool { + GITHUB_CLOUD_HOSTS.contains(&host) + || GITHUB_CLOUD_HOST_SUFFIXES + .iter() + .any(|suffix| host.ends_with(suffix)) +} + +fn github_token_prefix(value: &str) -> &str { + GITHUB_TOKEN_PREFIXES + .iter() + .copied() + .find(|prefix| value.starts_with(prefix)) + .unwrap_or("ghp_") +} + +fn github_host_hint(env: &HashMap) -> Option { + env.get(GH_HOST_ENV_VAR) + .map(String::as_str) + .map(normalize_host) + .filter(|host| !host.is_empty()) +} diff --git a/codex-rs/network-proxy/src/credential_broker/providers/openai.rs b/codex-rs/network-proxy/src/credential_broker/providers/openai.rs new file mode 100644 index 000000000..c3af9261f --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/providers/openai.rs @@ -0,0 +1,59 @@ +use super::CredentialHostBinding; +use super::CredentialProvider; +use super::CredentialSource; +use super::shaped_dummy_value; +use rama_http::HeaderMap; +use rama_http::HeaderValue; +use rama_http::header::AUTHORIZATION; +use std::collections::HashMap; + +const OPENAI_API_KEY_ENV_VARS: &[&str] = &["OPENAI_API_KEY"]; +const OPENAI_API_KEY_MIN_LEN: usize = 51; +const OPENAI_API_HOST: &str = "api.openai.com"; + +pub(super) static PROVIDER: CredentialProvider = CredentialProvider { + context_env_vars: &[], + sources: &[CredentialSource { + env_vars: OPENAI_API_KEY_ENV_VARS, + host_binding, + }], + dummy_value, + request_header, + request_header_value, + insert_request_header, +}; + +fn dummy_value(real_value: &str) -> String { + shaped_dummy_value( + real_value, + openai_api_key_prefix(real_value), + OPENAI_API_KEY_MIN_LEN, + ) +} + +fn request_header(headers: &HeaderMap) -> Option<&HeaderValue> { + headers.get(AUTHORIZATION) +} + +fn request_header_value(value: &str) -> Option { + HeaderValue::from_str(&format!("Bearer {value}")).ok() +} + +fn insert_request_header(headers: &mut HeaderMap, value: HeaderValue) { + headers.insert(AUTHORIZATION, value); +} + +fn host_binding(_: &HashMap) -> Option { + Some(CredentialHostBinding::ExactHost( + OPENAI_API_HOST.to_string(), + )) +} + +fn openai_api_key_prefix(value: &str) -> &str { + let Some(suffix) = value.strip_prefix("sk-") else { + return "sk-"; + }; + suffix + .find('-') + .map_or("sk-", |separator| &value[..separator + 4]) +} diff --git a/codex-rs/network-proxy/src/credential_broker_tests.rs b/codex-rs/network-proxy/src/credential_broker_tests.rs new file mode 100644 index 000000000..186e89163 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker_tests.rs @@ -0,0 +1,220 @@ +use super::*; + +use pretty_assertions::assert_eq; +use rama_http::HeaderValue; +use rama_http::header::AUTHORIZATION; + +fn env_map(entries: [(&str, &str); N]) -> HashMap { + entries + .into_iter() + .map(|(key, value)| (key.to_string(), value.to_string())) + .collect() +} + +fn headers_with_bearer(value: &str) -> HeaderMap { + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {value}")).expect("valid bearer header"), + ); + headers +} + +fn authorization(headers: &HeaderMap) -> Option<&str> { + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()) +} + +fn assert_credential_shape(real_value: &str, dummy_value: &str, prefix: &str) { + assert_ne!(dummy_value, real_value); + assert_eq!(dummy_value.len(), real_value.len()); + assert_eq!(&dummy_value[..prefix.len()], prefix); + let same_shape = real_value + .bytes() + .zip(dummy_value.bytes()) + .skip(prefix.len()) + .all(|(real, dummy)| { + real.is_ascii_alphanumeric() && dummy.is_ascii_alphanumeric() || real == dummy + }); + assert!(same_shape); +} + +#[test] +fn virtualize_child_env_replaces_supported_credentials() { + let broker = CredentialBroker::new(/*enabled*/ true); + let github_token = "github_pat_11AA0bbCC_abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGH"; + let openai_api_key = "sk-proj-abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789-_"; + let mut env = env_map([ + ("GH_TOKEN", github_token), + ("OPENAI_API_KEY", openai_api_key), + ("GH_ENTERPRISE_TOKEN", "ghp-enterprise-real"), + ]); + + broker.virtualize_child_env(&mut env); + + let github_dummy = env.get("GH_TOKEN").expect("dummy GitHub token"); + let openai_dummy = env.get("OPENAI_API_KEY").expect("dummy OpenAI API key"); + assert_credential_shape(github_token, github_dummy, "github_pat_"); + assert_credential_shape(openai_api_key, openai_dummy, "sk-proj-"); + env.insert("OPENAI_API_KEY".to_string(), "sk-user-override".to_string()); + assert_eq!( + brokered_credential_dummy_env_keys(&env), + vec!["GH_TOKEN".to_string()] + ); +} + +#[test] +fn virtualize_child_env_preserves_live_dummy_mappings() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut first_env = env_map([("GH_TOKEN", "ghp-real-one")]); + let mut second_env = env_map([("GH_TOKEN", "ghp-real-two")]); + + broker.virtualize_child_env(&mut first_env); + broker.virtualize_child_env(&mut second_env); + let first_dummy = first_env.get("GH_TOKEN").expect("first dummy token"); + let second_dummy = second_env.get("GH_TOKEN").expect("second dummy token"); + let mut first_headers = headers_with_bearer(first_dummy); + let mut second_headers = headers_with_bearer(second_dummy); + + broker.inject_request_headers("api.github.com", &mut first_headers); + broker.inject_request_headers("api.github.com", &mut second_headers); + + assert_eq!(authorization(&first_headers), Some("Bearer ghp-real-one")); + assert_eq!(authorization(&second_headers), Some("Bearer ghp-real-two")); +} + +#[test] +fn virtualize_child_env_uses_fresh_dummy_capabilities() { + let mut first_env = env_map([("OPENAI_API_KEY", "sk-proj-abcdefghijklmnopqrstuvwxyz")]); + let mut second_env = first_env.clone(); + + CredentialBroker::new(/*enabled*/ true).virtualize_child_env(&mut first_env); + CredentialBroker::new(/*enabled*/ true).virtualize_child_env(&mut second_env); + + assert_ne!(first_env["OPENAI_API_KEY"], second_env["OPENAI_API_KEY"]); +} + +#[test] +fn child_without_dummy_cannot_use_previous_child_credential() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut first_env = env_map([("OPENAI_API_KEY", "sk-real")]); + let mut second_env = HashMap::new(); + + broker.virtualize_child_env(&mut first_env); + broker.virtualize_child_env(&mut second_env); + let mut headers = HeaderMap::new(); + + broker.inject_request_headers("api.openai.com", &mut headers); + + assert_eq!(authorization(&headers), None); +} + +#[test] +fn virtualize_child_env_preserves_unbound_enterprise_token() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([("GH_ENTERPRISE_TOKEN", "ghp-enterprise-real")]); + + broker.virtualize_child_env(&mut env); + let inert_token = "ghp_abcdefghijklmnopqrstuvwxyz1234567890"; + let mut headers = headers_with_bearer(inert_token); + broker.inject_request_headers("attacker.example", &mut headers); + + assert_eq!(env["GH_ENTERPRISE_TOKEN"], "ghp-enterprise-real"); + assert_eq!(headers, headers_with_bearer(inert_token)); + assert!(!broker.host_requires_mitm("attacker.example")); +} + +#[test] +fn inject_request_headers_requires_dummy_to_select_ambiguous_github_credential() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([ + ("GH_TOKEN", "ghp-real-one"), + ("GITHUB_TOKEN", "ghp-real-two"), + ]); + broker.virtualize_child_env(&mut env); + let github_token = env.get("GITHUB_TOKEN").expect("dummy github token"); + let mut headers = HeaderMap::new(); + + broker.inject_request_headers("api.github.com", &mut headers); + assert_eq!(authorization(&headers), None); + + headers = headers_with_bearer(github_token); + + broker.inject_request_headers("api.github.com", &mut headers); + + assert_eq!(authorization(&headers), Some("Bearer ghp-real-two")); +} + +#[test] +fn inject_request_headers_requires_dummy_and_preserves_explicit_authorization() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([("OPENAI_API_KEY", "sk-real")]); + broker.virtualize_child_env(&mut env); + let openai_api_key = env.get("OPENAI_API_KEY").expect("dummy OpenAI API key"); + let mut headers = HeaderMap::new(); + + broker.inject_request_headers("api.openai.com", &mut headers); + assert_eq!(authorization(&headers), None); + + headers = headers_with_bearer(openai_api_key); + broker.inject_request_headers("api.openai.com", &mut headers); + assert_eq!(authorization(&headers), Some("Bearer sk-real")); + + let mut explicit_headers = headers_with_bearer("sk-explicit"); + broker.inject_request_headers("api.openai.com", &mut explicit_headers); + + assert_eq!(authorization(&explicit_headers), Some("Bearer sk-explicit")); +} + +#[test] +fn github_cloud_credentials_match_ghe_com_host_hint() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([("GH_HOST", "astemu.ghe.com"), ("GH_TOKEN", "ghp-real")]); + broker.virtualize_child_env(&mut env); + let github_token = env.get("GH_TOKEN").expect("dummy GitHub token"); + let mut headers = headers_with_bearer(github_token); + + broker.inject_request_headers("api.astemu.ghe.com", &mut headers); + + assert_eq!(authorization(&headers), Some("Bearer ghp-real")); +} + +#[test] +fn github_cloud_credentials_do_not_bind_to_ghes_host_hint() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([("GH_HOST", "github.example.com"), ("GH_TOKEN", "ghp-real")]); + broker.virtualize_child_env(&mut env); + let github_token = env.get("GH_TOKEN").expect("dummy github token"); + let expected_authorization = format!("Bearer {github_token}"); + let mut headers = headers_with_bearer(github_token); + + broker.inject_request_headers("github.example.com", &mut headers); + + assert_eq!( + authorization(&headers), + Some(expected_authorization.as_str()) + ); + assert!(!broker.host_requires_mitm("github.example.com")); + assert!(broker.host_requires_mitm("api.github.com")); +} + +#[test] +fn github_enterprise_credentials_bind_to_gh_host() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([ + ("GH_HOST", "github.example.com"), + ("GH_ENTERPRISE_TOKEN", "ghp-enterprise-real"), + ]); + broker.virtualize_child_env(&mut env); + let github_token = env + .get("GH_ENTERPRISE_TOKEN") + .expect("dummy GitHub enterprise token"); + let mut headers = headers_with_bearer(github_token); + + broker.inject_request_headers("github.example.com", &mut headers); + + assert_eq!(authorization(&headers), Some("Bearer ghp-enterprise-real")); + assert!(broker.host_requires_mitm("github.example.com")); + assert!(!broker.host_requires_mitm("api.github.com")); +} diff --git a/codex-rs/network-proxy/src/http_proxy.rs b/codex-rs/network-proxy/src/http_proxy.rs index 3f02ad3fc..180ef41c5 100644 --- a/codex-rs/network-proxy/src/http_proxy.rs +++ b/codex-rs/network-proxy/src/http_proxy.rs @@ -23,6 +23,7 @@ use crate::responses::blocked_header_value; use crate::responses::blocked_message_with_policy; use crate::responses::blocked_text_response_with_policy; use crate::responses::json_response; +use crate::runtime::HostMitmRequirement; use crate::runtime::unix_socket_permissions_supported; use crate::state::BlockedRequest; use crate::state::BlockedRequestArgs; @@ -34,13 +35,13 @@ use anyhow::Result; use codex_utils_rustls_provider::ensure_rustls_crypto_provider; use rama_core::Layer; use rama_core::Service; -use rama_core::error::BoxError; use rama_core::error::ErrorExt as _; use rama_core::error::OpaqueError; use rama_core::extensions::ExtensionsMut; use rama_core::extensions::ExtensionsRef; use rama_core::layer::AddInputExtensionLayer; use rama_core::service::service_fn; +use rama_core::stream::Stream; use rama_http::Body; use rama_http::HeaderMap; use rama_http::HeaderName; @@ -58,7 +59,6 @@ use rama_http_backend::server::HttpServer; use rama_http_backend::server::layer::upgrade::UpgradeLayer; use rama_http_backend::server::layer::upgrade::Upgraded; use rama_net::Protocol; -use rama_net::address::ProxyAddress; use rama_net::client::ConnectorService; use rama_net::client::EstablishedClientConnection; use rama_net::http::RequestContext; @@ -80,8 +80,12 @@ use tracing::error; use tracing::info; use tracing::warn; -#[derive(Clone, Copy, Debug)] -struct ConnectMitmEnabled(bool); +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum ConnectMitmMode { + Disabled, + Enabled, + DetectTls, +} pub async fn run_http_proxy( state: Arc, @@ -268,18 +272,26 @@ async fn http_connect_accept( return Err(text_response(StatusCode::INTERNAL_SERVER_ERROR, "error")); } }; - let host_has_mitm_hooks = match app_state.host_has_mitm_hooks(&host).await { - Ok(has_hooks) => has_hooks, + let host_mitm_requirement = match app_state.host_mitm_requirement(&host).await { + Ok(requirement) => requirement, Err(err) => { - error!("failed to inspect MITM hooks for {host}: {err}"); + error!("failed to inspect MITM requirements for {host}: {err}"); return Err(text_response(StatusCode::INTERNAL_SERVER_ERROR, "error")); } }; - let connect_needs_mitm = mode == NetworkMode::Limited || host_has_mitm_hooks; + let connect_mitm_mode = if mode == NetworkMode::Limited { + ConnectMitmMode::Enabled + } else { + match host_mitm_requirement { + HostMitmRequirement::None => ConnectMitmMode::Disabled, + HostMitmRequirement::Tls => ConnectMitmMode::DetectTls, + HostMitmRequirement::Always => ConnectMitmMode::Enabled, + } + }; - if connect_needs_mitm && mitm_state.is_none() { - // CONNECT needs MITM whenever HTTPS policy depends on inner-request inspection, either for - // limited-mode method enforcement or for host-specific MITM hooks. + if connect_mitm_mode == ConnectMitmMode::Enabled && mitm_state.is_none() { + // Limited-mode enforcement and host-specific hooks require interception. Credential-only + // interception is deferred until the upgraded stream presents a TLS ClientHello. emit_http_block_decision_audit_event( &app_state, BlockDecisionAuditEventArgs { @@ -315,16 +327,17 @@ async fn http_connect_accept( .await; let client = client.as_deref().unwrap_or_default(); warn!( - "CONNECT blocked; MITM required to enforce HTTPS policy (client={client}, host={host}, mode={mode:?}, hooked_host={host_has_mitm_hooks})" + "CONNECT blocked; MITM required to enforce HTTPS policy (client={client}, host={host}, mode={mode:?}, host_mitm_requirement={host_mitm_requirement:?})" ); return Err(blocked_text_with_details(REASON_MITM_REQUIRED, &details)); } req.extensions_mut().insert(ProxyTarget(authority)); - req.extensions_mut() - .insert(ConnectMitmEnabled(connect_needs_mitm)); + req.extensions_mut().insert(connect_mitm_mode); req.extensions_mut().insert(mode); - if connect_needs_mitm && let Some(mitm_state) = mitm_state { + if connect_mitm_mode != ConnectMitmMode::Disabled + && let Some(mitm_state) = mitm_state + { req.extensions_mut().insert(mitm_state); } @@ -338,51 +351,68 @@ async fn http_connect_accept( } async fn http_connect_proxy(upgraded: Upgraded) -> Result<(), Infallible> { - let mode = upgraded + let connect_mitm_mode = upgraded + .extensions() + .get::() + .copied() + .unwrap_or(ConnectMitmMode::Disabled); + let result: Result<(), OpaqueError> = match connect_mitm_mode { + ConnectMitmMode::Disabled => forward_connect_tunnel(upgraded).await, + ConnectMitmMode::Enabled => mitm_connect_tunnel(upgraded).await, + ConnectMitmMode::DetectTls => match mitm::peek_tls_prefix(upgraded).await { + Ok((true, stream)) => mitm_connect_tunnel(stream).await, + Ok((false, stream)) => forward_connect_tunnel(stream).await, + Err(err) => Err(OpaqueError::from_display(format!("detect TLS: {err:#}"))), + }, + }; + if let Err(err) = result { + warn!("CONNECT tunnel error: {err}"); + } + Ok(()) +} + +async fn mitm_connect_tunnel(stream: S) -> Result<(), OpaqueError> +where + S: Stream + Unpin + ExtensionsMut, +{ + let target = stream + .extensions() + .get::() + .map(|target| target.0.clone()) + .ok_or_else(|| OpaqueError::from_display("missing MITM authority"))?; + let host = normalize_host(&target.host.to_string()); + let port = target.port; + let mode = stream .extensions() .get::() .copied() .unwrap_or(NetworkMode::Full); - - let Some(target) = upgraded - .extensions() - .get::() - .map(|t| t.0.clone()) - else { - warn!("CONNECT missing proxy target"); - return Ok(()); - }; - - if upgraded - .extensions() - .get::() - .is_some_and(|enabled| enabled.0) - && upgraded - .extensions() - .get::>() - .is_some() - { - let host = normalize_host(&target.host.to_string()); - let port = target.port; - info!("CONNECT MITM enabled (host={host}, port={port}, mode={mode:?})"); - if let Err(err) = mitm::mitm_tunnel(upgraded).await { - warn!("MITM tunnel error: {err}"); - } - return Ok(()); + if stream.extensions().get::>().is_none() { + return Err(OpaqueError::from_display(format!( + "cannot enable MITM without state (host={host}, port={port})" + ))); } - let app_state = match upgraded + info!("CONNECT MITM enabled (host={host}, port={port}, mode={mode:?})"); + mitm::mitm_stream(stream) + .await + .map_err(|err| OpaqueError::from_display(format!("MITM tunnel error: {err}"))) +} + +async fn forward_connect_tunnel(upgraded: S) -> Result<(), OpaqueError> +where + S: Stream + Unpin + ExtensionsMut, +{ + let authority = upgraded + .extensions() + .get::() + .map(|target| target.0.clone()) + .ok_or_else(|| OpaqueError::from_display("missing forward authority"))?; + let app_state = upgraded .extensions() .get::>() .cloned() - { - Some(state) => state, - None => { - error!("missing app state"); - return Ok(()); - } - }; - + .ok_or_else(|| OpaqueError::from_display("missing app state"))?; let allow_upstream_proxy = match app_state.allow_upstream_proxy().await { Ok(allowed) => allowed, Err(err) => { @@ -390,7 +420,6 @@ async fn http_connect_proxy(upgraded: Upgraded) -> Result<(), Infallible> { false } }; - let proxy = if allow_upstream_proxy { proxy_for_connect() } else { @@ -399,31 +428,14 @@ async fn http_connect_proxy(upgraded: Upgraded) -> Result<(), Infallible> { match proxy.as_ref() { Some(proxy) => info!( "CONNECT route selected (host={}, port={}, route=upstream_proxy, proxy={})", - target.host, target.port, proxy.address + authority.host, authority.port, proxy.address ), None => info!( "CONNECT route selected (host={}, port={}, route=direct)", - target.host, target.port + authority.host, authority.port ), } - if let Err(err) = forward_connect_tunnel(upgraded, proxy, app_state).await { - warn!("tunnel error: {err}"); - } - Ok(()) -} - -async fn forward_connect_tunnel( - upgraded: Upgraded, - proxy: Option, - app_state: Arc, -) -> Result<(), BoxError> { - let authority = upgraded - .extensions() - .get::() - .map(|target| target.0.clone()) - .ok_or_else(|| OpaqueError::from_display("missing forward authority").into_boxed())?; - let mut extensions = upgraded.extensions().clone(); if let Some(proxy) = proxy { extensions.insert(proxy); @@ -454,8 +466,7 @@ async fn forward_connect_tunnel( connect_started_at.elapsed().as_millis() ); return Err(OpaqueError::from_boxed(err) - .with_context(|| format!("establish CONNECT tunnel to {authority}")) - .into_boxed()); + .with_context(|| format!("establish CONNECT tunnel to {authority}"))); } }; @@ -481,7 +492,6 @@ async fn forward_connect_tunnel( ); OpaqueError::from_boxed(err.into()) .with_context(|| format!("forward CONNECT tunnel to {authority}")) - .into_boxed() }) } @@ -784,6 +794,15 @@ async fn http_plain_proxy( )); } + if let Err(err) = + inject_plaintext_credentials_if_enabled(app_state.as_ref(), &host, req.headers_mut()).await + { + return Ok(internal_error( + "failed to read plaintext credential injection config", + err, + )); + } + let client = client.as_deref().unwrap_or_default(); let method = req.method(); info!("request allowed (client={client}, host={host}, method={method})"); @@ -813,6 +832,17 @@ async fn http_plain_proxy( } } +async fn inject_plaintext_credentials_if_enabled( + app_state: &NetworkProxyState, + host: &str, + headers: &mut HeaderMap, +) -> Result<()> { + if app_state.plaintext_credential_injection_enabled().await? { + app_state.inject_request_credentials(host, headers); + } + Ok(()) +} + async fn proxy_via_unix_socket(req: Request, socket_path: &str) -> Result { #[cfg(target_os = "macos")] { @@ -1058,6 +1088,7 @@ mod tests { use pretty_assertions::assert_eq; use rama_http::Method; use rama_http::Request; + use std::collections::HashMap; use std::net::Ipv4Addr; use std::net::TcpListener as StdTcpListener; use std::sync::Arc; @@ -1165,6 +1196,99 @@ mod tests { ); } + #[tokio::test] + async fn http_connect_accept_defers_brokered_host_mitm_until_protocol_detection() { + let mut policy = NetworkProxySettings { + credential_broker: true, + mitm: true, + ..NetworkProxySettings::default() + }; + policy.set_allowed_domains(vec!["github.com".to_string()]); + let state = Arc::new(network_proxy_state_for_policy(policy)); + let mut env = HashMap::from([("GH_TOKEN".to_string(), "ghp-real".to_string())]); + state.virtualize_child_credentials(&mut env); + + let mut req = Request::builder() + .method(Method::CONNECT) + .uri("https://github.com:22") + .header("host", "github.com:22") + .body(Body::empty()) + .unwrap(); + req.extensions_mut().insert(state); + + let (response, request) = http_connect_accept( + /*policy_decider*/ None, /*environment_id*/ None, req, + ) + .await + .expect("brokered credentials should defer MITM until protocol detection"); + assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + request.extensions().get::().copied(), + Some(ConnectMitmMode::DetectTls) + ); + } + + #[tokio::test] + async fn plaintext_credential_injection_requires_explicit_opt_in() { + let real_token = "ghp-real"; + let mut disabled_network = NetworkProxySettings { + credential_broker: true, + mitm: true, + ..NetworkProxySettings::default() + }; + disabled_network.set_allowed_domains(vec!["api.github.com".to_string()]); + let disabled_state = Arc::new(network_proxy_state_for_policy(disabled_network)); + let mut disabled_env = HashMap::from([("GH_TOKEN".to_string(), real_token.to_string())]); + disabled_state.virtualize_child_credentials(&mut disabled_env); + let dummy_token = disabled_env.get("GH_TOKEN").expect("dummy GitHub token"); + let mut disabled_headers = HeaderMap::from_iter([( + header::AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy_token}")) + .expect("valid authorization header"), + )]); + + inject_plaintext_credentials_if_enabled( + disabled_state.as_ref(), + "api.github.com", + &mut disabled_headers, + ) + .await + .expect("disabled plaintext injection check should succeed"); + assert_eq!( + disabled_headers.get(header::AUTHORIZATION), + Some(&HeaderValue::from_str(&format!("Bearer {dummy_token}")).unwrap()) + ); + + let mut enabled_network = NetworkProxySettings { + credential_broker: true, + dangerously_allow_plaintext_credential_injection: true, + mitm: true, + ..NetworkProxySettings::default() + }; + enabled_network.set_allowed_domains(vec!["api.github.com".to_string()]); + let enabled_state = Arc::new(network_proxy_state_for_policy(enabled_network)); + let mut enabled_env = HashMap::from([("GH_TOKEN".to_string(), real_token.to_string())]); + enabled_state.virtualize_child_credentials(&mut enabled_env); + let enabled_dummy = enabled_env.get("GH_TOKEN").expect("dummy GitHub token"); + let mut enabled_headers = HeaderMap::from_iter([( + header::AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {enabled_dummy}")) + .expect("valid authorization header"), + )]); + + inject_plaintext_credentials_if_enabled( + enabled_state.as_ref(), + "api.github.com", + &mut enabled_headers, + ) + .await + .expect("enabled plaintext injection check should succeed"); + assert_eq!( + enabled_headers.get(header::AUTHORIZATION), + Some(&HeaderValue::from_str(&format!("Bearer {real_token}")).unwrap()) + ); + } + #[tokio::test] async fn http_connect_accept_blocks_hooked_host_in_full_mode_without_mitm_state() { let mut policy = NetworkProxySettings { @@ -1185,8 +1309,8 @@ mod tests { let mut req = Request::builder() .method(Method::CONNECT) - .uri("https://api.github.com:443") - .header("host", "api.github.com:443") + .uri("https://api.github.com:8443") + .header("host", "api.github.com:8443") .body(Body::empty()) .unwrap(); req.extensions_mut().insert(state); @@ -1204,7 +1328,8 @@ mod tests { } #[tokio::test] - async fn http_proxy_listener_accepts_plain_http1_connect_requests() { + async fn brokered_connect_forwards_server_first_opaque_protocol_without_mitm() { + let server_banner = b"SSH-2.0-server\r\n"; let target_listener = TokioTcpListener::bind((Ipv4Addr::LOCALHOST, 0)) .await .expect("target listener should bind"); @@ -1216,16 +1341,30 @@ mod tests { .accept() .await .expect("target listener should accept"); - let mut buf = [0_u8; 1]; - let _ = timeout(Duration::from_secs(1), stream.read(&mut buf)).await; + stream + .write_all(server_banner) + .await + .expect("target should write opaque server bytes"); }); let state = Arc::new(network_proxy_state_for_policy({ - let mut network = NetworkProxySettings::default(); + let mut network = NetworkProxySettings { + credential_broker: true, + mitm: true, + ..NetworkProxySettings::default() + }; network.set_allowed_domains(vec!["127.0.0.1".to_string()]); network.allow_local_binding = true; network })); + let mut env = HashMap::from([ + ("GH_HOST".to_string(), "127.0.0.1".to_string()), + ( + "GH_ENTERPRISE_TOKEN".to_string(), + "ghp-enterprise-real".to_string(), + ), + ]); + state.virtualize_child_credentials(&mut env); let listener = StdTcpListener::bind((Ipv4Addr::LOCALHOST, 0)).expect("proxy listener should bind"); let proxy_addr = listener @@ -1258,11 +1397,17 @@ mod tests { "unexpected proxy response: {response:?}" ); + let mut buf = vec![0_u8; server_banner.len()]; + timeout(Duration::from_secs(2), stream.read_exact(&mut buf)) + .await + .expect("opaque server bytes should arrive before timeout") + .expect("client should read opaque server bytes"); + assert_eq!(buf, server_banner); + drop(stream); proxy_task.abort(); let _ = proxy_task.await; - target_task.abort(); - let _ = target_task.await; + target_task.await.expect("target task should finish"); } #[tokio::test(flavor = "current_thread")] diff --git a/codex-rs/network-proxy/src/lib.rs b/codex-rs/network-proxy/src/lib.rs index 5bce99030..82ab48ba8 100644 --- a/codex-rs/network-proxy/src/lib.rs +++ b/codex-rs/network-proxy/src/lib.rs @@ -3,6 +3,7 @@ mod certs; mod config; mod connect_policy; +mod credential_broker; mod http_proxy; mod mitm; mod mitm_hook; @@ -27,6 +28,9 @@ pub use config::NetworkProxyConfig; pub use config::NetworkUnixSocketPermission; pub use config::NetworkUnixSocketPermissions; pub use config::host_and_port_from_network_addr; +pub use credential_broker::CREDENTIAL_BROKER_ACTIVE_ENV_KEY; +pub use credential_broker::brokered_credential_dummy_env_keys; +pub use credential_broker::brokered_credential_env_keys; pub use mitm_hook::InjectedHeaderConfig; pub use mitm_hook::MitmHookActionsConfig; pub use mitm_hook::MitmHookBodyConfig; diff --git a/codex-rs/network-proxy/src/mitm.rs b/codex-rs/network-proxy/src/mitm.rs index 6607d2ac0..4e7243487 100644 --- a/codex-rs/network-proxy/src/mitm.rs +++ b/codex-rs/network-proxy/src/mitm.rs @@ -26,6 +26,8 @@ use rama_core::extensions::ExtensionsRef; use rama_core::futures::stream::Stream as FuturesStream; use rama_core::rt::Executor; use rama_core::service::service_fn; +use rama_core::stream::PeekStream; +use rama_core::stream::StackReader; use rama_core::stream::Stream; use rama_http::Body; use rama_http::BodyDataStream; @@ -39,15 +41,18 @@ use rama_http::header::HOST; use rama_http::layer::remove_header::RemoveRequestHeaderLayer; use rama_http::layer::remove_header::RemoveResponseHeaderLayer; use rama_http_backend::server::HttpServer; -use rama_http_backend::server::layer::upgrade::Upgraded; use rama_net::proxy::ProxyTarget; use rama_net::stream::SocketInfo; +use rama_net::tls::server::TlsPeekStream; use rama_tls_rustls::server::TlsAcceptorData; use rama_tls_rustls::server::TlsAcceptorLayer; use std::pin::Pin; use std::sync::Arc; use std::task::Context as TaskContext; use std::task::Poll; +use std::time::Duration; +use tokio::io::AsyncReadExt; +use tokio::time::timeout; use tracing::info; use tracing::warn; @@ -87,6 +92,51 @@ enum MitmPolicyDecision { const MITM_INSPECT_BODIES: bool = false; const MITM_MAX_BODY_BYTES: usize = 4096; +const TLS_PREFIX_LEN: usize = 5; +const TLS_PREFIX_FIRST_BYTE_TIMEOUT: Duration = Duration::from_millis(250); + +/// Peeks enough bytes to distinguish a TLS handshake from an opaque CONNECT stream. +/// +/// The first-byte timeout preserves server-first protocols. Once the client starts a possible TLS +/// prefix, all five record-header bytes are accumulated so fragmented handshakes cannot bypass +/// interception. Every byte read is replayed through `TlsPeekStream`. +pub(crate) async fn peek_tls_prefix(mut stream: S) -> Result<(bool, TlsPeekStream)> +where + S: Stream + Unpin + ExtensionsMut, +{ + let mut peek_buf = [0_u8; TLS_PREFIX_LEN]; + let mut bytes_read = + match timeout(TLS_PREFIX_FIRST_BYTE_TIMEOUT, stream.read(&mut peek_buf)).await { + Ok(result) => result.context("read TLS prefix")?, + Err(_) => 0, + }; + while bytes_read > 0 && bytes_read < TLS_PREFIX_LEN { + let possible_tls_prefix = matches!( + &peek_buf[..bytes_read], + [0x16] | [0x16, 0x03] | [0x16, 0x03, 0x00..=0x04, ..] + ); + if !possible_tls_prefix { + break; + } + let read = stream + .read(&mut peek_buf[bytes_read..]) + .await + .context("read TLS prefix")?; + if read == 0 { + break; + } + bytes_read += read; + } + + let is_tls = bytes_read == TLS_PREFIX_LEN && matches!(peek_buf, [0x16, 0x03, 0x00..=0x04, ..]); + let offset = TLS_PREFIX_LEN - bytes_read; + if offset > 0 { + peek_buf.copy_within(0..bytes_read, offset); + } + let mut peek = StackReader::new(peek_buf); + peek.skip(offset); + Ok((is_tls, PeekStream::new(peek, stream))) +} impl std::fmt::Debug for MitmState { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { @@ -143,11 +193,6 @@ impl MitmState { } } -/// Terminate the upgraded CONNECT stream with a generated leaf cert and proxy inner HTTPS traffic. -pub(crate) async fn mitm_tunnel(upgraded: Upgraded) -> Result<()> { - mitm_stream(upgraded).await -} - /// Terminate a raw client stream with a generated leaf cert and proxy inner HTTPS traffic. pub(crate) async fn mitm_stream(stream: S) -> Result<()> where @@ -247,6 +292,10 @@ async fn forward_request(req: Request, request_ctx: &MitmRequestContext) -> Resu let log_path = path_for_log(req.uri()); let (mut parts, body) = req.into_parts(); + request_ctx + .policy + .app_state + .inject_request_credentials(&target_host, &mut parts.headers); apply_mitm_hook_actions(&mut parts.headers, hook_actions.as_ref()); let authority = authority_header_value(&target_host, target_port); parts.uri = build_https_uri(&authority, &path)?; diff --git a/codex-rs/network-proxy/src/mitm_tests.rs b/codex-rs/network-proxy/src/mitm_tests.rs index 0823be517..18e650baa 100644 --- a/codex-rs/network-proxy/src/mitm_tests.rs +++ b/codex-rs/network-proxy/src/mitm_tests.rs @@ -7,6 +7,9 @@ use crate::reasons::REASON_NOT_ALLOWED_LOCAL; use crate::runtime::network_proxy_state_for_policy; use codex_utils_absolute_path::AbsolutePathBuf; use pretty_assertions::assert_eq; +use rama_core::extensions::Extensions; +use rama_core::extensions::ExtensionsMut; +use rama_core::extensions::ExtensionsRef; use rama_http::Body; use rama_http::HeaderMap; use rama_http::HeaderValue; @@ -14,7 +17,90 @@ use rama_http::Method; use rama_http::Request; use rama_http::StatusCode; use rama_http::header::HeaderName; +use std::pin::Pin; +use std::task::Context; +use std::task::Poll; use tempfile::NamedTempFile; +use tokio::io::AsyncRead; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWrite; +use tokio::io::AsyncWriteExt; +use tokio::io::DuplexStream; +use tokio::io::ReadBuf; +use tokio::time::Duration; + +struct TestStream { + inner: DuplexStream, + extensions: Extensions, +} + +impl TestStream { + fn new(inner: DuplexStream) -> Self { + Self { + inner, + extensions: Extensions::new(), + } + } +} + +impl AsyncRead for TestStream { + fn poll_read( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_read(cx, buf) + } +} + +impl AsyncWrite for TestStream { + fn poll_write( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_write(cx, buf) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_flush(cx) + } + + fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_shutdown(cx) + } +} + +impl ExtensionsRef for TestStream { + fn extensions(&self) -> &Extensions { + &self.extensions + } +} + +impl ExtensionsMut for TestStream { + fn extensions_mut(&mut self) -> &mut Extensions { + &mut self.extensions + } +} + +#[tokio::test] +async fn tls_prefix_detection_accumulates_fragmented_reads() { + let tls_prefix = [0x16, 0x03, 0x03, 0x00, 0x80]; + let (mut writer, reader) = tokio::io::duplex(16); + let writer_task = tokio::spawn(async move { + writer.write_all(&tls_prefix[..1]).await.unwrap(); + tokio::time::sleep(TLS_PREFIX_FIRST_BYTE_TIMEOUT + Duration::from_millis(50)).await; + writer.write_all(&tls_prefix[1..]).await.unwrap(); + }); + + let (is_tls, mut stream) = peek_tls_prefix(TestStream::new(reader)).await.unwrap(); + let mut replayed = [0_u8; 5]; + stream.read_exact(&mut replayed).await.unwrap(); + + assert!(is_tls); + assert_eq!(replayed, tls_prefix); + writer_task.await.unwrap(); +} fn github_write_hook() -> crate::mitm_hook::MitmHookConfig { crate::mitm_hook::MitmHookConfig { diff --git a/codex-rs/network-proxy/src/proxy.rs b/codex-rs/network-proxy/src/proxy.rs index 7d2372f17..a338744b6 100644 --- a/codex-rs/network-proxy/src/proxy.rs +++ b/codex-rs/network-proxy/src/proxy.rs @@ -1,4 +1,6 @@ use crate::config; +use crate::credential_broker::BROKERED_CREDENTIALS_ENV_KEY; +use crate::credential_broker::CREDENTIAL_BROKER_ACTIVE_ENV_KEY; use crate::http_proxy; use crate::network_policy::NetworkPolicyDecider; use crate::runtime::BlockedRequestObserver; @@ -419,6 +421,8 @@ const NODE_USE_ENV_PROXY_ENV_KEY: &str = "NODE_USE_ENV_PROXY"; const GIT_SSH_COMMAND_ENV_KEY: &str = "GIT_SSH_COMMAND"; pub const PROXY_ENV_KEYS: &[&str] = &[ PROXY_ACTIVE_ENV_KEY, + CREDENTIAL_BROKER_ACTIVE_ENV_KEY, + BROKERED_CREDENTIALS_ENV_KEY, ALLOW_LOCAL_BINDING_ENV_KEY, ELECTRON_GET_USE_PROXY_ENV_KEY, NODE_USE_ENV_PROXY_ENV_KEY, @@ -692,6 +696,7 @@ impl NetworkProxy { runtime_settings.allow_local_binding, runtime_settings.mitm_ca_trust_bundle.as_ref(), ); + self.state.virtualize_child_credentials(&mut env); let mut loopback_ports = [ Some(addrs.http_addr), self.socks_enabled.then_some(addrs.socks_addr), diff --git a/codex-rs/network-proxy/src/runtime.rs b/codex-rs/network-proxy/src/runtime.rs index e721d0de7..b486ab057 100644 --- a/codex-rs/network-proxy/src/runtime.rs +++ b/codex-rs/network-proxy/src/runtime.rs @@ -2,6 +2,7 @@ use crate::config::NetworkDomainPermission; use crate::config::NetworkMode; use crate::config::NetworkProxyConfig; use crate::config::ValidatedUnixSocketPath; +use crate::credential_broker::CredentialBroker; use crate::mitm::MitmState; use crate::mitm_hook::HookEvaluation; use crate::mitm_hook::MitmHooksByHost; @@ -23,6 +24,7 @@ use anyhow::Result; use codex_utils_absolute_path::AbsolutePathBuf; use globset::GlobSet; use serde::Serialize; +use std::collections::HashMap; use std::collections::HashSet; use std::collections::VecDeque; use std::future::Future; @@ -207,9 +209,17 @@ pub struct NetworkProxyState { state: Arc>, reloader: Arc, blocked_request_observer: Arc>>>, + credential_broker: CredentialBroker, audit_metadata: NetworkProxyAuditMetadata, } +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum HostMitmRequirement { + None, + Tls, + Always, +} + impl std::fmt::Debug for NetworkProxyState { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { // Avoid logging internal state (config contents, derived globsets, etc.) which can be noisy @@ -224,6 +234,7 @@ impl Clone for NetworkProxyState { state: self.state.clone(), reloader: self.reloader.clone(), blocked_request_observer: self.blocked_request_observer.clone(), + credential_broker: self.credential_broker.clone(), audit_metadata: self.audit_metadata.clone(), } } @@ -271,6 +282,7 @@ impl NetworkProxyState { blocked_request_observer: Option>, ) -> Self { Self { + credential_broker: CredentialBroker::new(state.config.network.credential_broker), state: Arc::new(RwLock::new(state)), reloader, blocked_request_observer: Arc::new(RwLock::new(blocked_request_observer)), @@ -290,6 +302,23 @@ impl NetworkProxyState { &self.audit_metadata } + pub fn virtualize_child_credentials(&self, env: &mut HashMap) { + self.credential_broker.virtualize_child_env(env); + } + + pub fn inject_request_credentials(&self, host: &str, headers: &mut rama_http::HeaderMap) { + self.credential_broker.inject_request_headers(host, headers); + } + + pub async fn plaintext_credential_injection_enabled(&self) -> Result { + self.reload_if_needed().await?; + let guard = self.state.read().await; + Ok(guard + .config + .network + .dangerously_allow_plaintext_credential_injection) + } + pub async fn current_cfg(&self) -> Result { // Callers treat `NetworkProxyState` as a live view of policy. We reload-on-demand so edits to // `config.toml` (including Codex-managed writes) take effect without a restart. @@ -321,6 +350,7 @@ impl NetworkProxyState { match self.reloader.reload_now().await { Ok(mut new_state) => { + self.ensure_credential_broker_enablement_unchanged(&new_state)?; // Policy changes are operationally sensitive; logging diffs makes changes traceable // without needing to dump full config blobs (which can include unrelated settings). log_policy_changes(&previous_cfg, &new_state.config); @@ -343,6 +373,7 @@ impl NetworkProxyState { pub async fn replace_config_state(&self, mut new_state: ConfigState) -> Result<()> { self.reload_if_needed().await?; + self.ensure_credential_broker_enablement_unchanged(&new_state)?; let mut guard = self.state.write().await; log_policy_changes(&guard.config, &new_state.config); new_state.blocked = guard.blocked.clone(); @@ -599,10 +630,20 @@ impl NetworkProxyState { Ok(evaluate_mitm_hooks(&guard.mitm_hooks, host, req)) } - pub async fn host_has_mitm_hooks(&self, host: &str) -> Result { + pub(crate) async fn host_mitm_requirement(&self, host: &str) -> Result { self.reload_if_needed().await?; - let guard = self.state.read().await; - Ok(guard.mitm_hooks.contains_key(&normalize_host(host))) + let normalized_host = normalize_host(host); + let host_has_mitm_hooks = { + let guard = self.state.read().await; + guard.mitm_hooks.contains_key(&normalized_host) + }; + Ok(if host_has_mitm_hooks { + HostMitmRequirement::Always + } else if self.credential_broker.host_requires_mitm(&normalized_host) { + HostMitmRequirement::Tls + } else { + HostMitmRequirement::None + }) } pub async fn add_allowed_domain(&self, host: &str) -> Result<()> { @@ -676,6 +717,7 @@ impl NetworkProxyState { match self.reloader.maybe_reload().await? { None => Ok(()), Some(mut new_state) => { + self.ensure_credential_broker_enablement_unchanged(&new_state)?; let (previous_cfg, blocked, blocked_total) = { let guard = self.state.read().await; ( @@ -697,6 +739,14 @@ impl NetworkProxyState { } } } + + fn ensure_credential_broker_enablement_unchanged(&self, new_state: &ConfigState) -> Result<()> { + anyhow::ensure!( + self.credential_broker.enabled() == new_state.config.network.credential_broker, + "network.credential_broker cannot change while the proxy is running" + ); + Ok(()) + } } #[derive(Clone, Copy)] @@ -918,6 +968,27 @@ mod tests { use crate::state::validate_policy_against_constraints; use pretty_assertions::assert_eq; + #[derive(Clone)] + struct StaticReloader { + state: ConfigState, + } + + impl ConfigReloader for StaticReloader { + fn source_label(&self) -> String { + "static test reloader".to_string() + } + + fn maybe_reload(&self) -> ConfigReloaderFuture<'_, Option> { + let state = self.state.clone(); + Box::pin(async move { Ok(Some(state)) }) + } + + fn reload_now(&self) -> ConfigReloaderFuture<'_, ConfigState> { + let state = self.state.clone(); + Box::pin(async move { Ok(state) }) + } + } + fn strings(entries: &[&str]) -> Vec { entries.iter().map(|entry| (*entry).to_string()).collect() } @@ -945,6 +1016,40 @@ mod tests { network } + #[tokio::test] + async fn reload_rejects_credential_broker_enablement_changes() { + let initial_state = build_config_state( + NetworkProxyConfig::default(), + NetworkProxyConstraints::default(), + ) + .unwrap(); + let mut reloaded_state = initial_state.clone(); + reloaded_state + .config + .set_credential_broker_enabled(/*enabled*/ true); + let state = NetworkProxyState::with_reloader( + initial_state, + Arc::new(StaticReloader { + state: reloaded_state, + }), + ); + + let err = state + .force_reload() + .await + .expect_err("credential broker enablement should require a proxy restart"); + let mut env = HashMap::from([("OPENAI_API_KEY".to_string(), "sk-real".to_string())]); + state.virtualize_child_credentials(&mut env); + + assert!( + format!("{err:#}") + .contains("network.credential_broker cannot change while the proxy is running"), + "unexpected error: {err:#}" + ); + assert_eq!(env["OPENAI_API_KEY"], "sk-real"); + assert!(!state.credential_broker.enabled()); + } + #[tokio::test] async fn host_blocked_denied_wins_over_allowed() { let state = diff --git a/codex-rs/network-proxy/src/socks5.rs b/codex-rs/network-proxy/src/socks5.rs index 5cc7b2eed..956187704 100644 --- a/codex-rs/network-proxy/src/socks5.rs +++ b/codex-rs/network-proxy/src/socks5.rs @@ -17,6 +17,7 @@ use crate::reasons::REASON_MITM_REQUIRED; use crate::reasons::REASON_PROXY_DISABLED; use crate::responses::PolicyDecisionDetails; use crate::responses::blocked_message_with_policy; +use crate::runtime::HostMitmRequirement; use crate::state::BlockedRequest; use crate::state::BlockedRequestArgs; use crate::state::NetworkProxyState; @@ -341,10 +342,10 @@ async fn handle_socks5_tcp( } } - let host_has_mitm_hooks = match app_state.host_has_mitm_hooks(&host).await { - Ok(has_hooks) => has_hooks, + let host_mitm_requirement = match app_state.host_mitm_requirement(&host).await { + Ok(requirement) => requirement, Err(err) => { - error!("failed to inspect MITM hooks for {host}: {err}"); + error!("failed to inspect MITM requirements for {host}: {err}"); return Err(io::Error::other("proxy error").into()); } }; @@ -355,10 +356,19 @@ async fn handle_socks5_tcp( return Err(io::Error::other("proxy error").into()); } }; - let socks_needs_mitm = - socks5_tcp_target_is_https && (mode == NetworkMode::Limited || host_has_mitm_hooks); - if (host_has_mitm_hooks && !socks5_tcp_target_is_https) - || (socks_needs_mitm && mitm_state.is_none()) + let socks_mitm_mode = if mode == NetworkMode::Limited { + SocksMitmMode::Enabled + } else { + match host_mitm_requirement { + HostMitmRequirement::None => SocksMitmMode::Disabled, + HostMitmRequirement::Tls => SocksMitmMode::DetectTls, + HostMitmRequirement::Always => SocksMitmMode::Enabled, + } + }; + let unsupported_hook_protocol = + host_mitm_requirement == HostMitmRequirement::Always && !socks5_tcp_target_is_https; + if unsupported_hook_protocol + || (socks_mitm_mode != SocksMitmMode::Disabled && mitm_state.is_none()) { emit_socks_block_decision_audit_event( &app_state, @@ -392,23 +402,35 @@ async fn handle_socks5_tcp( .await; let client = client.as_deref().unwrap_or_default(); warn!( - "SOCKS blocked; MITM required to enforce HTTPS policy (client={client}, host={host}, mode={mode:?}, hooked_host={host_has_mitm_hooks}, https_target={socks5_tcp_target_is_https})" + "SOCKS blocked; MITM required to enforce HTTPS policy (client={client}, host={host}, mode={mode:?}, host_mitm_requirement={host_mitm_requirement:?}, https_target={socks5_tcp_target_is_https})" ); return Err(policy_denied_error(REASON_MITM_REQUIRED, &details).into()); } - if socks_needs_mitm && let Some(mitm_state) = mitm_state { + if let Some(mitm_state) = mitm_state { let client = client.as_deref().unwrap_or_default(); - info!("SOCKS MITM enabled (client={client}, host={host}, port={port}, mode={mode:?})"); - return Ok(EstablishedClientConnection { - input: req, - conn: Socks5TcpConnection::Mitm { + let conn = match socks_mitm_mode { + SocksMitmMode::Disabled => None, + SocksMitmMode::Enabled => Some(Socks5TcpConnection::Mitm { target, mode, mitm: mitm_state, extensions: Extensions::new(), - }, - }); + }), + SocksMitmMode::DetectTls => Some(Socks5TcpConnection::DetectTls { + target, + mode, + mitm: mitm_state, + state: app_state, + extensions: Extensions::new(), + }), + }; + if let Some(conn) = conn { + info!( + "SOCKS MITM selected (client={client}, host={host}, port={port}, mode={mode:?}, mitm_mode={socks_mitm_mode:?})" + ); + return Ok(EstablishedClientConnection { input: req, conn }); + } } info!("SOCKS upstream dial started (host={host}, port={port})"); @@ -435,6 +457,13 @@ async fn handle_socks5_tcp( /// Internal connector output for SOCKS5 TCP. MITM requests do not dial upstream before the /// inner HTTPS request is inspected, so they carry the target metadata instead of a socket. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum SocksMitmMode { + Disabled, + Enabled, + DetectTls, +} + #[derive(Debug)] enum Socks5TcpConnection { Direct(TcpStream), @@ -444,6 +473,13 @@ enum Socks5TcpConnection { mitm: Arc, extensions: Extensions, }, + DetectTls { + target: HostWithPort, + mode: NetworkMode, + mitm: Arc, + state: Arc, + extensions: Extensions, + }, } impl AsyncRead for Socks5TcpConnection { @@ -454,7 +490,7 @@ impl AsyncRead for Socks5TcpConnection { ) -> Poll> { match self.get_mut() { Self::Direct(stream) => Pin::new(stream).poll_read(cx, buf), - Self::Mitm { .. } => Poll::Ready(Ok(())), + Self::Mitm { .. } | Self::DetectTls { .. } => Poll::Ready(Ok(())), } } } @@ -467,21 +503,21 @@ impl AsyncWrite for Socks5TcpConnection { ) -> Poll> { match self.get_mut() { Self::Direct(stream) => Pin::new(stream).poll_write(cx, buf), - Self::Mitm { .. } => Poll::Ready(Ok(buf.len())), + Self::Mitm { .. } | Self::DetectTls { .. } => Poll::Ready(Ok(buf.len())), } } fn poll_flush(self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll> { match self.get_mut() { Self::Direct(stream) => Pin::new(stream).poll_flush(cx), - Self::Mitm { .. } => Poll::Ready(Ok(())), + Self::Mitm { .. } | Self::DetectTls { .. } => Poll::Ready(Ok(())), } } fn poll_shutdown(self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll> { match self.get_mut() { Self::Direct(stream) => Pin::new(stream).poll_shutdown(cx), - Self::Mitm { .. } => Poll::Ready(Ok(())), + Self::Mitm { .. } | Self::DetectTls { .. } => Poll::Ready(Ok(())), } } } @@ -490,14 +526,14 @@ impl Socket for Socks5TcpConnection { fn local_addr(&self) -> io::Result { match self { Self::Direct(stream) => stream.local_addr(), - Self::Mitm { .. } => Ok(SocketAddr::from(([0, 0, 0, 0], 0))), + Self::Mitm { .. } | Self::DetectTls { .. } => Ok(SocketAddr::from(([0, 0, 0, 0], 0))), } } fn peer_addr(&self) -> io::Result { match self { Self::Direct(stream) => stream.peer_addr(), - Self::Mitm { .. } => Ok(SocketAddr::from(([0, 0, 0, 0], 0))), + Self::Mitm { .. } | Self::DetectTls { .. } => Ok(SocketAddr::from(([0, 0, 0, 0], 0))), } } } @@ -506,7 +542,7 @@ impl ExtensionsRef for Socks5TcpConnection { fn extensions(&self) -> &Extensions { match self { Self::Direct(stream) => stream.extensions(), - Self::Mitm { extensions, .. } => extensions, + Self::Mitm { extensions, .. } | Self::DetectTls { extensions, .. } => extensions, } } } @@ -515,7 +551,7 @@ impl ExtensionsMut for Socks5TcpConnection { fn extensions_mut(&mut self) -> &mut Extensions { match self { Self::Direct(stream) => stream.extensions_mut(), - Self::Mitm { extensions, .. } => extensions, + Self::Mitm { extensions, .. } | Self::DetectTls { extensions, .. } => extensions, } } } @@ -537,6 +573,41 @@ async fn proxy_socks5_tcp( source.extensions_mut().insert(mitm); mitm::mitm_stream(source).await.map_err(Into::into) } + Socks5TcpConnection::DetectTls { + target, + mode, + mitm, + state, + .. + } => { + source.extensions_mut().insert(ProxyTarget(target.clone())); + source.extensions_mut().insert(mode); + source.extensions_mut().insert(mitm); + let (is_tls, source) = mitm::peek_tls_prefix(source) + .await + .map_err(|err| -> BoxError { err.into() })?; + if is_tls { + mitm::mitm_stream(source).await.map_err(Into::into) + } else { + info!("SOCKS opaque upstream dial started (target={target})"); + let connect_started_at = Instant::now(); + let EstablishedClientConnection { conn: upstream, .. } = + TargetCheckedTcpConnector::new(state) + .serve(TcpRequest::new(target.clone())) + .await?; + info!( + "SOCKS opaque upstream dial established (target={target}, elapsed_ms={})", + connect_started_at.elapsed().as_millis() + ); + StreamForwardService::default() + .serve(ProxyRequest { + source, + target: upstream, + }) + .await + .map_err(Into::into) + } + } } } @@ -752,6 +823,7 @@ mod tests { use rama_net::address::HostWithPort; use rama_net::address::SocketAddress; use rama_socks5::server::udp::RelayDirection; + use std::collections::HashMap; use std::net::IpAddr; use std::net::Ipv4Addr; use std::sync::Arc; @@ -908,6 +980,36 @@ mod tests { assert_eq!(event.field("client.address"), Some("unknown")); } + #[tokio::test(flavor = "current_thread")] + async fn handle_socks5_tcp_detects_tls_for_brokered_nonstandard_port_in_full_mode() { + let mut settings = NetworkProxySettings { + enabled: true, + mode: NetworkMode::Full, + mitm: true, + credential_broker: true, + ..NetworkProxySettings::default() + }; + settings.set_allowed_domains(vec!["api.openai.com".to_string()]); + let state = state_for_settings(settings); + let mut env = HashMap::from([("OPENAI_API_KEY".to_string(), "sk-real".to_string())]); + state.virtualize_child_credentials(&mut env); + let mut request = TcpRequest::new( + HostWithPort::try_from("api.openai.com:8443").expect("valid authority"), + ); + request.extensions_mut().insert(state.clone()); + + let result = handle_socks5_tcp( + request, + TargetCheckedTcpConnector::new(state), + /*policy_decider*/ None, + /*environment_id*/ None, + ) + .await + .expect("brokered TLS should defer MITM until protocol detection"); + + assert!(matches!(result.conn, Socks5TcpConnection::DetectTls { .. })); + } + #[tokio::test(flavor = "current_thread")] async fn handle_socks5_tcp_blocks_limited_mode_without_mitm_state() { let mut settings = NetworkProxySettings { diff --git a/codex-rs/network-proxy/src/state.rs b/codex-rs/network-proxy/src/state.rs index 32cdfab14..9a5287dac 100644 --- a/codex-rs/network-proxy/src/state.rs +++ b/codex-rs/network-proxy/src/state.rs @@ -57,6 +57,8 @@ pub struct PartialNetworkConfig { pub unix_sockets: Option, pub allow_local_binding: Option, pub mitm: Option, + pub credential_broker: Option, + pub dangerously_allow_plaintext_credential_injection: Option, #[serde(default)] pub mitm_hooks: Option>, } @@ -66,6 +68,10 @@ pub fn build_config_state( constraints: NetworkProxyConstraints, ) -> anyhow::Result { crate::config::validate_unix_socket_allowlist_paths(&config)?; + anyhow::ensure!( + !config.network.credential_broker || config.network.mitm, + "network.credential_broker requires network.mitm = true" + ); let allowed_domains = config.network.allowed_domains().unwrap_or_default(); let denied_domains = config.network.denied_domains().unwrap_or_default(); validate_non_global_wildcard_domain_patterns("network.denied_domains", &denied_domains)