1use std::collections::HashMap;
5use std::time::Instant;
6
7use llm_guard::{Pipeline, PipelineMode, ScanResult, Severity};
8use serde_json::Value;
9
10use crate::error::GatewayError;
11use crate::routing::request::{ChatRequest, Message};
12
13use super::policy::{GuardrailsConfig, PolicyMode};
14use super::scanners::{build_scanners, BoxedScanner};
15
16struct CompiledPolicy {
17 mode: PolicyMode,
18 pipeline: Pipeline,
19}
20
21#[derive(Default)]
24pub struct GuardEngine {
25 policies: HashMap<String, CompiledPolicy>,
26}
27
28impl GuardEngine {
29 pub fn empty() -> Self {
31 Self::default()
32 }
33
34 pub fn from_config(cfg: &GuardrailsConfig) -> anyhow::Result<Self> {
37 let mut policies = HashMap::new();
38 for (name, policy) in &cfg.guardrails {
39 let mut pipeline = Pipeline::new(PipelineMode::All);
40 for spec in &policy.scanners {
41 for scanner in build_scanners(spec)? {
42 pipeline = pipeline.with(BoxedScanner(scanner));
43 }
44 }
45 policies.insert(
46 name.clone(),
47 CompiledPolicy {
48 mode: policy.mode,
49 pipeline,
50 },
51 );
52 }
53 Ok(Self { policies })
54 }
55
56 pub fn guard(&self, policy_name: &str, req: &ChatRequest) -> Result<(), GatewayError> {
61 let Some(policy) = self.policies.get(policy_name) else {
62 return Ok(());
63 };
64 let text = input_text(req);
65 let started = Instant::now();
66 let result = policy.pipeline.scan(&text);
67 let outcome = outcome_label(&result, policy.mode);
68 record_metrics(policy_name, &result, outcome, started);
69
70 if result.should_refuse() && policy.mode == PolicyMode::Block {
71 return Err(GatewayError::ContentBlocked {
72 policy: policy_name.to_string(),
73 scanners: blocking_scanners(&result),
74 });
75 }
76 Ok(())
77 }
78}
79
80fn input_text(req: &ChatRequest) -> String {
83 req.messages
84 .iter()
85 .filter(|m| matches!(m.role.as_str(), "system" | "user" | "tool"))
86 .filter_map(text_of)
87 .collect::<Vec<_>>()
88 .join("\n")
89}
90
91fn text_of(msg: &Message) -> Option<String> {
94 match &msg.content {
95 Value::String(s) if !s.is_empty() => Some(s.clone()),
96 Value::Array(parts) => {
97 let joined = parts
98 .iter()
99 .filter_map(|p| p.get("text").and_then(Value::as_str))
100 .collect::<Vec<_>>()
101 .join("\n");
102 (!joined.is_empty()).then_some(joined)
103 }
104 _ => None,
105 }
106}
107
108fn outcome_label(result: &ScanResult, mode: PolicyMode) -> &'static str {
109 if !result.flagged() {
110 "pass"
111 } else if result.should_refuse() {
112 match mode {
113 PolicyMode::Block => "block",
114 PolicyMode::Observe => "observe",
115 }
116 } else {
117 "flag"
118 }
119}
120
121fn blocking_scanners(result: &ScanResult) -> Vec<String> {
123 let mut names: Vec<String> = result
124 .matches
125 .iter()
126 .filter(|m| m.severity == Severity::Block)
127 .map(|m| m.scanner.to_string())
128 .collect();
129 names.sort();
130 names.dedup();
131 names
132}
133
134fn severity_label(s: Severity) -> &'static str {
135 match s {
136 Severity::Info => "info",
137 Severity::Warn => "warn",
138 Severity::Block => "block",
139 }
140}
141
142fn record_metrics(policy: &str, result: &ScanResult, outcome: &'static str, started: Instant) {
143 metrics::counter!(
144 "synapse_guard_scans_total",
145 "policy" => policy.to_string(),
146 "outcome" => outcome,
147 )
148 .increment(1);
149 for m in &result.matches {
150 metrics::counter!(
151 "synapse_guard_matches_total",
152 "policy" => policy.to_string(),
153 "scanner" => m.scanner,
154 "severity" => severity_label(m.severity),
155 )
156 .increment(1);
157 }
158 metrics::histogram!(
159 "synapse_guard_scan_duration_seconds",
160 "policy" => policy.to_string(),
161 )
162 .record(started.elapsed().as_secs_f64());
163}
164
165#[cfg(test)]
166mod tests {
167 use super::*;
168 use crate::guard::policy::GuardrailsConfig;
169
170 fn engine(toml: &str) -> GuardEngine {
171 GuardEngine::from_config(&GuardrailsConfig::from_toml_str(toml).unwrap()).unwrap()
172 }
173
174 fn req(content: &str) -> ChatRequest {
175 serde_json::from_value(serde_json::json!({
176 "model": "m",
177 "messages": [{ "role": "user", "content": content }]
178 }))
179 .unwrap()
180 }
181
182 const BLOCKING: &str = r#"[guardrails.strict]
183 scanners = [{ type = "ban_substrings", substrings = ["forbidden"] }]"#;
184
185 #[test]
186 fn unknown_policy_is_noop() {
187 let e = engine(BLOCKING);
188 assert!(e.guard("nonexistent", &req("forbidden")).is_ok());
189 }
190
191 #[test]
192 fn clean_input_passes() {
193 let e = engine(BLOCKING);
194 assert!(e.guard("strict", &req("hello world")).is_ok());
195 }
196
197 #[test]
198 fn block_mode_refuses_on_block_severity() {
199 let e = engine(BLOCKING);
200 let err = e.guard("strict", &req("this is forbidden")).unwrap_err();
201 match err {
202 GatewayError::ContentBlocked { policy, scanners } => {
203 assert_eq!(policy, "strict");
204 assert_eq!(scanners, vec!["ban_substrings".to_string()]);
205 }
206 other => panic!("expected ContentBlocked, got {other:?}"),
207 }
208 }
209
210 #[test]
211 fn observe_mode_proceeds_despite_block_severity() {
212 let e = engine(
213 r#"[guardrails.canary]
214 mode = "observe"
215 scanners = [{ type = "ban_substrings", substrings = ["forbidden"] }]"#,
216 );
217 assert!(e.guard("canary", &req("this is forbidden")).is_ok());
218 }
219
220 #[test]
221 fn extracts_text_from_array_content_parts() {
222 let r: ChatRequest = serde_json::from_value(serde_json::json!({
223 "model": "m",
224 "messages": [{ "role": "user",
225 "content": [{ "type": "text", "text": "forbidden" },
226 { "type": "image_url", "image_url": { "url": "x" } }] }]
227 }))
228 .unwrap();
229 let e = engine(BLOCKING);
230 assert!(matches!(
231 e.guard("strict", &r).unwrap_err(),
232 GatewayError::ContentBlocked { .. }
233 ));
234 }
235}