1use async_trait::async_trait;
5
6use crate::error::{WorkflowError, WorkflowErrorCode, WorkflowResult};
7use crate::guard::GuardEvaluator;
8
9pub 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}