1use serde::Deserialize;
4use std::collections::HashMap;
5
6#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Deserialize)]
8#[serde(rename_all = "lowercase")]
9pub enum PolicyMode {
10 #[default]
12 Block,
13 Observe,
15}
16
17#[derive(Debug, Clone, Deserialize)]
19#[serde(untagged)]
20pub enum ScannerSpec {
21 Named(String),
22 Detailed(DetailedScanner),
23}
24
25#[derive(Debug, Clone, Deserialize)]
28pub struct DetailedScanner {
29 #[serde(rename = "type")]
30 pub kind: String,
31 #[serde(default)]
32 pub substrings: Vec<String>,
33 #[serde(default)]
34 pub severity: Option<String>,
35 #[serde(default)]
36 pub max_chars: Option<usize>,
37 #[serde(default)]
38 pub threshold: Option<usize>,
39}
40
41#[derive(Debug, Clone, Deserialize)]
43pub struct Policy {
44 #[serde(default)]
45 pub mode: PolicyMode,
46 pub scanners: Vec<ScannerSpec>,
47}
48
49#[derive(Debug, Clone, Default, Deserialize)]
51pub struct GuardrailsConfig {
52 #[serde(default)]
53 pub guardrails: HashMap<String, Policy>,
54}
55
56impl GuardrailsConfig {
57 pub fn from_toml_str(s: &str) -> anyhow::Result<Self> {
58 toml::from_str(s).map_err(Into::into)
59 }
60}
61
62#[cfg(test)]
63mod tests {
64 use super::*;
65
66 const SAMPLE: &str = r#"
67 [guardrails.default]
68 scanners = [
69 "secrets",
70 { type = "token_limit", max_chars = 32000 },
71 { type = "ban_substrings", substrings = ["BEGIN RSA PRIVATE KEY"], severity = "block" },
72 ]
73 [guardrails.canary]
74 mode = "observe"
75 scanners = ["prompt_injection"]
76 "#;
77
78 #[test]
79 fn parses_named_and_detailed_scanners_and_modes() {
80 let cfg = GuardrailsConfig::from_toml_str(SAMPLE).unwrap();
81 let default = cfg.guardrails.get("default").unwrap();
82 assert_eq!(default.mode, PolicyMode::Block); assert_eq!(default.scanners.len(), 3);
84 assert!(matches!(&default.scanners[0], ScannerSpec::Named(n) if n == "secrets"));
85 match &default.scanners[1] {
86 ScannerSpec::Detailed(d) => {
87 assert_eq!(d.kind, "token_limit");
88 assert_eq!(d.max_chars, Some(32000));
89 }
90 _ => panic!("expected detailed scanner"),
91 }
92 let canary = cfg.guardrails.get("canary").unwrap();
93 assert_eq!(canary.mode, PolicyMode::Observe);
94 }
95
96 #[test]
97 fn empty_config_parses_to_no_policies() {
98 let cfg = GuardrailsConfig::from_toml_str("").unwrap();
99 assert!(cfg.guardrails.is_empty());
100 }
101}