agentshield/egress/policy/
merge.rs1use serde::{Deserialize, Serialize};
2use std::collections::HashMap;
3use std::path::PathBuf;
4
5use super::domain::domain_matches;
6use super::{DomainPolicy, EgressPolicy, NetworkPolicy};
7
8#[derive(Debug, Clone, Serialize, Deserialize)]
10pub struct RateLimitPolicy {
11 #[serde(default = "default_rate_limit")]
13 pub max_requests_per_minute: u32,
14 #[serde(default)]
16 pub per_domain: HashMap<String, u32>,
17}
18
19fn default_rate_limit() -> u32 {
20 60
21}
22
23impl Default for RateLimitPolicy {
24 fn default() -> Self {
25 Self {
26 max_requests_per_minute: default_rate_limit(),
27 per_domain: HashMap::new(),
28 }
29 }
30}
31
32impl RateLimitPolicy {
33 pub(super) fn rate_limit_for(&self, domain: &str) -> u32 {
37 self.per_domain
38 .get(domain)
39 .copied()
40 .unwrap_or(self.max_requests_per_minute)
41 }
42}
43
44#[derive(Debug, Clone, Serialize, Deserialize)]
46pub struct AuditPolicy {
47 #[serde(default)]
49 pub log_path: Option<PathBuf>,
50 #[serde(default = "default_log_format")]
52 pub log_format: String,
53 #[serde(default)]
55 pub log_allowed: bool,
56}
57
58fn default_log_format() -> String {
59 "json".to_string()
60}
61
62impl Default for AuditPolicy {
63 fn default() -> Self {
64 Self {
65 log_path: None,
66 log_format: default_log_format(),
67 log_allowed: false,
68 }
69 }
70}
71
72pub(super) fn merge_override(base: &EgressPolicy, operator: &EgressPolicy) -> EgressPolicy {
84 let allow = if operator.domains.allow.is_empty() {
86 base.domains.allow.clone()
88 } else if base.domains.allow.is_empty() {
89 operator.domains.allow.clone()
91 } else {
92 base.domains
94 .allow
95 .iter()
96 .filter(|d| {
97 operator
98 .domains
99 .allow
100 .iter()
101 .any(|o| domain_matches(d, o) || domain_matches(o, d))
102 })
103 .cloned()
104 .collect()
105 };
106
107 let mut deny = base.domains.deny.clone();
109 for d in &operator.domains.deny {
110 if !deny.contains(d) {
111 deny.push(d.clone());
112 }
113 }
114
115 let global_min = base
117 .rate_limits
118 .max_requests_per_minute
119 .min(operator.rate_limits.max_requests_per_minute);
120
121 let mut per_domain = base.rate_limits.per_domain.clone();
122 for (domain, &op_rate) in &operator.rate_limits.per_domain {
123 let entry = per_domain
124 .entry(domain.clone())
125 .or_insert(base.rate_limits.max_requests_per_minute);
126 *entry = (*entry).min(op_rate);
127 }
128
129 EgressPolicy {
130 schema_version: base.schema_version,
131 domains: DomainPolicy { allow, deny },
132 networks: NetworkPolicy {
133 block_private: base.networks.block_private || operator.networks.block_private,
134 block_link_local: base.networks.block_link_local || operator.networks.block_link_local,
135 block_localhost: base.networks.block_localhost || operator.networks.block_localhost,
136 block_metadata: base.networks.block_metadata || operator.networks.block_metadata,
137 },
138 rate_limits: RateLimitPolicy {
139 max_requests_per_minute: global_min,
140 per_domain,
141 },
142 audit: operator.audit.clone(),
143 }
144}