1use std::collections::HashSet;
36
37use crate::error::{KnowReason, KnowledgeResult};
38use orion_error::conversion::ToStructError;
39
40fn is_name_start(c: char) -> bool {
41 c.is_ascii_alphabetic() || c == '_'
42}
43
44fn is_name_char(c: char) -> bool {
45 c.is_ascii_alphanumeric() || c == '_'
46}
47
48#[derive(Debug, Clone, PartialEq)]
50pub enum VarDef {
51 Literal { name: String, value: String },
53 PhaseBucket {
55 name: String,
56 period_s: u64,
57 bucket_s: u64,
58 offset_slots: i64,
59 prefix: String,
60 },
61}
62
63pub fn parse(code: &str) -> KnowledgeResult<Vec<VarDef>> {
65 let mut out = Vec::new();
66 let mut seen: HashSet<String> = HashSet::new();
67 for (idx, raw) in code.lines().enumerate() {
68 let line = raw.trim();
69 if line.is_empty() || line.starts_with('#') {
70 continue;
71 }
72 let Some((lhs, rhs)) = line.split_once('=') else {
73 return err_at(idx, format!("缺 '=' 的赋值行: {line}"));
74 };
75 let name = lhs.trim();
76 let valid = name.starts_with('$')
77 && name.len() > 1
78 && is_name_start(name[1..].chars().next().unwrap_or('_'))
79 && name[1..].chars().all(is_name_char);
80 if !valid {
81 return err_at(
82 idx,
83 format!("左侧应为 $name([A-Za-z_][A-Za-z0-9_]*),实际 {name}"),
84 );
85 }
86 let key = name[1..].to_string();
87 if !seen.insert(key.clone()) {
88 return err_at(idx, format!("变量重复定义: {name}"));
89 }
90 let expr = rhs.trim();
91 let var = parse_expr(&key, expr).map_err(|e| vel_err(format!("第{}行: {}", idx + 1, e)))?;
92 out.push(var);
93 }
94 Ok(out)
95}
96
97fn parse_expr(name: &str, expr: &str) -> Result<VarDef, String> {
98 let expr = expr.trim();
99 if let Some(after_open) = expr.strip_prefix('"') {
101 let Some(q) = after_open.find('"') else {
102 return Err(format!("字符串字面量未闭合: {expr}"));
103 };
104 validate_tail(&after_open[q + 1..], expr)?;
105 return Ok(VarDef::Literal {
106 name: name.to_string(),
107 value: after_open[..q].to_string(),
108 });
109 }
110 let Some(open) = expr.find('(') else {
112 return Err(format!(
113 "不支持的表达式: {expr}(支持 字符串字面量 / VEL 内建函数)"
114 ));
115 };
116 let Some(close) = expr.rfind(')') else {
117 return Err(format!("函数调用缺右括号: {expr}"));
118 };
119 if close < open {
120 return Err(format!("函数调用缺右括号: {expr}"));
121 }
122 validate_tail(&expr[close + 1..], expr)?;
123 let core = &expr[..=close];
124 let fname = core[..open].trim();
125 let args_raw = core[open + 1..close].trim();
126 let args: Vec<&str> = if args_raw.is_empty() {
127 Vec::new()
128 } else {
129 args_raw.split(',').map(|a| a.trim()).collect()
130 };
131 let num = |a: &str| -> Result<u64, String> {
132 a.parse::<u64>()
133 .map_err(|_| format!("{fname}() 参数应为正整数秒,实际 {a:?}"))
134 };
135 match (fname, args.len()) {
136 ("phase_now", 2..=3) => Ok(VarDef::PhaseBucket {
137 name: name.to_string(),
138 period_s: num(args[0])?,
139 bucket_s: num(args[1])?,
140 offset_slots: 0,
141 prefix: prefix_arg(args.get(2).copied())?,
142 }),
143 ("phase_next", 2..=3) => Ok(VarDef::PhaseBucket {
144 name: name.to_string(),
145 period_s: num(args[0])?,
146 bucket_s: num(args[1])?,
147 offset_slots: 1,
148 prefix: prefix_arg(args.get(2).copied())?,
149 }),
150 ("phase_now" | "phase_next", n) => Err(format!(
151 "{fname}() 需 2~3 参数 (period_s, bucket_s[, prefix]),实际 {n}"
152 )),
153 _ => Err(format!("未知 VEL 函数: {fname}")),
154 }
155}
156
157fn validate_tail(tail: &str, whole: &str) -> Result<(), String> {
159 let t = tail.trim();
160 if t.is_empty() || t.starts_with('#') {
161 Ok(())
162 } else {
163 Err(format!("表达式后只允许 # 注释,实际尾缀 {t:?}({whole})"))
164 }
165}
166
167fn prefix_arg(arg: Option<&str>) -> Result<String, String> {
168 match arg {
169 None => Ok("p".to_string()),
170 Some(a) if a.starts_with('"') && a.ends_with('"') && a.len() >= 2 => {
171 Ok(a[1..a.len() - 1].to_string())
172 }
173 Some(a) => Err(format!("prefix 应为字符串字面量,实际 {a:?}")),
174 }
175}
176
177fn err_at(idx: usize, msg: String) -> KnowledgeResult<Vec<VarDef>> {
178 Err(vel_err(format!("第{}行: {}", idx + 1, msg)))
179}
180
181fn vel_err(msg: String) -> crate::error::KnowledgeError {
182 KnowReason::from_res()
183 .to_err()
184 .with_detail(format!("VEL: {msg}"))
185}
186
187pub fn eval(code: &str, now_ns: u64) -> KnowledgeResult<Vec<(String, String)>> {
189 let defs = parse(code)?;
190 let mut out = Vec::with_capacity(defs.len());
191 for d in &defs {
192 match d {
193 VarDef::Literal { name, value } => out.push((name.clone(), value.clone())),
194 VarDef::PhaseBucket {
195 name,
196 period_s,
197 bucket_s,
198 offset_slots,
199 prefix,
200 } => {
201 let period_ns = period_s.saturating_mul(1_000_000_000);
202 let bucket_ns = bucket_s.saturating_mul(1_000_000_000);
203 if period_ns == 0 || bucket_ns == 0 || bucket_ns > period_ns {
204 return Err(vel_err(format!(
205 "${name}: 相位参数非法(须 0<桶≤周期): period={period_s} bucket={bucket_s}"
206 )));
207 }
208 let t = if *offset_slots >= 0 {
209 now_ns.saturating_add((*offset_slots as u64).saturating_mul(bucket_ns))
210 } else {
211 now_ns.saturating_sub((-(*offset_slots) as u64).saturating_mul(bucket_ns))
212 };
213 let idx = (t % period_ns) / bucket_ns;
214 out.push((name.clone(), format!("{prefix}{idx}")));
215 }
216 }
217 }
218 Ok(out)
219}
220
221pub fn resolve_vars(sql: &str, vars: &[(String, String)]) -> String {
227 if vars.is_empty() {
228 return sql.to_string();
229 }
230 let map: std::collections::HashMap<&str, &str> =
231 vars.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect();
232 let mut out = String::with_capacity(sql.len());
233 let mut rest = sql;
234 while let Some(pos) = rest.find('$') {
235 out.push_str(&rest[..pos]);
236 rest = &rest[pos + 1..]; let Some(first) = rest.chars().next() else {
238 out.push('$');
239 break;
240 };
241 if !is_name_start(first) {
242 out.push('$'); continue;
244 }
245 let mut end = 0usize;
247 for (idx, c) in rest.char_indices() {
248 if !is_name_char(c) {
249 break;
250 }
251 end = idx + c.len_utf8();
252 }
253 let ident = &rest[..end];
254 match map.get(ident) {
255 Some(value) => {
256 out.push_str(value);
257 rest = &rest[end..];
258 }
259 None => {
260 out.push('$');
261 out.push_str(ident);
262 rest = &rest[end..];
263 }
264 }
265 }
266 out.push_str(rest);
267 out
268}
269
270pub fn render(sql: &str, code: &str, now_ns: u64) -> KnowledgeResult<String> {
272 if code.trim().is_empty() {
273 return Ok(sql.to_string());
274 }
275 let vars = eval(code, now_ns)?;
276 Ok(resolve_vars(sql, &vars))
277}
278
279pub fn current_wall_nanos() -> u64 {
281 std::time::SystemTime::now()
282 .duration_since(std::time::UNIX_EPOCH)
283 .map(|d| d.as_nanos() as u64)
284 .unwrap_or(0)
285}
286
287#[cfg(test)]
288mod tests {
289 use super::*;
290
291 fn ns(s: u64) -> u64 {
292 s.saturating_mul(1_000_000_000)
293 }
294
295 #[test]
296 fn eval_phase_functions_and_literals() {
297 let code = r#"
299# 注释 + 空行应忽略
300$max_age = "2 hours"
301$cur = phase_now(240, 15)
302$next = phase_next(240, 15)
303"#;
304 assert_eq!(
305 eval(code, ns(120)).unwrap(),
306 vec![
307 ("max_age".to_string(), "2 hours".to_string()),
308 ("cur".to_string(), "p8".to_string()),
309 ("next".to_string(), "p9".to_string()),
310 ]
311 );
312 let kv = eval(code, ns(225)).unwrap();
314 let map: std::collections::HashMap<&str, &str> =
315 kv.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect();
316 assert_eq!(map["cur"], "p15");
317 assert_eq!(map["next"], "p0");
318 let kv = eval(code, ns(360)).unwrap();
320 assert_eq!(kv[1], ("cur".to_string(), "p8".to_string()));
321 }
322
323 #[test]
324 fn custom_prefix_and_errors() {
325 let kv = eval("$b = phase_now(240, 15, \"slot\")", ns(120)).unwrap();
327 assert_eq!(kv, vec![("b".to_string(), "slot8".to_string())]);
328 assert!(eval("$x = foo(1)", ns(0)).is_err(), "未知函数应报错");
330 assert!(eval("no_assign", ns(0)).is_err(), "缺 = 应报错");
331 assert!(eval("", ns(0)).unwrap().is_empty());
332 assert!(eval("# 纯注释", ns(0)).unwrap().is_empty());
333 assert!(eval("$x = phase_now(15, 240)", ns(0)).is_err());
335 }
336
337 #[test]
338 fn render_substitutes_and_empty_passes_through() {
339 let code = "$cur = phase_now(240, 15)\n$max_age = \"2 hours\"";
340 let sql = "SELECT * FROM t WHERE phase_bucket = '$cur' AND win_start >= now() - interval '$max_age'";
341 assert_eq!(
342 render(sql, code, ns(120)).unwrap(),
343 "SELECT * FROM t WHERE phase_bucket = 'p8' AND win_start >= now() - interval '2 hours'"
344 );
345 assert_eq!(render(sql, "", ns(120)).unwrap(), sql);
347 assert_eq!(render(sql, " \n# note\n", ns(120)).unwrap(), sql);
348 }
349
350 #[test]
351 fn resolve_replaces_and_keeps_unknown() {
352 let vars = vec![
353 ("cur".to_string(), "p7".to_string()),
354 ("next".to_string(), "p8".to_string()),
355 ];
356 assert_eq!(
357 resolve_vars("WHERE k IN ('$cur','$next') AND z='$ghost'", &vars),
358 "WHERE k IN ('p7','p8') AND z='$ghost'"
359 );
360 assert_eq!(resolve_vars("WHERE 1", &[]), "WHERE 1");
361 }
362
363 #[test]
364 fn resolve_vars_is_identifier_aware_not_substring() {
365 let vars = vec![
367 ("cur".to_string(), "p7".to_string()),
368 ("cur2".to_string(), "x".to_string()),
369 ("c".to_string(), "y".to_string()),
370 ];
371 let sql = "IN ('$cur','$cur2','$cur_x','$c','$c2','pair:$cur$cur') AND '$9' AND '$中文'";
372 assert_eq!(
373 resolve_vars(sql, &vars),
374 "IN ('p7','x','$cur_x','y','$c2','pair:p7p7') AND '$9' AND '$中文'"
375 );
376 assert_eq!(resolve_vars("$ghost", &vars), "$ghost");
378 let vars2 = vec![("cur".to_string(), "p1".to_string())];
379 assert_eq!(resolve_vars("$$cur", &vars2), "$p1");
380 assert_eq!(resolve_vars("$cur$", &vars2), "p1$");
381 assert_eq!(resolve_vars("$中文$cur", &vars2), "$中文p1");
382 }
383
384 #[test]
385 fn parse_rejects_duplicate_and_invalid_names() {
386 assert!(parse("$a = \"1\"\n$a = \"2\"").is_err());
388 assert!(parse("$1x = \"v\"").is_err());
390 assert!(parse("$x-y = \"v\"").is_err());
391 assert!(parse("x = \"v\"").is_err());
392 assert!(parse("$_a = \"1\"\n$a_1 = \"2\"").is_ok());
394 }
395
396 #[test]
397 fn inline_comments_supported_and_trailing_garbage_rejected() {
398 let code = r#"
399$max_age = "30 days" # 保留期
400$cur = phase_now(240, 15) # 当前格
401$hint = "a # not comment" # 引号内 # 是值
402"#;
403 let kv = eval(code, ns(120)).unwrap();
404 let map: std::collections::HashMap<&str, &str> =
405 kv.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect();
406 assert_eq!(map["max_age"], "30 days");
407 assert_eq!(map["cur"], "p8");
408 assert_eq!(map["hint"], "a # not comment");
409 assert!(parse("$a = \"v\" junk").is_err());
411 assert!(parse("$a = phase_now(240, 15) junk").is_err());
412 assert!(parse("$a = phase_now(240, 15))").is_err());
413 }
414
415 #[test]
416 fn single_slot_period_equals_bucket_stays_stable() {
417 let code = "$cur = phase_now(60, 60)";
419 for t in [0u64, 1, 59, 60, 7_000_000_000_000_000] {
420 assert_eq!(
421 eval(code, t).unwrap(),
422 vec![("cur".to_string(), "p0".to_string())]
423 );
424 }
425 }
426
427 #[test]
428 fn eval_is_deterministic_over_now() {
429 let code = "$cur = phase_now(240, 15)";
431 let a = eval(code, ns(120)).unwrap();
432 let b = eval(code, ns(120)).unwrap();
433 assert_eq!(a, b);
434 let c = eval(code, ns(121)).unwrap();
435 assert_eq!(c, vec![("cur".to_string(), "p8".to_string())], "同格内不变");
436 let d = eval(code, ns(135)).unwrap();
437 assert_eq!(d, vec![("cur".to_string(), "p9".to_string())], "跨格推进");
438 }
439}