1use crate::taint::TaintLevel;
8use serde::{Deserialize, Serialize};
9
10#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
12pub struct AllowlistRule {
13 pub tool: String,
15 #[serde(default)]
18 pub to: Vec<String>,
19 #[serde(default)]
21 pub channel: Vec<String>,
22 pub max_sensitivity: TaintLevel,
24}
25
26impl AllowlistRule {
27 pub fn matches(&self, tool_name: &str, target: Option<&str>, taint: TaintLevel) -> bool {
29 if self.tool != tool_name {
30 return false;
31 }
32
33 if taint > self.max_sensitivity {
34 return false;
35 }
36
37 if self.to.is_empty() && self.channel.is_empty() {
39 return true;
40 }
41
42 if let Some(target_str) = target {
44 if self.to.iter().any(|p| pattern_matches(p, target_str)) {
45 return true;
46 }
47 if self.channel.iter().any(|p| pattern_matches(p, target_str)) {
48 return true;
49 }
50 }
51
52 if target.is_none() && (!self.to.is_empty() || !self.channel.is_empty()) {
54 return false;
55 }
56
57 false
58 }
59}
60
61fn pattern_matches(pattern: &str, value: &str) -> bool {
63 if pattern == "*" {
64 return true;
65 }
66 if let Some(suffix) = pattern.strip_prefix('*') {
67 return value.ends_with(suffix);
68 }
69 if let Some(prefix) = pattern.strip_suffix('*') {
70 return value.starts_with(prefix);
71 }
72 pattern == value
73}
74
75#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
77pub struct McpToolOverride {
78 #[serde(default)]
80 pub sensitivity: Option<TaintLevel>,
81 #[serde(default)]
83 pub direction: Option<String>,
84 #[serde(default)]
86 pub clearance: Option<TaintLevel>,
87}
88
89#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
91pub struct PolicyConfig {
92 #[serde(default)]
94 pub allowlist: Vec<AllowlistRule>,
95 #[serde(default)]
102 pub mcp_overrides: std::collections::HashMap<String, McpToolOverride>,
103}
104
105impl PolicyConfig {
106 pub fn from_toml(content: &str) -> Result<Self, String> {
108 let parsed: toml::Value =
110 toml::from_str(content).map_err(|e| format!("Failed to parse policy.toml: {}", e))?;
111
112 let mut config = PolicyConfig::default();
113
114 if let Some(allowlist_arr) = parsed.get("allowlist").and_then(|v| v.as_array()) {
116 for rule_val in allowlist_arr {
117 let tool = rule_val
118 .get("tool")
119 .and_then(|v| v.as_str())
120 .unwrap_or("")
121 .to_string();
122
123 let to: Vec<String> = rule_val
124 .get("to")
125 .and_then(|v| v.as_array())
126 .map(|arr| {
127 arr.iter()
128 .filter_map(|v| v.as_str().map(|s| s.to_string()))
129 .collect()
130 })
131 .unwrap_or_default();
132
133 let channel: Vec<String> = rule_val
134 .get("channel")
135 .and_then(|v| v.as_array())
136 .map(|arr| {
137 arr.iter()
138 .filter_map(|v| v.as_str().map(|s| s.to_string()))
139 .collect()
140 })
141 .unwrap_or_default();
142
143 let max_sensitivity = rule_val
144 .get("max_sensitivity")
145 .and_then(|v| v.as_str())
146 .and_then(TaintLevel::from_str_loose)
147 .unwrap_or(TaintLevel::Public);
148
149 config.allowlist.push(AllowlistRule {
150 tool,
151 to,
152 channel,
153 max_sensitivity,
154 });
155 }
156 }
157
158 if let Some(overrides_table) = parsed.get("mcp_overrides").and_then(|v| v.as_table()) {
175 for (entry_name, entry_val) in overrides_table {
176 match entry_val.get("tools").and_then(|v| v.as_table()) {
177 Some(tools_table) => {
178 for (tool_name, tool_val) in tools_table {
179 let key = crate::mcp_names::advertised_name(entry_name, tool_name);
180 config
181 .mcp_overrides
182 .insert(key, Self::read_override(tool_val));
183 }
184 }
185 None => {
189 if Self::classifies_something(entry_val) {
190 config
191 .mcp_overrides
192 .insert(entry_name.clone(), Self::read_override(entry_val));
193 }
194 }
195 }
196 }
197 }
198
199 Ok(config)
200 }
201
202 fn classifies_something(value: &toml::Value) -> bool {
208 ["sensitivity", "direction", "clearance"]
209 .iter()
210 .any(|field| value.get(field).and_then(|v| v.as_str()).is_some())
211 }
212
213 fn read_override(value: &toml::Value) -> McpToolOverride {
219 McpToolOverride {
220 sensitivity: value
221 .get("sensitivity")
222 .and_then(|v| v.as_str())
223 .and_then(TaintLevel::from_str_loose),
224 direction: value
225 .get("direction")
226 .and_then(|v| v.as_str())
227 .map(|s| s.to_string()),
228 clearance: value
229 .get("clearance")
230 .and_then(|v| v.as_str())
231 .and_then(TaintLevel::from_str_loose),
232 }
233 }
234
235 pub fn check_allowlist(
238 &self,
239 tool_name: &str,
240 target: Option<&str>,
241 taint: TaintLevel,
242 ) -> Option<usize> {
243 self.allowlist
244 .iter()
245 .position(|rule| rule.matches(tool_name, target, taint))
246 }
247}
248
249#[cfg(test)]
250mod tests {
251 use super::*;
252
253 #[test]
256 fn pattern_matches_exact() {
257 assert!(pattern_matches("hello", "hello"));
258 assert!(!pattern_matches("hello", "world"));
259 }
260
261 #[test]
262 fn pattern_matches_wildcard_all() {
263 assert!(pattern_matches("*", "anything"));
264 assert!(pattern_matches("*", ""));
265 }
266
267 #[test]
268 fn pattern_matches_wildcard_prefix() {
269 assert!(pattern_matches("*@example.com", "user@example.com"));
270 assert!(!pattern_matches("*@example.com", "user@other.com"));
271 }
272
273 #[test]
274 fn pattern_matches_wildcard_suffix() {
275 assert!(pattern_matches("megan@*", "megan@anywhere.com"));
276 assert!(!pattern_matches("megan@*", "bob@anywhere.com"));
277 }
278
279 #[test]
282 fn rule_matches_tool_and_sensitivity() {
283 let rule = AllowlistRule {
284 tool: "send_email".into(),
285 to: vec![],
286 channel: vec![],
287 max_sensitivity: TaintLevel::Private,
288 };
289 assert!(rule.matches("send_email", None, TaintLevel::Private));
290 assert!(rule.matches("send_email", None, TaintLevel::Public));
291 assert!(!rule.matches("other_tool", None, TaintLevel::Private));
292 }
293
294 #[test]
295 fn rule_blocks_above_max_sensitivity() {
296 let rule = AllowlistRule {
297 tool: "send_email".into(),
298 to: vec![],
299 channel: vec![],
300 max_sensitivity: TaintLevel::Internal,
301 };
302 assert!(!rule.matches("send_email", None, TaintLevel::Private));
303 }
304
305 #[test]
306 fn rule_matches_target_pattern() {
307 let rule = AllowlistRule {
308 tool: "send_email".into(),
309 to: vec!["megan@*".into(), "+17576306267".into()],
310 channel: vec![],
311 max_sensitivity: TaintLevel::Private,
312 };
313 assert!(rule.matches("send_email", Some("megan@work.com"), TaintLevel::Internal));
314 assert!(rule.matches("send_email", Some("+17576306267"), TaintLevel::Internal));
315 assert!(!rule.matches("send_email", Some("bob@work.com"), TaintLevel::Internal));
316 }
317
318 #[test]
319 fn rule_matches_channel_pattern() {
320 let rule = AllowlistRule {
321 tool: "post_to_slack".into(),
322 to: vec![],
323 channel: vec!["#team-standup".into()],
324 max_sensitivity: TaintLevel::Internal,
325 };
326 assert!(rule.matches("post_to_slack", Some("#team-standup"), TaintLevel::Internal));
327 assert!(!rule.matches("post_to_slack", Some("#general"), TaintLevel::Internal));
328 }
329
330 #[test]
331 fn rule_no_match_when_patterns_but_no_target() {
332 let rule = AllowlistRule {
333 tool: "send_email".into(),
334 to: vec!["megan@*".into()],
335 channel: vec![],
336 max_sensitivity: TaintLevel::Private,
337 };
338 assert!(!rule.matches("send_email", None, TaintLevel::Internal));
339 }
340
341 #[test]
344 fn parse_policy_with_allowlist() {
345 let toml = r##"
346[[allowlist]]
347tool = "send_email"
348to = ["megan@*", "+17576306267"]
349max_sensitivity = "private"
350
351[[allowlist]]
352tool = "post_to_slack"
353channel = ["#team-standup"]
354max_sensitivity = "internal"
355"##;
356 let config = PolicyConfig::from_toml(toml).unwrap();
357 assert_eq!(config.allowlist.len(), 2);
358 assert_eq!(config.allowlist[0].tool, "send_email");
359 assert_eq!(config.allowlist[0].to.len(), 2);
360 assert_eq!(config.allowlist[0].max_sensitivity, TaintLevel::Private);
361 assert_eq!(config.allowlist[1].tool, "post_to_slack");
362 assert_eq!(config.allowlist[1].channel, vec!["#team-standup"]);
363 }
364
365 #[test]
366 fn parse_policy_with_mcp_overrides() {
367 let toml = r#"
368[mcp_overrides."my-server".tools]
369read_customer_data = { sensitivity = "private" }
370search_public_docs = { sensitivity = "public" }
371"#;
372 let config = PolicyConfig::from_toml(toml).unwrap();
373 assert_eq!(config.mcp_overrides.len(), 2);
374 let cust = config
375 .mcp_overrides
376 .get("my-server__read_customer_data")
377 .unwrap();
378 assert_eq!(cust.sensitivity, Some(TaintLevel::Private));
379 let docs = config
380 .mcp_overrides
381 .get("my-server__search_public_docs")
382 .unwrap();
383 assert_eq!(docs.sensitivity, Some(TaintLevel::Public));
384 }
385
386 #[test]
387 fn parse_policy_mcp_override_with_direction_and_clearance() {
388 let toml = r#"
391[mcp_overrides."srv".tools]
392send_email = { sensitivity = "private", direction = "egress", clearance = "public" }
393"#;
394 let config = PolicyConfig::from_toml(toml).unwrap();
395 let ov = config.mcp_overrides.get("srv__send_email").unwrap();
396 assert_eq!(ov.sensitivity, Some(TaintLevel::Private));
397 assert_eq!(ov.direction.as_deref(), Some("egress"));
398 assert_eq!(ov.clearance, Some(TaintLevel::Public));
399 }
400
401 #[test]
402 fn parse_policy_empty() {
403 let config = PolicyConfig::from_toml("").unwrap();
404 assert!(config.allowlist.is_empty());
405 assert!(config.mcp_overrides.is_empty());
406 }
407
408 #[test]
409 fn parse_policy_invalid_toml() {
410 let result = PolicyConfig::from_toml("{{invalid}}");
411 assert!(result.is_err());
412 }
413
414 #[test]
415 fn check_allowlist_returns_matching_index() {
416 let config = PolicyConfig {
417 allowlist: vec![
418 AllowlistRule {
419 tool: "send_email".into(),
420 to: vec!["megan@*".into()],
421 channel: vec![],
422 max_sensitivity: TaintLevel::Private,
423 },
424 AllowlistRule {
425 tool: "post_to_slack".into(),
426 to: vec![],
427 channel: vec![],
428 max_sensitivity: TaintLevel::Internal,
429 },
430 ],
431 mcp_overrides: Default::default(),
432 };
433
434 assert_eq!(
435 config.check_allowlist("send_email", Some("megan@work.com"), TaintLevel::Internal),
436 Some(0)
437 );
438 assert_eq!(
439 config.check_allowlist("post_to_slack", None, TaintLevel::Internal),
440 Some(1)
441 );
442 assert_eq!(
443 config.check_allowlist("unknown", None, TaintLevel::Public),
444 None
445 );
446 }
447
448 #[test]
451 fn allowlist_rule_serde_roundtrip() {
452 let rule = AllowlistRule {
453 tool: "send_email".into(),
454 to: vec!["test@*".into()],
455 channel: vec![],
456 max_sensitivity: TaintLevel::Private,
457 };
458 let json = serde_json::to_string(&rule).unwrap();
459 let back: AllowlistRule = serde_json::from_str(&json).unwrap();
460 assert_eq!(rule, back);
461 }
462
463 #[test]
464 fn mcp_override_serde_roundtrip() {
465 let o = McpToolOverride {
466 sensitivity: Some(TaintLevel::Private),
467 direction: Some("outbound".into()),
468 clearance: Some(TaintLevel::Internal),
469 };
470 let json = serde_json::to_string(&o).unwrap();
471 let back: McpToolOverride = serde_json::from_str(&json).unwrap();
472 assert_eq!(o, back);
473 }
474
475 #[test]
476 fn test_matches_false_when_only_channel_pattern_set_but_no_target() {
477 let rule = AllowlistRule {
478 tool: "post_message".to_string(),
479 to: vec![],
480 channel: vec!["#general".to_string()],
481 max_sensitivity: TaintLevel::Private,
482 };
483 assert!(!rule.matches("post_message", None, TaintLevel::Public));
487 }
488
489 #[test]
490 fn two_policies_compare_by_what_they_say() {
491 let one = PolicyConfig::from_toml("[[allowlist]]\ntool = \"shell\"\n").unwrap();
495 let same = PolicyConfig::from_toml("[[allowlist]]\ntool = \"shell\"\n").unwrap();
496 let other = PolicyConfig::from_toml("[[allowlist]]\ntool = \"web_fetch\"\n").unwrap();
497 assert_eq!(one, same);
498 assert_ne!(one, other);
499 assert_ne!(one, PolicyConfig::default());
500 }
501
502 #[test]
503 fn test_from_toml_mcp_override_server_without_tools_table() {
504 let toml = r#"
507[mcp_overrides.emptyserver]
508note = "no tools declared here"
509"#;
510 let config = PolicyConfig::from_toml(toml).unwrap();
511 assert!(config.mcp_overrides.is_empty());
512 }
513
514 #[test]
529 fn a_nested_override_is_keyed_by_the_name_the_tool_dispatches_under() {
530 let toml = r#"
531[mcp_overrides.tracker.tools.create_issue]
532sensitivity = "internal"
533direction = "outbound"
534clearance = "internal"
535"#;
536 let config = PolicyConfig::from_toml(toml).unwrap();
537 assert_eq!(
538 config.mcp_overrides.keys().collect::<Vec<_>>(),
539 vec!["tracker__create_issue"],
540 "the key must be the advertised name, not a dotted one"
541 );
542 let over = &config.mcp_overrides["tracker__create_issue"];
543 assert_eq!(over.sensitivity, Some(TaintLevel::Internal));
544 assert_eq!(over.direction.as_deref(), Some("outbound"));
545 assert_eq!(over.clearance, Some(TaintLevel::Internal));
546 }
547
548 #[test]
552 fn a_nested_override_sanitizes_a_dotted_server_and_tool() {
553 let toml = r#"
554[mcp_overrides."my.tools".tools."find.all"]
555sensitivity = "private"
556"#;
557 let config = PolicyConfig::from_toml(toml).unwrap();
558 let keys: Vec<&String> = config.mcp_overrides.keys().collect();
559 assert!(
560 config.mcp_overrides.contains_key("my_tools__find_all"),
561 "keys: {keys:?}"
562 );
563 }
564
565 #[test]
569 fn a_policy_file_round_trips_through_serialization() {
570 let mut config = PolicyConfig::default();
571 config.mcp_overrides.insert(
572 "tracker__create_issue".to_string(),
573 McpToolOverride {
574 sensitivity: Some(TaintLevel::Internal),
575 direction: Some("outbound".to_string()),
576 clearance: Some(TaintLevel::Internal),
577 },
578 );
579 let written = toml::to_string_pretty(&config).expect("serializes");
580 let read_back = PolicyConfig::from_toml(&written).expect("parses");
581 assert_eq!(read_back, config, "written as:\n{written}");
582 }
583
584 #[test]
587 fn a_flat_override_keeps_its_name_verbatim() {
588 let toml = r#"
589[mcp_overrides.tracker__create_issue]
590sensitivity = "private"
591"#;
592 let config = PolicyConfig::from_toml(toml).unwrap();
593 assert_eq!(
594 config.mcp_overrides["tracker__create_issue"].sensitivity,
595 Some(TaintLevel::Private)
596 );
597 }
598
599 #[test]
602 fn an_unreadable_level_is_left_unset_rather_than_defaulted() {
603 let toml = r#"
604[mcp_overrides.tracker__create_issue]
605sensitivity = "banana"
606direction = "outbound"
607"#;
608 let config = PolicyConfig::from_toml(toml).unwrap();
609 let over = &config.mcp_overrides["tracker__create_issue"];
610 assert_eq!(over.sensitivity, None);
611 assert_eq!(over.clearance, None);
612 assert_eq!(over.direction.as_deref(), Some("outbound"));
613 }
614}