1use serde::{Deserialize, Serialize};
7use std::collections::HashMap;
8use std::path::{Path, PathBuf};
9
10use crate::error::ShieldError;
11use crate::ir::tool_surface::PermissionType;
12use crate::ir::{ArgumentSource, ScanTarget};
13
14const CURRENT_SCHEMA_VERSION: u32 = 1;
15
16#[derive(Debug, Clone, Serialize, Deserialize)]
18pub struct EgressPolicy {
19 pub schema_version: u32,
21 pub domains: DomainPolicy,
23 #[serde(default)]
25 pub networks: NetworkPolicy,
26 #[serde(default)]
28 pub rate_limits: RateLimitPolicy,
29 #[serde(default)]
31 pub audit: AuditPolicy,
32}
33
34#[derive(Debug, Clone, Serialize, Deserialize)]
36pub struct DomainPolicy {
37 #[serde(default)]
39 pub allow: Vec<String>,
40 #[serde(default)]
42 pub deny: Vec<String>,
43}
44
45#[derive(Debug, Clone, Serialize, Deserialize)]
47pub struct NetworkPolicy {
48 #[serde(default = "default_true")]
50 pub block_private: bool,
51 #[serde(default = "default_true")]
53 pub block_link_local: bool,
54 #[serde(default = "default_true")]
56 pub block_localhost: bool,
57 #[serde(default = "default_true")]
59 pub block_metadata: bool,
60}
61
62fn default_true() -> bool {
63 true
64}
65
66impl Default for NetworkPolicy {
67 fn default() -> Self {
68 Self {
69 block_private: true,
70 block_link_local: true,
71 block_localhost: true,
72 block_metadata: true,
73 }
74 }
75}
76
77#[derive(Debug, Clone, Serialize, Deserialize)]
79pub struct RateLimitPolicy {
80 #[serde(default = "default_rate_limit")]
82 pub max_requests_per_minute: u32,
83 #[serde(default)]
85 pub per_domain: HashMap<String, u32>,
86}
87
88fn default_rate_limit() -> u32 {
89 60
90}
91
92impl Default for RateLimitPolicy {
93 fn default() -> Self {
94 Self {
95 max_requests_per_minute: default_rate_limit(),
96 per_domain: HashMap::new(),
97 }
98 }
99}
100
101#[derive(Debug, Clone, Serialize, Deserialize)]
103pub struct AuditPolicy {
104 #[serde(default)]
106 pub log_path: Option<PathBuf>,
107 #[serde(default = "default_log_format")]
109 pub log_format: String,
110 #[serde(default)]
112 pub log_allowed: bool,
113}
114
115fn default_log_format() -> String {
116 "json".to_string()
117}
118
119impl Default for AuditPolicy {
120 fn default() -> Self {
121 Self {
122 log_path: None,
123 log_format: default_log_format(),
124 log_allowed: false,
125 }
126 }
127}
128
129impl EgressPolicy {
130 pub fn load(path: &Path) -> Result<Self, ShieldError> {
132 let content = std::fs::read_to_string(path).map_err(ShieldError::Io)?;
133 let policy: Self = toml::from_str(&content)?;
134 if policy.schema_version > CURRENT_SCHEMA_VERSION {
135 return Err(ShieldError::Config(format!(
136 "Egress policy schema version {} is newer than supported version {}",
137 policy.schema_version, CURRENT_SCHEMA_VERSION
138 )));
139 }
140 Ok(policy)
141 }
142
143 pub fn save(&self, path: &Path) -> Result<(), ShieldError> {
145 let content = toml::to_string_pretty(self)?;
146 std::fs::write(path, content).map_err(ShieldError::Io)?;
147 Ok(())
148 }
149
150 pub fn is_domain_allowed(&self, domain: &str) -> bool {
155 if self
157 .domains
158 .deny
159 .iter()
160 .any(|pattern| domain_matches(domain, pattern))
161 {
162 return false;
163 }
164 if self.domains.allow.is_empty() {
166 return true;
167 }
168 self.domains
170 .allow
171 .iter()
172 .any(|pattern| domain_matches(domain, pattern))
173 }
174
175 pub fn is_ip_blocked(&self, ip: &str) -> bool {
177 if self.networks.block_localhost && is_localhost(ip) {
178 return true;
179 }
180 if self.networks.block_private && is_private_ip(ip) {
181 return true;
182 }
183 if self.networks.block_link_local && is_link_local(ip) {
184 return true;
185 }
186 if self.networks.block_metadata && is_metadata_ip(ip) {
187 return true;
188 }
189 false
190 }
191
192 pub fn rate_limit_for(&self, domain: &str) -> u32 {
196 self.rate_limits
197 .per_domain
198 .get(domain)
199 .copied()
200 .unwrap_or(self.rate_limits.max_requests_per_minute)
201 }
202
203 pub fn from_scan_targets(targets: &[ScanTarget]) -> Self {
212 let mut domains = std::collections::HashSet::new();
213
214 for target in targets {
215 for net_op in &target.execution.network_operations {
217 if let ArgumentSource::Literal(ref url) = net_op.url_arg {
218 if let Some(domain) = extract_domain(url) {
219 domains.insert(domain);
220 }
221 }
222 }
223
224 for tool in &target.tools {
226 for perm in &tool.declared_permissions {
227 if matches!(perm.permission_type, PermissionType::NetworkAccess) {
228 if let Some(ref scope) = perm.target {
229 if let Some(domain) = extract_domain(scope) {
230 domains.insert(domain);
231 }
232 }
233 }
234 }
235 }
236 }
237
238 let mut allow: Vec<String> = domains.into_iter().collect();
239 allow.sort();
240
241 EgressPolicy {
242 schema_version: CURRENT_SCHEMA_VERSION,
243 domains: DomainPolicy {
244 allow,
245 deny: vec![],
246 },
247 networks: NetworkPolicy::default(),
248 rate_limits: RateLimitPolicy::default(),
249 audit: AuditPolicy::default(),
250 }
251 }
252
253 pub fn merge_override(&self, operator: &EgressPolicy) -> EgressPolicy {
265 let allow = if operator.domains.allow.is_empty() {
267 self.domains.allow.clone()
269 } else if self.domains.allow.is_empty() {
270 operator.domains.allow.clone()
272 } else {
273 self.domains
275 .allow
276 .iter()
277 .filter(|d| {
278 operator
279 .domains
280 .allow
281 .iter()
282 .any(|o| domain_matches(d, o) || domain_matches(o, d))
283 })
284 .cloned()
285 .collect()
286 };
287
288 let mut deny = self.domains.deny.clone();
290 for d in &operator.domains.deny {
291 if !deny.contains(d) {
292 deny.push(d.clone());
293 }
294 }
295
296 let global_min = self
298 .rate_limits
299 .max_requests_per_minute
300 .min(operator.rate_limits.max_requests_per_minute);
301
302 let mut per_domain = self.rate_limits.per_domain.clone();
303 for (domain, &op_rate) in &operator.rate_limits.per_domain {
304 let entry = per_domain
305 .entry(domain.clone())
306 .or_insert(self.rate_limits.max_requests_per_minute);
307 *entry = (*entry).min(op_rate);
308 }
309
310 EgressPolicy {
311 schema_version: self.schema_version,
312 domains: DomainPolicy { allow, deny },
313 networks: NetworkPolicy {
314 block_private: self.networks.block_private || operator.networks.block_private,
315 block_link_local: self.networks.block_link_local
316 || operator.networks.block_link_local,
317 block_localhost: self.networks.block_localhost || operator.networks.block_localhost,
318 block_metadata: self.networks.block_metadata || operator.networks.block_metadata,
319 },
320 rate_limits: RateLimitPolicy {
321 max_requests_per_minute: global_min,
322 per_domain,
323 },
324 audit: operator.audit.clone(),
325 }
326 }
327
328 pub fn starter_toml() -> &'static str {
330 r#"# AgentShield Egress Policy
331# See: https://github.com/aiconnai/agentshield
332
333schema_version = 1
334
335[domains]
336# Allowed domain patterns (glob-style)
337allow = ["*.example.com", "api.github.com"]
338# Explicitly denied (takes precedence over allow)
339deny = []
340
341[networks]
342block_private = true # 10.x, 172.16-31.x, 192.168.x
343block_link_local = true # 169.254.x
344block_localhost = true # 127.x, ::1
345block_metadata = true # 169.254.169.254, metadata.google.internal
346
347[rate_limits]
348max_requests_per_minute = 60
349
350[audit]
351# log_path = "agentshield-audit.jsonl"
352log_format = "json"
353log_allowed = false
354"#
355 }
356}
357
358pub fn extract_domain(url_or_domain: &str) -> Option<String> {
364 let rest = if let Some(r) = url_or_domain.strip_prefix("https://") {
366 r
367 } else if let Some(r) = url_or_domain.strip_prefix("http://") {
368 r
369 } else {
370 if url_or_domain.contains('.') && !url_or_domain.contains('/') {
372 return Some(url_or_domain.to_string());
373 }
374 return None;
375 };
376
377 let host = rest.split('/').next()?;
379 let host = host.split(':').next()?;
381
382 if host.is_empty() {
383 return None;
384 }
385 Some(host.to_string())
386}
387
388fn domain_matches(domain: &str, pattern: &str) -> bool {
393 if let Some(suffix) = pattern.strip_prefix('*') {
394 domain.ends_with(suffix) || domain == &suffix[1..]
396 } else {
397 domain == pattern
398 }
399}
400
401fn is_localhost(ip: &str) -> bool {
402 ip.starts_with("127.") || ip == "::1" || ip == "localhost"
403}
404
405fn is_private_ip(ip: &str) -> bool {
406 ip.starts_with("10.")
407 || (ip.starts_with("172.") && is_172_private(ip))
408 || ip.starts_with("192.168.")
409 || ip.starts_with("fd") }
411
412fn is_172_private(ip: &str) -> bool {
413 if let Some(second_octet) = ip
414 .strip_prefix("172.")
415 .and_then(|rest| rest.split('.').next())
416 {
417 if let Ok(n) = second_octet.parse::<u8>() {
418 return (16..=31).contains(&n);
419 }
420 }
421 false
422}
423
424fn is_link_local(ip: &str) -> bool {
425 ip.starts_with("169.254.") || ip.starts_with("fe80:")
426}
427
428fn is_metadata_ip(ip: &str) -> bool {
429 ip == "169.254.169.254"
430 || ip.contains("metadata.google.internal")
431 || ip == "100.100.100.200" || ip == "169.254.170.2" }
434
435#[cfg(test)]
436mod tests {
437 use super::*;
438 use tempfile::TempDir;
439
440 fn sample_policy() -> EgressPolicy {
441 EgressPolicy {
442 schema_version: 1,
443 domains: DomainPolicy {
444 allow: vec!["*.example.com".into(), "api.github.com".into()],
445 deny: vec!["evil.example.com".into()],
446 },
447 networks: NetworkPolicy::default(),
448 rate_limits: RateLimitPolicy {
449 max_requests_per_minute: 60,
450 per_domain: {
451 let mut m = HashMap::new();
452 m.insert("api.github.com".into(), 30);
453 m
454 },
455 },
456 audit: AuditPolicy::default(),
457 }
458 }
459
460 #[test]
461 fn test_load_and_save_roundtrip() {
462 let tmp = TempDir::new().unwrap();
463 let path = tmp.path().join("egress.toml");
464
465 let original = sample_policy();
466 original.save(&path).unwrap();
467
468 let loaded = EgressPolicy::load(&path).unwrap();
469
470 assert_eq!(loaded.schema_version, original.schema_version);
471 assert_eq!(loaded.domains.allow, original.domains.allow);
472 assert_eq!(loaded.domains.deny, original.domains.deny);
473 assert_eq!(
474 loaded.networks.block_private,
475 original.networks.block_private
476 );
477 assert_eq!(
478 loaded.networks.block_localhost,
479 original.networks.block_localhost
480 );
481 assert_eq!(
482 loaded.networks.block_link_local,
483 original.networks.block_link_local
484 );
485 assert_eq!(
486 loaded.networks.block_metadata,
487 original.networks.block_metadata
488 );
489 assert_eq!(
490 loaded.rate_limits.max_requests_per_minute,
491 original.rate_limits.max_requests_per_minute
492 );
493 assert_eq!(
494 loaded.rate_limits.per_domain,
495 original.rate_limits.per_domain
496 );
497 assert_eq!(loaded.audit.log_format, original.audit.log_format);
498 assert_eq!(loaded.audit.log_allowed, original.audit.log_allowed);
499 assert_eq!(loaded.audit.log_path, original.audit.log_path);
500 }
501
502 #[test]
503 fn test_domain_allowed() {
504 let policy = sample_policy();
505
506 assert!(policy.is_domain_allowed("api.github.com"));
508 assert!(policy.is_domain_allowed("sub.example.com"));
510 assert!(policy.is_domain_allowed("example.com"));
512 assert!(!policy.is_domain_allowed("random.org"));
514 }
515
516 #[test]
517 fn test_domain_denied_takes_precedence() {
518 let policy = sample_policy();
519
520 assert!(
522 !policy.is_domain_allowed("evil.example.com"),
523 "deny should take precedence over allow"
524 );
525 }
526
527 #[test]
528 fn test_empty_allow_list_allows_all() {
529 let policy = EgressPolicy {
530 schema_version: 1,
531 domains: DomainPolicy {
532 allow: vec![],
533 deny: vec!["blocked.com".into()],
534 },
535 networks: NetworkPolicy::default(),
536 rate_limits: RateLimitPolicy::default(),
537 audit: AuditPolicy::default(),
538 };
539
540 assert!(policy.is_domain_allowed("anything.com"));
541 assert!(policy.is_domain_allowed("whatever.org"));
542 assert!(
543 !policy.is_domain_allowed("blocked.com"),
544 "deny should still block even with empty allow"
545 );
546 }
547
548 #[test]
549 fn test_ip_blocking() {
550 let policy = sample_policy();
551
552 assert!(policy.is_ip_blocked("127.0.0.1"));
554 assert!(policy.is_ip_blocked("127.0.0.2"));
555 assert!(policy.is_ip_blocked("::1"));
556 assert!(policy.is_ip_blocked("localhost"));
557
558 assert!(policy.is_ip_blocked("10.0.0.1"));
560 assert!(policy.is_ip_blocked("172.16.0.1"));
561 assert!(policy.is_ip_blocked("172.31.255.255"));
562 assert!(policy.is_ip_blocked("192.168.1.1"));
563
564 assert!(!policy.is_ip_blocked("172.15.0.1"));
566 assert!(!policy.is_ip_blocked("172.32.0.1"));
567
568 assert!(policy.is_ip_blocked("169.254.1.1"));
570 assert!(policy.is_ip_blocked("fe80::1"));
571
572 assert!(policy.is_ip_blocked("169.254.169.254"));
574 assert!(policy.is_ip_blocked("metadata.google.internal"));
575 assert!(policy.is_ip_blocked("100.100.100.200"));
576 assert!(policy.is_ip_blocked("169.254.170.2"));
577
578 assert!(!policy.is_ip_blocked("8.8.8.8"));
580 assert!(!policy.is_ip_blocked("1.1.1.1"));
581 }
582
583 #[test]
584 fn test_rate_limit_per_domain() {
585 let policy = sample_policy();
586 assert_eq!(policy.rate_limit_for("api.github.com"), 30);
587 }
588
589 #[test]
590 fn test_rate_limit_default() {
591 let policy = sample_policy();
592 assert_eq!(policy.rate_limit_for("unknown.com"), 60);
593 }
594
595 #[test]
596 fn test_future_schema_rejected() {
597 let tmp = TempDir::new().unwrap();
598 let path = tmp.path().join("future.toml");
599
600 let content = r#"
601schema_version = 99
602
603[domains]
604allow = []
605deny = []
606"#;
607 std::fs::write(&path, content).unwrap();
608
609 let result = EgressPolicy::load(&path);
610 assert!(result.is_err());
611
612 let err_msg = result.unwrap_err().to_string();
613 assert!(
614 err_msg.contains("99") && err_msg.contains("newer"),
615 "Error should mention unsupported schema version, got: {err_msg}"
616 );
617 }
618
619 #[test]
620 fn test_starter_toml_parses() {
621 let toml_str = EgressPolicy::starter_toml();
622 let policy: EgressPolicy =
623 toml::from_str(toml_str).expect("starter_toml() should produce valid TOML");
624 assert_eq!(policy.schema_version, 1);
625 assert!(!policy.domains.allow.is_empty());
626 assert!(policy.networks.block_private);
627 assert!(policy.networks.block_metadata);
628 assert_eq!(policy.rate_limits.max_requests_per_minute, 60);
629 assert_eq!(policy.audit.log_format, "json");
630 }
631
632 #[test]
635 fn test_extract_domain_from_url() {
636 assert_eq!(
638 extract_domain("https://api.example.com/v1/items"),
639 Some("api.example.com".into())
640 );
641 assert_eq!(
642 extract_domain("http://api.example.com:8080/path"),
643 Some("api.example.com".into())
644 );
645 assert_eq!(
646 extract_domain("https://api.github.com"),
647 Some("api.github.com".into())
648 );
649 assert_eq!(
651 extract_domain("api.example.com"),
652 Some("api.example.com".into())
653 );
654 assert_eq!(extract_domain("localhost"), None);
656 assert_eq!(extract_domain("/some/path"), None);
658 assert_eq!(extract_domain(""), None);
660 }
661
662 #[test]
663 fn test_from_scan_targets_extracts_domains() {
664 use crate::ir::execution_surface::{ExecutionSurface, NetworkOperation};
665 use crate::ir::tool_surface::{DeclaredPermission, PermissionType, ToolSurface};
666 use crate::ir::{
667 ArgumentSource, DataSurface, DependencySurface, Framework, ProvenanceSurface,
668 ScanTarget, SourceLocation,
669 };
670 use std::path::PathBuf;
671
672 let make_loc = || SourceLocation {
673 file: PathBuf::from("server.py"),
674 line: 1,
675 column: 0,
676 end_line: None,
677 end_column: None,
678 };
679
680 let target = ScanTarget {
681 name: "test-server".into(),
682 framework: Framework::Mcp,
683 root_path: PathBuf::from("/tmp/test"),
684 tools: vec![ToolSurface {
685 name: "fetch_data".into(),
686 description: None,
687 input_schema: None,
688 output_schema: None,
689 declared_permissions: vec![DeclaredPermission {
690 permission_type: PermissionType::NetworkAccess,
691 target: Some("https://api.stripe.com/v1".into()),
692 description: None,
693 }],
694 defined_at: None,
695 declared_capabilities: Default::default(),
696 capability_declarations: Vec::new(),
697 observed_capabilities: Default::default(),
698 capability_observation_complete: false,
699 capability_evidence: Vec::new(),
700 }],
701 execution: ExecutionSurface {
702 network_operations: vec![
703 NetworkOperation {
704 function: "requests.get".into(),
705 url_arg: ArgumentSource::Literal("https://api.openai.com/v1/chat".into()),
706 method: Some("GET".into()),
707 sends_data: false,
708 location: make_loc(),
709 },
710 NetworkOperation {
711 function: "requests.post".into(),
712 url_arg: ArgumentSource::Parameter { name: "url".into() },
714 method: Some("POST".into()),
715 sends_data: true,
716 location: make_loc(),
717 },
718 ],
719 ..ExecutionSurface::default()
720 },
721 data: DataSurface::default(),
722 dependencies: DependencySurface::default(),
723 provenance: ProvenanceSurface::default(),
724 source_files: vec![],
725 };
726
727 let policy = EgressPolicy::from_scan_targets(&[target]);
728
729 assert_eq!(policy.schema_version, 1);
731 assert!(policy.domains.deny.is_empty());
733 assert!(
735 policy.domains.allow.contains(&"api.openai.com".to_string()),
736 "Expected api.openai.com in allow list, got: {:?}",
737 policy.domains.allow
738 );
739 assert!(
740 policy.domains.allow.contains(&"api.stripe.com".to_string()),
741 "Expected api.stripe.com in allow list, got: {:?}",
742 policy.domains.allow
743 );
744 assert_eq!(
746 policy.domains.allow,
747 {
748 let mut sorted = policy.domains.allow.clone();
749 sorted.sort();
750 sorted
751 },
752 "Allow list should be sorted"
753 );
754 assert!(policy.networks.block_private);
756 assert!(policy.networks.block_localhost);
757 assert!(policy.networks.block_link_local);
758 assert!(policy.networks.block_metadata);
759 assert_eq!(policy.rate_limits.max_requests_per_minute, 60);
761 }
762
763 fn base_policy() -> EgressPolicy {
766 EgressPolicy {
767 schema_version: 1,
768 domains: DomainPolicy {
769 allow: vec![
770 "api.example.com".into(),
771 "api.github.com".into(),
772 "api.openai.com".into(),
773 ],
774 deny: vec!["evil.com".into()],
775 },
776 networks: NetworkPolicy {
777 block_private: false,
778 block_link_local: true,
779 block_localhost: true,
780 block_metadata: false,
781 },
782 rate_limits: RateLimitPolicy {
783 max_requests_per_minute: 60,
784 per_domain: {
785 let mut m = HashMap::new();
786 m.insert("api.openai.com".into(), 20);
787 m
788 },
789 },
790 audit: AuditPolicy {
791 log_path: Some(PathBuf::from("/tmp/base-audit.jsonl")),
792 log_format: "json".into(),
793 log_allowed: false,
794 },
795 }
796 }
797
798 #[test]
799 fn test_merge_deny_union() {
800 let base = base_policy();
801 let operator = EgressPolicy {
802 schema_version: 1,
803 domains: DomainPolicy {
804 allow: vec![],
805 deny: vec!["extra-bad.com".into()],
806 },
807 networks: NetworkPolicy::default(),
808 rate_limits: RateLimitPolicy::default(),
809 audit: AuditPolicy::default(),
810 };
811
812 let merged = base.merge_override(&operator);
813
814 assert!(
815 merged.domains.deny.contains(&"evil.com".to_string()),
816 "base deny entry must be preserved"
817 );
818 assert!(
819 merged.domains.deny.contains(&"extra-bad.com".to_string()),
820 "operator deny entry must be added"
821 );
822 assert_eq!(merged.domains.deny.len(), 2);
823 }
824
825 #[test]
826 fn test_merge_allow_intersection() {
827 let base = base_policy();
828 let operator = EgressPolicy {
829 schema_version: 1,
830 domains: DomainPolicy {
831 allow: vec![
833 "api.github.com".into(),
834 "api.openai.com".into(),
835 "api.stripe.com".into(),
836 ],
837 deny: vec![],
838 },
839 networks: NetworkPolicy::default(),
840 rate_limits: RateLimitPolicy::default(),
841 audit: AuditPolicy::default(),
842 };
843
844 let merged = base.merge_override(&operator);
845
846 assert!(
847 merged.domains.allow.contains(&"api.github.com".to_string()),
848 "intersection: api.github.com must be in result"
849 );
850 assert!(
851 merged.domains.allow.contains(&"api.openai.com".to_string()),
852 "intersection: api.openai.com must be in result"
853 );
854 assert!(
855 !merged
856 .domains
857 .allow
858 .contains(&"api.example.com".to_string()),
859 "api.example.com not in operator allow → must be excluded"
860 );
861 assert!(
862 !merged.domains.allow.contains(&"api.stripe.com".to_string()),
863 "api.stripe.com not in base allow → must be excluded"
864 );
865 }
866
867 #[test]
868 fn test_merge_rate_limits_min() {
869 let base = base_policy(); let operator = EgressPolicy {
871 schema_version: 1,
872 domains: DomainPolicy {
873 allow: vec![],
874 deny: vec![],
875 },
876 networks: NetworkPolicy::default(),
877 rate_limits: RateLimitPolicy {
878 max_requests_per_minute: 30,
879 per_domain: {
880 let mut m = HashMap::new();
881 m.insert("api.openai.com".into(), 10);
882 m.insert("api.github.com".into(), 5);
883 m
884 },
885 },
886 audit: AuditPolicy::default(),
887 };
888
889 let merged = base.merge_override(&operator);
890
891 assert_eq!(
892 merged.rate_limits.max_requests_per_minute, 30,
893 "global rate: min(60, 30) = 30"
894 );
895 assert_eq!(
896 merged.rate_limits.per_domain["api.openai.com"], 10,
897 "per-domain rate: min(20, 10) = 10"
898 );
899 assert_eq!(
900 merged.rate_limits.per_domain["api.github.com"], 5,
901 "operator-only per-domain: min(60, 5) = 5"
902 );
903 }
904
905 #[test]
906 fn test_merge_network_blocks_or() {
907 let base = base_policy(); let operator = EgressPolicy {
909 schema_version: 1,
910 domains: DomainPolicy {
911 allow: vec![],
912 deny: vec![],
913 },
914 networks: NetworkPolicy {
915 block_private: true,
916 block_link_local: false,
917 block_localhost: false,
918 block_metadata: true,
919 },
920 rate_limits: RateLimitPolicy::default(),
921 audit: AuditPolicy::default(),
922 };
923
924 let merged = base.merge_override(&operator);
925
926 assert!(merged.networks.block_private, "false || true = true");
927 assert!(
928 merged.networks.block_link_local,
929 "true || false = true (base had it)"
930 );
931 assert!(
932 merged.networks.block_localhost,
933 "true || false = true (base had it)"
934 );
935 assert!(merged.networks.block_metadata, "false || true = true");
936 }
937
938 #[test]
939 fn test_merge_empty_override_allow_keeps_base() {
940 let base = base_policy(); let operator = EgressPolicy {
942 schema_version: 1,
943 domains: DomainPolicy {
944 allow: vec![], deny: vec![],
946 },
947 networks: NetworkPolicy::default(),
948 rate_limits: RateLimitPolicy::default(),
949 audit: AuditPolicy::default(),
950 };
951
952 let merged = base.merge_override(&operator);
953
954 assert_eq!(
955 merged.domains.allow, base.domains.allow,
956 "empty operator allow must not restrict base allow list"
957 );
958 }
959
960 #[test]
961 fn test_merge_audit_override_wins() {
962 let base = base_policy(); let operator = EgressPolicy {
964 schema_version: 1,
965 domains: DomainPolicy {
966 allow: vec![],
967 deny: vec![],
968 },
969 networks: NetworkPolicy::default(),
970 rate_limits: RateLimitPolicy::default(),
971 audit: AuditPolicy {
972 log_path: Some(PathBuf::from("/var/log/agentshield/operator.jsonl")),
973 log_format: "text".into(),
974 log_allowed: true,
975 },
976 };
977
978 let merged = base.merge_override(&operator);
979
980 assert_eq!(
981 merged.audit.log_path,
982 Some(PathBuf::from("/var/log/agentshield/operator.jsonl")),
983 "operator audit log_path must win"
984 );
985 assert_eq!(
986 merged.audit.log_format, "text",
987 "operator audit log_format must win"
988 );
989 assert!(
990 merged.audit.log_allowed,
991 "operator audit log_allowed must win"
992 );
993 }
994
995 #[test]
996 fn test_emit_egress_policy_integration() {
997 use crate::{ScanOptions, scan};
1000 use std::path::Path;
1001
1002 let opts = ScanOptions::default();
1003 let report = scan(Path::new("tests/fixtures/mcp_servers/vuln_ssrf"), &opts)
1004 .expect("scan should succeed");
1005
1006 let policy = EgressPolicy::from_scan_targets(&report.targets);
1007
1008 let tmp = TempDir::new().unwrap();
1010 let policy_path = tmp.path().join("agentshield.egress.toml");
1011 policy.save(&policy_path).unwrap();
1012
1013 let loaded = EgressPolicy::load(&policy_path).unwrap();
1014 assert_eq!(loaded.schema_version, 1);
1015 assert!(loaded.networks.block_private);
1016 assert!(loaded.networks.block_metadata);
1017 assert!(loaded.domains.deny.is_empty());
1019 }
1020}