Skip to main content

gateway_core/
guardrail.rs

1use std::collections::HashMap;
2
3use regex::Regex;
4use serde::{Deserialize, Serialize};
5
6#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
7#[serde(rename_all = "snake_case")]
8pub enum GuardrailAction {
9    Block,
10    Redact,
11}
12
13#[derive(Debug, Clone, PartialEq, Eq)]
14pub struct GuardrailRequest<'a> {
15    pub scope_id: &'a str,
16    pub prompts: &'a [String],
17}
18
19#[derive(Debug, Clone, PartialEq, Eq)]
20pub enum GuardrailVerdict {
21    Allow,
22    Block {
23        policy_id: String,
24    },
25    Redact {
26        policy_ids: Vec<String>,
27        prompts: Vec<String>,
28    },
29}
30
31pub trait Guardrail: Send + Sync {
32    fn inspect(&self, request: GuardrailRequest<'_>) -> GuardrailVerdict;
33}
34
35#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
36pub struct GuardrailRule {
37    pub id: String,
38    pub pattern: String,
39    pub action: GuardrailAction,
40    #[serde(default = "default_redaction")]
41    pub redaction: String,
42}
43
44#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
45pub struct GuardrailPolicy {
46    pub scope_id: String,
47    pub rules: Vec<GuardrailRule>,
48}
49
50struct CompiledRule {
51    id: String,
52    regex: Regex,
53    action: GuardrailAction,
54    redaction: String,
55}
56
57pub struct RegexGuardrail {
58    policies: HashMap<String, Vec<CompiledRule>>,
59}
60
61impl RegexGuardrail {
62    pub fn compile(policies: &[GuardrailPolicy]) -> Result<Self, regex::Error> {
63        let mut compiled = HashMap::new();
64        for policy in policies {
65            let rules = policy
66                .rules
67                .iter()
68                .map(|rule| {
69                    Ok(CompiledRule {
70                        id: rule.id.clone(),
71                        regex: Regex::new(&rule.pattern)?,
72                        action: rule.action,
73                        redaction: rule.redaction.clone(),
74                    })
75                })
76                .collect::<Result<_, regex::Error>>()?;
77            compiled.insert(policy.scope_id.clone(), rules);
78        }
79        Ok(Self { policies: compiled })
80    }
81}
82
83impl Guardrail for RegexGuardrail {
84    fn inspect(&self, request: GuardrailRequest<'_>) -> GuardrailVerdict {
85        let Some(rules) = self
86            .policies
87            .get(request.scope_id)
88            .or_else(|| self.policies.get("*"))
89        else {
90            return GuardrailVerdict::Allow;
91        };
92        for rule in rules
93            .iter()
94            .filter(|rule| rule.action == GuardrailAction::Block)
95        {
96            if request
97                .prompts
98                .iter()
99                .any(|prompt| rule.regex.is_match(prompt))
100            {
101                return GuardrailVerdict::Block {
102                    policy_id: rule.id.clone(),
103                };
104            }
105        }
106        let mut prompts = request.prompts.to_vec();
107        let mut policy_ids = Vec::new();
108        for rule in rules
109            .iter()
110            .filter(|rule| rule.action == GuardrailAction::Redact)
111        {
112            let mut matched = false;
113            for prompt in &mut prompts {
114                let replaced = rule
115                    .regex
116                    .replace_all(prompt, regex::NoExpand(&rule.redaction));
117                if replaced != prompt.as_str() {
118                    *prompt = replaced.into_owned();
119                    matched = true;
120                }
121            }
122            if matched {
123                policy_ids.push(rule.id.clone());
124            }
125        }
126        if policy_ids.is_empty() {
127            GuardrailVerdict::Allow
128        } else {
129            GuardrailVerdict::Redact {
130                policy_ids,
131                prompts,
132            }
133        }
134    }
135}
136
137fn default_redaction() -> String {
138    "[REDACTED]".into()
139}
140
141#[cfg(test)]
142mod tests {
143    use super::*;
144
145    #[test]
146    fn block_precedes_redaction_and_scope_policy_falls_back() {
147        let guardrail = RegexGuardrail::compile(&[GuardrailPolicy {
148            scope_id: "*".into(),
149            rules: vec![
150                GuardrailRule {
151                    id: "email".into(),
152                    pattern: r"\S+@\S+".into(),
153                    action: GuardrailAction::Redact,
154                    redaction: "[EMAIL]".into(),
155                },
156                GuardrailRule {
157                    id: "deny".into(),
158                    pattern: "forbidden".into(),
159                    action: GuardrailAction::Block,
160                    redaction: default_redaction(),
161                },
162            ],
163        }])
164        .unwrap();
165        assert!(matches!(
166            guardrail.inspect(GuardrailRequest {
167                scope_id: "scope-a",
168                prompts: &["a@example.com forbidden".into()],
169            }),
170            GuardrailVerdict::Block { policy_id } if policy_id == "deny"
171        ));
172        assert_eq!(
173            guardrail.inspect(GuardrailRequest {
174                scope_id: "scope-a",
175                prompts: &["email a@example.com".into()],
176            }),
177            GuardrailVerdict::Redact {
178                policy_ids: vec!["email".into()],
179                prompts: vec!["email [EMAIL]".into()],
180            }
181        );
182    }
183}