Skip to main content

agentshield/egress/policy/
merge.rs

1use serde::{Deserialize, Serialize};
2use std::collections::HashMap;
3use std::path::PathBuf;
4
5use super::domain::domain_matches;
6use super::{DomainPolicy, EgressPolicy, NetworkPolicy};
7
8/// Rate limiting configuration for outbound requests.
9#[derive(Debug, Clone, Serialize, Deserialize)]
10pub struct RateLimitPolicy {
11    /// Maximum requests per minute per domain. 0 = unlimited.
12    #[serde(default = "default_rate_limit")]
13    pub max_requests_per_minute: u32,
14    /// Per-domain overrides (domain string -> requests per minute).
15    #[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    /// Get rate limit for a domain (requests per minute).
34    ///
35    /// Returns the per-domain override if one exists, otherwise the global default.
36    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/// Audit logging configuration for egress events.
45#[derive(Debug, Clone, Serialize, Deserialize)]
46pub struct AuditPolicy {
47    /// Path to write audit log.
48    #[serde(default)]
49    pub log_path: Option<PathBuf>,
50    /// Log format: `"json"` or `"text"`.
51    #[serde(default = "default_log_format")]
52    pub log_format: String,
53    /// Log allowed requests too (not just blocked). Default: false.
54    #[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
72/// Merge with an operator override policy. The override can only restrict, never expand.
73///
74/// Merge rules:
75/// - `domains.allow` = intersection(base.allow, override.allow)
76///   If override.allow is empty, base.allow is kept (empty means "no restriction").
77///   If base.allow is empty (allow all), operator's allow list becomes the effective list.
78/// - `domains.deny` = union(base.deny, override.deny)
79/// - `networks`: if either policy blocks a range, it is blocked in the result
80/// - `rate_limits.max_requests_per_minute` = min(self, override)
81/// - `rate_limits.per_domain`: min rate per domain; missing entries inherit the global min
82/// - `audit`: operator override wins (operator controls where logs go)
83pub(super) fn merge_override(base: &EgressPolicy, operator: &EgressPolicy) -> EgressPolicy {
84    // Allow list: intersection when both are non-empty; operator restricts further
85    let allow = if operator.domains.allow.is_empty() {
86        // Empty override allow = "no additional restriction on allow"
87        base.domains.allow.clone()
88    } else if base.domains.allow.is_empty() {
89        // Self allows all; operator restricts to its list
90        operator.domains.allow.clone()
91    } else {
92        // Both have allow lists: intersection (only domains in BOTH lists)
93        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    // Deny list: union (operator can only add more denials)
108    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    // Rate limits: take the minimum (more restrictive wins)
116    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}