Skip to main content

sz_rust_workflow/guard/
default.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2024-2026 SZ-Rust Team
3//
4use async_trait::async_trait;
5
6use crate::error::{WorkflowError, WorkflowErrorCode, WorkflowResult};
7use crate::guard::GuardEvaluator;
8
9/// 默认守卫求值器,支持纯函数表达式子集。
10///
11/// 支持的语法:
12/// - 字段访问:`$.field` 或 `$.field.subfield`
13/// - 比较:`==`、`!=`、`>`、`<`、`>=`、`<=`
14/// - 逻辑:`and`、`or`、`not`
15/// - 字面量:数字、字符串(单引号)、`true`/`false`/`null`
16///
17/// 不支持函数调用(副作用检测)。
18pub struct DefaultGuardEvaluator {
19    max_expr_length: usize,
20}
21
22impl DefaultGuardEvaluator {
23    pub fn new(max_expr_length: usize) -> Self {
24        Self { max_expr_length }
25    }
26}
27
28impl Default for DefaultGuardEvaluator {
29    fn default() -> Self {
30        Self::new(1024)
31    }
32}
33
34#[async_trait]
35impl GuardEvaluator for DefaultGuardEvaluator {
36    async fn evaluate(&self, expr: &str, context: &serde_json::Value) -> WorkflowResult<bool> {
37        if expr.len() > self.max_expr_length {
38            return Err(WorkflowError::with_field(
39                WorkflowErrorCode::GuardSideEffect,
40                "表达式超长",
41                "expr_len",
42                &expr.len().to_string(),
43            ));
44        }
45        if contains_function_call(expr) {
46            return Err(WorkflowError::with_field(
47                WorkflowErrorCode::GuardSideEffect,
48                "表达式含函数调用",
49                "expr",
50                expr,
51            ));
52        }
53        let result = eval_expr(expr.trim(), context)?;
54        if let serde_json::Value::Bool(b) = result {
55            Ok(b)
56        } else {
57            Err(WorkflowError::with_field(
58                WorkflowErrorCode::GuardTypeError,
59                "求值结果非布尔",
60                "result",
61                &result.to_string(),
62            ))
63        }
64    }
65}
66
67fn contains_function_call(expr: &str) -> bool {
68    let bytes = expr.as_bytes();
69    let mut i = 0;
70    while i < bytes.len() {
71        let c = bytes[i] as char;
72        if c.is_alphabetic() || c == '_' {
73            let start = i;
74            while i < bytes.len() && (bytes[i] as char).is_alphanumeric()
75                || (i < bytes.len() && bytes[i] == b'_')
76            {
77                i += 1;
78            }
79            let word = &expr[start..i];
80            while i < bytes.len() && bytes[i] as char == ' ' {
81                i += 1;
82            }
83            if i < bytes.len() && bytes[i] == b'(' && word != "not" {
84                return true;
85            }
86        } else {
87            i += 1;
88        }
89    }
90    false
91}
92
93fn eval_expr(expr: &str, ctx: &serde_json::Value) -> WorkflowResult<serde_json::Value> {
94    let expr = expr.trim();
95    if let Some(pos) = find_top_level(expr, " or ") {
96        let left = eval_expr(&expr[..pos], ctx)?;
97        let right = eval_expr(&expr[pos + 4..], ctx)?;
98        return Ok(serde_json::Value::Bool(as_bool(&left)? || as_bool(&right)?));
99    }
100    if let Some(pos) = find_top_level(expr, " and ") {
101        let left = eval_expr(&expr[..pos], ctx)?;
102        let right = eval_expr(&expr[pos + 5..], ctx)?;
103        return Ok(serde_json::Value::Bool(as_bool(&left)? && as_bool(&right)?));
104    }
105    if expr.starts_with("not ") || expr.starts_with("not(") {
106        let inner = expr
107            .strip_prefix("not ")
108            .or_else(|| expr.strip_prefix("not(").map(|s| &s[..s.len() - 1]))
109            .unwrap_or(expr);
110        let val = eval_expr(inner, ctx)?;
111        return Ok(serde_json::Value::Bool(!as_bool(&val)?));
112    }
113    for op in &[" == ", " != ", " >= ", " <= ", " > ", " < "] {
114        if let Some(pos) = find_top_level(expr, op) {
115            let left = eval_value(&expr[..pos], ctx)?;
116            let right = eval_value(&expr[pos + op.len()..], ctx)?;
117            return Ok(serde_json::Value::Bool(compare(&left, &right, op.trim())));
118        }
119    }
120    eval_value(expr, ctx)
121}
122
123fn find_top_level(expr: &str, sep: &str) -> Option<usize> {
124    let mut depth = 0i32;
125    let mut in_string = false;
126    let bytes = expr.as_bytes();
127    let sep_bytes = sep.as_bytes();
128    let sep_len = sep_bytes.len();
129    for i in 0..bytes.len().saturating_sub(sep_len) {
130        let c = bytes[i] as char;
131        if c == '\'' {
132            in_string = !in_string;
133        }
134        if !in_string {
135            if c == '(' {
136                depth += 1;
137            } else if c == ')' {
138                depth -= 1;
139            }
140            if depth == 0 && &expr[i..i + sep_len] == sep {
141                return Some(i);
142            }
143        }
144    }
145    None
146}
147
148fn eval_value(expr: &str, ctx: &serde_json::Value) -> WorkflowResult<serde_json::Value> {
149    let expr = expr.trim();
150    if expr.starts_with('\'') && expr.ends_with('\'') && expr.len() >= 2 {
151        return Ok(serde_json::Value::String(
152            expr[1..expr.len() - 1].to_string(),
153        ));
154    }
155    if expr == "true" {
156        return Ok(serde_json::Value::Bool(true));
157    }
158    if expr == "false" {
159        return Ok(serde_json::Value::Bool(false));
160    }
161    if expr == "null" {
162        return Ok(serde_json::Value::Null);
163    }
164    if let Ok(n) = expr.parse::<i64>() {
165        return Ok(serde_json::Value::Number(n.into()));
166    }
167    if let Ok(n) = expr.parse::<f64>() {
168        return Ok(serde_json::Number::from_f64(n)
169            .map(serde_json::Value::Number)
170            .unwrap_or(serde_json::Value::Null));
171    }
172    if let Some(path) = expr.strip_prefix("$.") {
173        return lookup_path(path, ctx);
174    }
175    Err(WorkflowError::with_field(
176        WorkflowErrorCode::GuardEvalFailed,
177        "无法求值表达式",
178        "expr",
179        expr,
180    ))
181}
182
183fn lookup_path(path: &str, ctx: &serde_json::Value) -> WorkflowResult<serde_json::Value> {
184    let mut current = ctx;
185    for part in path.split('.') {
186        if let serde_json::Value::Object(obj) = current {
187            match obj.get(part) {
188                Some(v) => current = v,
189                None => {
190                    return Err(WorkflowError::with_field(
191                        WorkflowErrorCode::GuardEvalFailed,
192                        "引用不存在的字段",
193                        "field",
194                        part,
195                    ))
196                }
197            }
198        } else {
199            return Err(WorkflowError::with_field(
200                WorkflowErrorCode::GuardEvalFailed,
201                "路径访问非对象",
202                "path",
203                path,
204            ));
205        }
206    }
207    Ok(current.clone())
208}
209
210fn as_bool(v: &serde_json::Value) -> WorkflowResult<bool> {
211    match v {
212        serde_json::Value::Bool(b) => Ok(*b),
213        _ => Err(WorkflowError::with_field(
214            WorkflowErrorCode::GuardTypeError,
215            "非布尔值",
216            "value",
217            &v.to_string(),
218        )),
219    }
220}
221
222fn compare(left: &serde_json::Value, right: &serde_json::Value, op: &str) -> bool {
223    match op {
224        "==" => left == right,
225        "!=" => left != right,
226        ">" | "<" | ">=" | "<=" => {
227            let l = left
228                .as_f64()
229                .or_else(|| left.as_str().and_then(|s| s.parse::<f64>().ok()));
230            let r = right
231                .as_f64()
232                .or_else(|| right.as_str().and_then(|s| s.parse::<f64>().ok()));
233            match (l, r) {
234                (Some(l), Some(r)) => match op {
235                    ">" => l > r,
236                    "<" => l < r,
237                    ">=" => l >= r,
238                    "<=" => l <= r,
239                    _ => false,
240                },
241                _ => false,
242            }
243        }
244        _ => false,
245    }
246}
247
248#[cfg(test)]
249mod tests {
250    use super::*;
251
252    fn ctx() -> serde_json::Value {
253        serde_json::json!({"amount": 200, "name": "test", "flag": true, "nested": {"value": 50}})
254    }
255
256    #[tokio::test]
257    async fn field_access_and_compare() {
258        let ev = DefaultGuardEvaluator::default();
259        assert!(ev.evaluate("$.amount > 100", &ctx()).await.unwrap());
260        assert!(!ev.evaluate("$.amount > 300", &ctx()).await.unwrap());
261        assert!(ev.evaluate("$.amount == 200", &ctx()).await.unwrap());
262        assert!(ev.evaluate("$.amount != 100", &ctx()).await.unwrap());
263        assert!(ev.evaluate("$.amount >= 200", &ctx()).await.unwrap());
264        assert!(ev.evaluate("$.amount <= 200", &ctx()).await.unwrap());
265    }
266
267    #[tokio::test]
268    async fn nested_field_access() {
269        let ev = DefaultGuardEvaluator::default();
270        assert!(ev.evaluate("$.nested.value > 40", &ctx()).await.unwrap());
271        assert!(!ev.evaluate("$.nested.value > 60", &ctx()).await.unwrap());
272    }
273
274    #[tokio::test]
275    async fn logical_and_or() {
276        let ev = DefaultGuardEvaluator::default();
277        assert!(ev
278            .evaluate("$.amount > 100 and $.flag == true", &ctx())
279            .await
280            .unwrap());
281        assert!(!ev
282            .evaluate("$.amount > 100 and $.flag == false", &ctx())
283            .await
284            .unwrap());
285        assert!(ev
286            .evaluate("$.amount > 300 or $.flag == true", &ctx())
287            .await
288            .unwrap());
289        assert!(!ev
290            .evaluate("$.amount > 300 or $.flag == false", &ctx())
291            .await
292            .unwrap());
293    }
294
295    #[tokio::test]
296    async fn logical_not() {
297        let ev = DefaultGuardEvaluator::default();
298        assert!(ev.evaluate("not $.flag == false", &ctx()).await.unwrap());
299        assert!(!ev.evaluate("not $.flag == true", &ctx()).await.unwrap());
300    }
301
302    #[tokio::test]
303    async fn string_compare() {
304        let ev = DefaultGuardEvaluator::default();
305        assert!(ev.evaluate("$.name == 'test'", &ctx()).await.unwrap());
306        assert!(!ev.evaluate("$.name == 'other'", &ctx()).await.unwrap());
307    }
308
309    #[tokio::test]
310    async fn missing_field_error() {
311        let ev = DefaultGuardEvaluator::default();
312        let result = ev.evaluate("$.nonexistent > 100", &ctx()).await;
313        assert!(result.is_err());
314        assert_eq!(result.unwrap_err().code, WorkflowErrorCode::GuardEvalFailed);
315    }
316
317    #[tokio::test]
318    async fn function_call_rejected() {
319        let ev = DefaultGuardEvaluator::default();
320        let result = ev.evaluate("eval('1+1') == 2", &ctx()).await;
321        assert!(result.is_err());
322        assert_eq!(result.unwrap_err().code, WorkflowErrorCode::GuardSideEffect);
323    }
324
325    #[tokio::test]
326    async fn expr_too_long() {
327        let ev = DefaultGuardEvaluator::new(10);
328        let result = ev
329            .evaluate("$.amount > 100 and $.flag == true", &ctx())
330            .await;
331        assert!(result.is_err());
332        assert_eq!(result.unwrap_err().code, WorkflowErrorCode::GuardSideEffect);
333    }
334
335    #[tokio::test]
336    async fn non_boolean_result_error() {
337        let ev = DefaultGuardEvaluator::default();
338        let result = ev.evaluate("$.amount", &ctx()).await;
339        assert!(result.is_err());
340        assert_eq!(result.unwrap_err().code, WorkflowErrorCode::GuardTypeError);
341    }
342}