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::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/// Holds compiled pipelines keyed by policy name. Construct once at startup
22/// and share behind an `Arc`. An empty engine is a no-op (guard returns `Ok`).
23#[derive(Default)]
24pub struct GuardEngine {
25    policies: HashMap<String, CompiledPolicy>,
26}
27
28impl GuardEngine {
29    /// An engine with no policies — every `guard()` call is a no-op.
30    pub fn empty() -> Self {
31        Self::default()
32    }
33
34    /// Compile every policy in `cfg` into a pipeline. Errors on unknown
35    /// scanner names or bad params (fail-fast at startup).
36    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    /// Scan `req`'s input under `policy_name`. Returns `Ok(())` when the
57    /// policy is unknown (no-op), clean, in observe mode, or only flags
58    /// non-block matches; returns `ContentBlocked` when a block-severity
59    /// match fires under a `Block`-mode policy.
60    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
80/// Concatenate the text of system/user/tool messages (assistant turns and
81/// non-text content parts are skipped) into one buffer to scan.
82fn 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
91/// Extract scannable text from one message: a string content, or the joined
92/// `text` fields of an array-of-parts content. `None` when there's no text.
93fn 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
121/// Unique scanner names of block-severity matches (for the error body).
122fn 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}