Skip to main content

codex_network_proxy/
state.rs

1use crate::config::NetworkDomainPermissions;
2use crate::config::NetworkMode;
3use crate::config::NetworkProxyConfig;
4use crate::config::NetworkUnixSocketPermissions;
5use crate::mitm::MitmState;
6use crate::mitm::MitmUpstreamConfig;
7use crate::mitm_hook::MitmHookConfig;
8use crate::mitm_hook::compile_mitm_hooks;
9use crate::mitm_hook::validate_mitm_hook_config;
10use crate::policy::DomainPattern;
11use crate::policy::compile_allowlist_globset;
12use crate::policy::compile_denylist_globset;
13use crate::policy::is_global_wildcard_domain_pattern;
14use crate::runtime::ConfigState;
15use serde::Deserialize;
16use std::collections::HashSet;
17use std::sync::Arc;
18
19pub use crate::runtime::BlockedRequest;
20pub use crate::runtime::BlockedRequestArgs;
21pub use crate::runtime::NetworkProxyAuditMetadata;
22pub use crate::runtime::NetworkProxyState;
23#[cfg(test)]
24pub(crate) use crate::runtime::network_proxy_state_for_policy;
25
26#[derive(Debug, Default, Clone, PartialEq, Eq)]
27pub struct NetworkProxyConstraints {
28    pub enabled: Option<bool>,
29    pub mode: Option<NetworkMode>,
30    pub allow_upstream_proxy: Option<bool>,
31    pub dangerously_allow_non_loopback_proxy: Option<bool>,
32    pub dangerously_allow_all_unix_sockets: Option<bool>,
33    pub allowed_domains: Option<Vec<String>>,
34    pub allowlist_expansion_enabled: Option<bool>,
35    pub denied_domains: Option<Vec<String>>,
36    pub denylist_expansion_enabled: Option<bool>,
37    pub allow_unix_sockets: Option<Vec<String>>,
38    pub allow_local_binding: Option<bool>,
39}
40
41#[derive(Debug, Clone, Deserialize)]
42pub struct PartialNetworkProxyConfig {
43    pub enabled: Option<bool>,
44    pub mode: Option<NetworkMode>,
45    pub allow_upstream_proxy: Option<bool>,
46    pub dangerously_allow_non_loopback_proxy: Option<bool>,
47    pub dangerously_allow_all_unix_sockets: Option<bool>,
48    #[serde(default)]
49    pub domains: Option<NetworkDomainPermissions>,
50    #[serde(default)]
51    pub unix_sockets: Option<NetworkUnixSocketPermissions>,
52    pub allow_local_binding: Option<bool>,
53    pub mitm: Option<bool>,
54    pub credential_broker: Option<bool>,
55    pub dangerously_allow_plaintext_credential_injection: Option<bool>,
56    #[serde(default)]
57    pub mitm_hooks: Option<Vec<MitmHookConfig>>,
58}
59
60pub fn build_config_state(
61    config: NetworkProxyConfig,
62    constraints: NetworkProxyConstraints,
63) -> anyhow::Result<ConfigState> {
64    crate::config::validate_unix_socket_allowlist_paths(&config)?;
65    anyhow::ensure!(
66        !config.credential_broker || config.mitm,
67        "network.credential_broker requires network.mitm = true"
68    );
69    let allowed_domains = config.allowed_domains().unwrap_or_default();
70    let denied_domains = config.denied_domains().unwrap_or_default();
71    validate_non_global_wildcard_domain_patterns("network.denied_domains", &denied_domains)
72        .map_err(NetworkProxyConstraintError::into_anyhow)?;
73    let deny_set = compile_denylist_globset(&denied_domains)?;
74    let allow_set = compile_allowlist_globset(&allowed_domains)?;
75    let mitm_hooks = compile_mitm_hooks(&config)?;
76    let mitm = if config.mitm {
77        Some(Arc::new(MitmState::new(MitmUpstreamConfig {
78            allow_upstream_proxy: config.allow_upstream_proxy,
79        })?))
80    } else {
81        None
82    };
83    Ok(ConfigState {
84        config,
85        allow_set,
86        deny_set,
87        mitm,
88        mitm_hooks,
89        constraints,
90        blocked: std::collections::VecDeque::new(),
91        blocked_total: 0,
92    })
93}
94
95pub fn validate_policy_against_constraints(
96    config: &NetworkProxyConfig,
97    constraints: &NetworkProxyConstraints,
98) -> Result<(), NetworkProxyConstraintError> {
99    fn invalid_value(
100        field_name: &'static str,
101        candidate: impl Into<String>,
102        allowed: impl Into<String>,
103    ) -> NetworkProxyConstraintError {
104        NetworkProxyConstraintError::InvalidValue {
105            field_name,
106            candidate: candidate.into(),
107            allowed: allowed.into(),
108        }
109    }
110
111    fn validate<T>(
112        candidate: T,
113        validator: impl FnOnce(&T) -> Result<(), NetworkProxyConstraintError>,
114    ) -> Result<(), NetworkProxyConstraintError> {
115        validator(&candidate)
116    }
117
118    let enabled = config.enabled;
119    let config_allowed_domains = config.allowed_domains().unwrap_or_default();
120    let config_denied_domains = config.denied_domains().unwrap_or_default();
121    let denied_domain_overrides: HashSet<String> = config_denied_domains
122        .iter()
123        .map(|entry| entry.to_ascii_lowercase())
124        .collect();
125    let config_allow_unix_sockets = config.allow_unix_sockets();
126    validate_mitm_hook_config(config).map_err(invalid_mitm_hook_configuration)?;
127    validate_non_global_wildcard_domain_patterns("network.denied_domains", &config_denied_domains)?;
128    if let Some(max_enabled) = constraints.enabled {
129        validate(enabled, move |candidate| {
130            if *candidate && !max_enabled {
131                Err(invalid_value(
132                    "network.enabled",
133                    "true",
134                    "false (disabled by managed config)",
135                ))
136            } else {
137                Ok(())
138            }
139        })?;
140    }
141
142    if let Some(max_mode) = constraints.mode {
143        validate(config.mode, move |candidate| {
144            if network_mode_rank(*candidate) > network_mode_rank(max_mode) {
145                Err(invalid_value(
146                    "network.mode",
147                    format!("{candidate:?}"),
148                    format!("{max_mode:?} or more restrictive"),
149                ))
150            } else {
151                Ok(())
152            }
153        })?;
154    }
155
156    let allow_upstream_proxy = constraints.allow_upstream_proxy;
157    validate(
158        config.allow_upstream_proxy,
159        move |candidate| match allow_upstream_proxy {
160            Some(true) | None => Ok(()),
161            Some(false) => {
162                if *candidate {
163                    Err(invalid_value(
164                        "network.allow_upstream_proxy",
165                        "true",
166                        "false (disabled by managed config)",
167                    ))
168                } else {
169                    Ok(())
170                }
171            }
172        },
173    )?;
174
175    let allow_non_loopback_proxy = constraints.dangerously_allow_non_loopback_proxy;
176    validate(
177        config.dangerously_allow_non_loopback_proxy,
178        move |candidate| match allow_non_loopback_proxy {
179            Some(true) | None => Ok(()),
180            Some(false) => {
181                if *candidate {
182                    Err(invalid_value(
183                        "network.dangerously_allow_non_loopback_proxy",
184                        "true",
185                        "false (disabled by managed config)",
186                    ))
187                } else {
188                    Ok(())
189                }
190            }
191        },
192    )?;
193
194    let allow_all_unix_sockets = constraints
195        .dangerously_allow_all_unix_sockets
196        .unwrap_or(constraints.allow_unix_sockets.is_none());
197    validate(
198        config.dangerously_allow_all_unix_sockets,
199        move |candidate| {
200            if *candidate && !allow_all_unix_sockets {
201                Err(invalid_value(
202                    "network.dangerously_allow_all_unix_sockets",
203                    "true",
204                    "false (disabled by managed config)",
205                ))
206            } else {
207                Ok(())
208            }
209        },
210    )?;
211
212    if let Some(allow_local_binding) = constraints.allow_local_binding {
213        validate(config.allow_local_binding, move |candidate| {
214            if *candidate && !allow_local_binding {
215                Err(invalid_value(
216                    "network.allow_local_binding",
217                    "true",
218                    "false (disabled by managed config)",
219                ))
220            } else {
221                Ok(())
222            }
223        })?;
224    }
225
226    if let Some(allowed_domains) = &constraints.allowed_domains {
227        validate_non_global_wildcard_domain_patterns("network.allowed_domains", allowed_domains)?;
228        match constraints.allowlist_expansion_enabled {
229            Some(true) => {
230                let required_set: HashSet<String> = allowed_domains
231                    .iter()
232                    .map(|entry| entry.to_ascii_lowercase())
233                    .collect();
234                validate(config_allowed_domains, |candidate| {
235                    let candidate_set: HashSet<String> = candidate
236                        .iter()
237                        .map(|entry| entry.to_ascii_lowercase())
238                        .collect();
239                    let missing: Vec<String> = required_set
240                        .iter()
241                        .filter(|entry| {
242                            !candidate_set.contains(*entry)
243                                && !denied_domain_overrides.contains(*entry)
244                        })
245                        .cloned()
246                        .collect();
247                    if missing.is_empty() {
248                        Ok(())
249                    } else {
250                        Err(invalid_value(
251                            "network.allowed_domains",
252                            "missing managed allowed_domains entries",
253                            format!("{missing:?}"),
254                        ))
255                    }
256                })?;
257            }
258            Some(false) => {
259                let required_set: HashSet<String> = allowed_domains
260                    .iter()
261                    .map(|entry| entry.to_ascii_lowercase())
262                    .collect();
263                validate(config_allowed_domains, |candidate| {
264                    let candidate_set: HashSet<String> = candidate
265                        .iter()
266                        .map(|entry| entry.to_ascii_lowercase())
267                        .collect();
268                    let expected_set: HashSet<String> = required_set
269                        .difference(&denied_domain_overrides)
270                        .cloned()
271                        .collect();
272                    if candidate_set == expected_set {
273                        Ok(())
274                    } else {
275                        Err(invalid_value(
276                            "network.allowed_domains",
277                            format!("{candidate:?}"),
278                            "must match managed allowed_domains",
279                        ))
280                    }
281                })?;
282            }
283            None => {
284                let managed_patterns: Vec<DomainPattern> = allowed_domains
285                    .iter()
286                    .map(|entry| DomainPattern::parse_for_constraints(entry))
287                    .collect();
288                validate(config_allowed_domains, move |candidate| {
289                    let mut invalid = Vec::new();
290                    for entry in candidate {
291                        let candidate_pattern = DomainPattern::parse_for_constraints(entry);
292                        if !managed_patterns
293                            .iter()
294                            .any(|managed| managed.allows(&candidate_pattern))
295                        {
296                            invalid.push(entry.clone());
297                        }
298                    }
299                    if invalid.is_empty() {
300                        Ok(())
301                    } else {
302                        Err(invalid_value(
303                            "network.allowed_domains",
304                            format!("{invalid:?}"),
305                            "subset of managed allowed_domains",
306                        ))
307                    }
308                })?;
309            }
310        }
311    }
312
313    if let Some(denied_domains) = &constraints.denied_domains {
314        validate_non_global_wildcard_domain_patterns("network.denied_domains", denied_domains)?;
315        let required_set: HashSet<String> = denied_domains
316            .iter()
317            .map(|s| s.to_ascii_lowercase())
318            .collect();
319        match constraints.denylist_expansion_enabled {
320            Some(false) => {
321                validate(config_denied_domains, move |candidate| {
322                    let candidate_set: HashSet<String> = candidate
323                        .iter()
324                        .map(|entry| entry.to_ascii_lowercase())
325                        .collect();
326                    if candidate_set == required_set {
327                        Ok(())
328                    } else {
329                        Err(invalid_value(
330                            "network.denied_domains",
331                            format!("{candidate:?}"),
332                            "must match managed denied_domains",
333                        ))
334                    }
335                })?;
336            }
337            Some(true) | None => {
338                validate(config_denied_domains, move |candidate| {
339                    let candidate_set: HashSet<String> =
340                        candidate.iter().map(|s| s.to_ascii_lowercase()).collect();
341                    let missing: Vec<String> = required_set
342                        .iter()
343                        .filter(|entry| !candidate_set.contains(*entry))
344                        .cloned()
345                        .collect();
346                    if missing.is_empty() {
347                        Ok(())
348                    } else {
349                        Err(invalid_value(
350                            "network.denied_domains",
351                            "missing managed denied_domains entries",
352                            format!("{missing:?}"),
353                        ))
354                    }
355                })?;
356            }
357        }
358    }
359
360    if let Some(allow_unix_sockets) = &constraints.allow_unix_sockets {
361        let allowed_set: HashSet<String> = allow_unix_sockets
362            .iter()
363            .map(|s| s.to_ascii_lowercase())
364            .collect();
365        validate(config_allow_unix_sockets, move |candidate| {
366            let mut invalid = Vec::new();
367            for entry in candidate {
368                if !allowed_set.contains(&entry.to_ascii_lowercase()) {
369                    invalid.push(entry.clone());
370                }
371            }
372            if invalid.is_empty() {
373                Ok(())
374            } else {
375                Err(invalid_value(
376                    "network.allow_unix_sockets",
377                    format!("{invalid:?}"),
378                    "subset of managed allow_unix_sockets",
379                ))
380            }
381        })?;
382    }
383
384    Ok(())
385}
386
387fn invalid_mitm_hook_configuration(err: anyhow::Error) -> NetworkProxyConstraintError {
388    NetworkProxyConstraintError::InvalidValue {
389        field_name: "network.mitm_hooks",
390        candidate: err.to_string(),
391        allowed: "valid MITM hook configuration".to_string(),
392    }
393}
394
395fn validate_non_global_wildcard_domain_patterns(
396    field_name: &'static str,
397    patterns: &[String],
398) -> Result<(), NetworkProxyConstraintError> {
399    if let Some(pattern) = patterns
400        .iter()
401        .find(|pattern| is_global_wildcard_domain_pattern(pattern))
402    {
403        return Err(NetworkProxyConstraintError::InvalidValue {
404            field_name,
405            candidate: pattern.trim().to_string(),
406            allowed: "exact hosts or scoped wildcards like *.example.com or **.example.com"
407                .to_string(),
408        });
409    }
410    Ok(())
411}
412
413#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
414pub enum NetworkProxyConstraintError {
415    #[error("invalid value for {field_name}: {candidate} (allowed {allowed})")]
416    InvalidValue {
417        field_name: &'static str,
418        candidate: String,
419        allowed: String,
420    },
421}
422
423impl NetworkProxyConstraintError {
424    pub fn into_anyhow(self) -> anyhow::Error {
425        anyhow::anyhow!(self)
426    }
427}
428
429fn network_mode_rank(mode: NetworkMode) -> u8 {
430    match mode {
431        NetworkMode::Limited => 0,
432        NetworkMode::Full => 1,
433    }
434}
435
436#[cfg(test)]
437mod tests {}