Skip to main content

synapse/guard/
policy.rs

1//! Guardrails config: named policies parsed from `guardrails.toml`.
2
3use serde::Deserialize;
4use std::collections::HashMap;
5
6/// Enforcement mode for a policy.
7#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Deserialize)]
8#[serde(rename_all = "lowercase")]
9pub enum PolicyMode {
10    /// Reject the request when a block-severity scanner fires (default).
11    #[default]
12    Block,
13    /// Never reject; record a would-block and proceed (safe rollout).
14    Observe,
15}
16
17/// One scanner entry: either a bare name (defaults) or a table with params.
18#[derive(Debug, Clone, Deserialize)]
19#[serde(untagged)]
20pub enum ScannerSpec {
21    Named(String),
22    Detailed(DetailedScanner),
23}
24
25/// Table form of a scanner entry. Only the fields relevant to the named
26/// scanner are read; unknown combinations are rejected when the scanner is built.
27#[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/// A named policy: a mode plus an ordered list of scanners.
42#[derive(Debug, Clone, Deserialize)]
43pub struct Policy {
44    #[serde(default)]
45    pub mode: PolicyMode,
46    pub scanners: Vec<ScannerSpec>,
47}
48
49/// Top-level `guardrails.toml`: `[guardrails.<name>]` tables.
50#[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); // default when omitted
83        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}