gateway_core/
guardrail.rs1use 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}