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