Skip to main content

synapse/guard/
engine.rs

1//! The guard engine: compiles one `llm-guard` pipeline per policy and
2//! runs a request's input text through the route's selected policy.
3
4use 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/// Holds compiled pipelines keyed by policy name. Construct once at startup
24/// and share behind an `Arc`. An empty engine is a no-op (guard returns `Ok`).
25#[derive(Default)]
26pub struct GuardEngine {
27    policies: HashMap<String, CompiledPolicy>,
28    metrics: Arc<GatewayMetrics>,
29}
30
31impl GuardEngine {
32    /// An engine with no policies — every `guard()` call is a no-op.
33    pub fn empty() -> Self {
34        Self::default()
35    }
36
37    /// Compile every policy in `cfg` into a pipeline. Errors on unknown
38    /// scanner names or bad params (fail-fast at startup).
39    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    /// Record scans and matches on `metrics`.
63    pub fn with_metrics(self, metrics: Arc<GatewayMetrics>) -> Self {
64        Self { metrics, ..self }
65    }
66
67    /// Scan `req`'s input under `policy_name`. Returns `Ok(())` when the
68    /// policy is unknown (no-op), clean, in observe mode, or only flags
69    /// non-block matches; returns `ContentBlocked` when a block-severity
70    /// match fires under a `Block`-mode policy.
71    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
106/// Concatenate the text of system/user/tool messages (assistant turns and
107/// non-text content parts are skipped) into one buffer to scan.
108fn 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
117/// Extract scannable text from one message: a string content, or the joined
118/// `text` fields of an array-of-parts content. `None` when there's no text.
119fn 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
147/// Unique scanner names of block-severity matches (for the error body).
148fn 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}