Files
ai-agent-admin/backend-fastapi/zq_smart_table/formula.py
T

613 lines
19 KiB
Python

"""
Safe formula engine for smart table.
Supports field references {FieldName}, arithmetic, comparisons, and built-in functions.
No eval/exec - all evaluation done via AST traversal.
"""
import math
from datetime import date, datetime, timedelta
from enum import Enum
from typing import Any, Dict, List, Optional
# ==================== Tokenizer ====================
class TokenType(Enum):
NUMBER = "NUMBER"
STRING = "STRING"
FIELD_REF = "FIELD_REF"
FUNCTION = "FUNCTION"
OPERATOR = "OPERATOR"
LPAREN = "LPAREN"
RPAREN = "RPAREN"
COMMA = "COMMA"
BOOLEAN = "BOOLEAN"
EOF = "EOF"
class Token:
__slots__ = ("type", "value")
def __init__(self, type_: TokenType, value: Any):
self.type = type_
self.value = value
def __repr__(self):
return f"Token({self.type.name}, {self.value!r})"
_FUNC_NAMES = {
"IF", "AND", "OR", "NOT",
"CONCATENATE", "CONCAT",
"ABS", "ROUND", "CEIL", "FLOOR", "INT", "MOD", "POWER", "SQRT",
"UPPER", "LOWER", "LEN", "LEFT", "RIGHT", "MID", "TRIM", "SUBSTITUTE",
"NOW", "TODAY", "DATEDIFF", "DATEADD", "YEAR", "MONTH", "DAY",
"MIN", "MAX", "SUM", "AVERAGE",
"ISNULL", "VALUE", "TEXT", "FIXED",
}
_TWO_CHAR_OPS = {"!=", ">=", "<=", "<>", "&&", "||"}
_ONE_CHAR_OPS = {"+", "-", "*", "/", "%", "=", ">", "<", "&"}
def tokenize(formula: str) -> List[Token]:
tokens: List[Token] = []
i = 0
n = len(formula)
while i < n:
ch = formula[i]
if ch in (" ", "\t", "\n", "\r"):
i += 1
continue
if ch == "{":
end = formula.find("}", i + 1)
if end == -1:
raise FormulaError(f"未闭合的字段引用 '{{' 在位置 {i}")
tokens.append(Token(TokenType.FIELD_REF, formula[i + 1:end]))
i = end + 1
continue
if ch == '"' or ch == "'":
quote = ch
j = i + 1
parts = []
while j < n:
if formula[j] == "\\" and j + 1 < n:
parts.append(formula[j + 1])
j += 2
elif formula[j] == quote:
break
else:
parts.append(formula[j])
j += 1
if j >= n:
raise FormulaError(f"未闭合的字符串在位置 {i}")
tokens.append(Token(TokenType.STRING, "".join(parts)))
i = j + 1
continue
if ch.isdigit() or (ch == "." and i + 1 < n and formula[i + 1].isdigit()):
j = i
has_dot = False
while j < n and (formula[j].isdigit() or (formula[j] == "." and not has_dot)):
if formula[j] == ".":
has_dot = True
j += 1
tokens.append(Token(TokenType.NUMBER, float(formula[i:j])))
i = j
continue
if ch == "(":
tokens.append(Token(TokenType.LPAREN, "("))
i += 1
continue
if ch == ")":
tokens.append(Token(TokenType.RPAREN, ")"))
i += 1
continue
if ch == ",":
tokens.append(Token(TokenType.COMMA, ","))
i += 1
continue
two = formula[i:i + 2] if i + 1 < n else ""
if two in _TWO_CHAR_OPS:
tokens.append(Token(TokenType.OPERATOR, two))
i += 2
continue
if ch in _ONE_CHAR_OPS:
tokens.append(Token(TokenType.OPERATOR, ch))
i += 1
continue
if ch.isalpha() or ch == "_":
j = i
while j < n and (formula[j].isalnum() or formula[j] == "_"):
j += 1
word = formula[i:j]
upper = word.upper()
if upper in ("TRUE", "FALSE"):
tokens.append(Token(TokenType.BOOLEAN, upper == "TRUE"))
elif upper in _FUNC_NAMES:
tokens.append(Token(TokenType.FUNCTION, upper))
else:
tokens.append(Token(TokenType.FIELD_REF, word))
i = j
continue
raise FormulaError(f"无法识别的字符 '{ch}' 在位置 {i}")
tokens.append(Token(TokenType.EOF, None))
return tokens
# ==================== AST Nodes ====================
class ASTNode:
pass
class NumberLiteral(ASTNode):
__slots__ = ("value",)
def __init__(self, value: float):
self.value = value
class StringLiteral(ASTNode):
__slots__ = ("value",)
def __init__(self, value: str):
self.value = value
class BooleanLiteral(ASTNode):
__slots__ = ("value",)
def __init__(self, value: bool):
self.value = value
class FieldReference(ASTNode):
__slots__ = ("name",)
def __init__(self, name: str):
self.name = name
class BinaryOp(ASTNode):
__slots__ = ("op", "left", "right")
def __init__(self, op: str, left: ASTNode, right: ASTNode):
self.op = op
self.left = left
self.right = right
class UnaryOp(ASTNode):
__slots__ = ("op", "operand")
def __init__(self, op: str, operand: ASTNode):
self.op = op
self.operand = operand
class FunctionCall(ASTNode):
__slots__ = ("name", "args")
def __init__(self, name: str, args: List[ASTNode]):
self.name = name
self.args = args
# ==================== Parser ====================
class FormulaError(Exception):
pass
class Parser:
def __init__(self, tokens: List[Token]):
self.tokens = tokens
self.pos = 0
def _current(self) -> Token:
return self.tokens[self.pos]
def _eat(self, expected_type: Optional[TokenType] = None) -> Token:
tok = self._current()
if expected_type and tok.type != expected_type:
raise FormulaError(f"期望 {expected_type.name},实际 {tok.type.name}({tok.value!r})")
self.pos += 1
return tok
def parse(self) -> ASTNode:
node = self._expr()
if self._current().type != TokenType.EOF:
raise FormulaError(f"意外的 token: {self._current()}")
return node
def _expr(self) -> ASTNode:
return self._logic_or()
def _logic_or(self) -> ASTNode:
node = self._logic_and()
while self._current().type == TokenType.OPERATOR and self._current().value in ("||",):
op = self._eat().value
right = self._logic_and()
node = BinaryOp(op, node, right)
return node
def _logic_and(self) -> ASTNode:
node = self._comparison()
while self._current().type == TokenType.OPERATOR and self._current().value in ("&&",):
op = self._eat().value
right = self._comparison()
node = BinaryOp(op, node, right)
return node
def _comparison(self) -> ASTNode:
node = self._concat()
while self._current().type == TokenType.OPERATOR and self._current().value in ("=", "!=", "<>", ">", "<", ">=", "<="):
op = self._eat().value
right = self._concat()
node = BinaryOp(op, node, right)
return node
def _concat(self) -> ASTNode:
node = self._addition()
while self._current().type == TokenType.OPERATOR and self._current().value == "&":
self._eat()
right = self._addition()
node = BinaryOp("&", node, right)
return node
def _addition(self) -> ASTNode:
node = self._multiplication()
while self._current().type == TokenType.OPERATOR and self._current().value in ("+", "-"):
op = self._eat().value
right = self._multiplication()
node = BinaryOp(op, node, right)
return node
def _multiplication(self) -> ASTNode:
node = self._unary()
while self._current().type == TokenType.OPERATOR and self._current().value in ("*", "/", "%"):
op = self._eat().value
right = self._unary()
node = BinaryOp(op, node, right)
return node
def _unary(self) -> ASTNode:
if self._current().type == TokenType.OPERATOR and self._current().value == "-":
self._eat()
operand = self._unary()
return UnaryOp("-", operand)
return self._primary()
def _primary(self) -> ASTNode:
tok = self._current()
if tok.type == TokenType.NUMBER:
self._eat()
return NumberLiteral(tok.value)
if tok.type == TokenType.STRING:
self._eat()
return StringLiteral(tok.value)
if tok.type == TokenType.BOOLEAN:
self._eat()
return BooleanLiteral(tok.value)
if tok.type == TokenType.FIELD_REF:
self._eat()
return FieldReference(tok.value)
if tok.type == TokenType.FUNCTION:
return self._function_call()
if tok.type == TokenType.LPAREN:
self._eat()
node = self._expr()
self._eat(TokenType.RPAREN)
return node
raise FormulaError(f"意外的 token: {tok}")
def _function_call(self) -> ASTNode:
name = self._eat(TokenType.FUNCTION).value
self._eat(TokenType.LPAREN)
args: List[ASTNode] = []
if self._current().type != TokenType.RPAREN:
args.append(self._expr())
while self._current().type == TokenType.COMMA:
self._eat()
args.append(self._expr())
self._eat(TokenType.RPAREN)
return FunctionCall(name, args)
# ==================== Evaluator ====================
def _to_number(v: Any) -> float:
if v is None or v == "":
return 0.0
try:
return float(v)
except (ValueError, TypeError):
return 0.0
def _to_string(v: Any) -> str:
if v is None:
return ""
if isinstance(v, bool):
return "TRUE" if v else "FALSE"
if isinstance(v, float) and v == int(v):
return str(int(v))
return str(v)
def _to_bool(v: Any) -> bool:
if isinstance(v, bool):
return v
if isinstance(v, (int, float)):
return v != 0
if isinstance(v, str):
return v.upper() not in ("", "FALSE", "0")
return bool(v)
def _parse_date(v: Any) -> Optional[datetime]:
if isinstance(v, datetime):
return v
if isinstance(v, date):
return datetime.combine(v, datetime.min.time())
if isinstance(v, str):
for fmt in ("%Y-%m-%d %H:%M:%S", "%Y-%m-%dT%H:%M:%S", "%Y-%m-%d", "%Y/%m/%d"):
try:
return datetime.strptime(v.strip()[:19], fmt)
except ValueError:
continue
return None
def _evaluate_function(name: str, args: List[Any]) -> Any:
n = len(args)
if name == "IF":
if n < 2:
raise FormulaError("IF 需要至少 2 个参数")
cond = _to_bool(args[0])
return args[1] if cond else (args[2] if n > 2 else "")
if name == "AND":
return all(_to_bool(a) for a in args)
if name == "OR":
return any(_to_bool(a) for a in args)
if name == "NOT":
return not _to_bool(args[0]) if n > 0 else True
if name in ("CONCATENATE", "CONCAT"):
return "".join(_to_string(a) for a in args)
if name == "ABS":
return abs(_to_number(args[0])) if n > 0 else 0
if name == "ROUND":
digits = int(_to_number(args[1])) if n > 1 else 0
return round(_to_number(args[0]), digits) if n > 0 else 0
if name == "CEIL":
return math.ceil(_to_number(args[0])) if n > 0 else 0
if name == "FLOOR":
return math.floor(_to_number(args[0])) if n > 0 else 0
if name == "INT":
return int(_to_number(args[0])) if n > 0 else 0
if name == "MOD":
if n < 2:
return 0
divisor = _to_number(args[1])
return _to_number(args[0]) % divisor if divisor != 0 else 0
if name == "POWER":
return _to_number(args[0]) ** _to_number(args[1]) if n >= 2 else 0
if name == "SQRT":
val = _to_number(args[0]) if n > 0 else 0
return math.sqrt(val) if val >= 0 else None
if name == "UPPER":
return _to_string(args[0]).upper() if n > 0 else ""
if name == "LOWER":
return _to_string(args[0]).lower() if n > 0 else ""
if name == "LEN":
return len(_to_string(args[0])) if n > 0 else 0
if name == "LEFT":
s = _to_string(args[0]) if n > 0 else ""
count = int(_to_number(args[1])) if n > 1 else 1
return s[:count]
if name == "RIGHT":
s = _to_string(args[0]) if n > 0 else ""
count = int(_to_number(args[1])) if n > 1 else 1
return s[-count:] if count > 0 else ""
if name == "MID":
s = _to_string(args[0]) if n > 0 else ""
start = max(1, int(_to_number(args[1]))) if n > 1 else 1
length = int(_to_number(args[2])) if n > 2 else 1
return s[start - 1:start - 1 + length]
if name == "TRIM":
return _to_string(args[0]).strip() if n > 0 else ""
if name == "SUBSTITUTE":
if n < 3:
return _to_string(args[0]) if n > 0 else ""
s = _to_string(args[0])
old = _to_string(args[1])
new = _to_string(args[2])
return s.replace(old, new)
if name == "NOW":
return datetime.now().strftime("%Y-%m-%d %H:%M:%S")
if name == "TODAY":
return date.today().isoformat()
if name == "YEAR":
d = _parse_date(args[0]) if n > 0 else None
return d.year if d else None
if name == "MONTH":
d = _parse_date(args[0]) if n > 0 else None
return d.month if d else None
if name == "DAY":
d = _parse_date(args[0]) if n > 0 else None
return d.day if d else None
if name == "DATEDIFF":
if n < 2:
return None
d1 = _parse_date(args[0])
d2 = _parse_date(args[1])
if d1 and d2:
unit = _to_string(args[2]).upper() if n > 2 else "DAYS"
diff = d2 - d1
if unit in ("DAYS", "D"):
return diff.days
if unit in ("HOURS", "H"):
return diff.total_seconds() / 3600
if unit in ("MONTHS", "M"):
return (d2.year - d1.year) * 12 + (d2.month - d1.month)
if unit in ("YEARS", "Y"):
return d2.year - d1.year
return diff.days
return None
if name == "DATEADD":
if n < 2:
return None
d = _parse_date(args[0])
amount = int(_to_number(args[1]))
unit = _to_string(args[2]).upper() if n > 2 else "DAYS"
if d:
if unit in ("DAYS", "D"):
return (d + timedelta(days=amount)).strftime("%Y-%m-%d")
if unit in ("HOURS", "H"):
return (d + timedelta(hours=amount)).strftime("%Y-%m-%d %H:%M:%S")
if unit in ("MONTHS", "M"):
month = d.month + amount
year = d.year + (month - 1) // 12
month = (month - 1) % 12 + 1
day = min(d.day, 28)
return date(year, month, day).isoformat()
return None
if name == "MIN":
nums = [_to_number(a) for a in args if a is not None and a != ""]
return min(nums) if nums else None
if name == "MAX":
nums = [_to_number(a) for a in args if a is not None and a != ""]
return max(nums) if nums else None
if name == "SUM":
return sum(_to_number(a) for a in args if a is not None and a != "")
if name == "AVERAGE":
nums = [_to_number(a) for a in args if a is not None and a != ""]
return sum(nums) / len(nums) if nums else None
if name == "ISNULL":
return args[0] is None or args[0] == "" if n > 0 else True
if name == "VALUE":
return _to_number(args[0]) if n > 0 else 0
if name == "TEXT":
return _to_string(args[0]) if n > 0 else ""
if name == "FIXED":
val = _to_number(args[0]) if n > 0 else 0
digits = int(_to_number(args[1])) if n > 1 else 2
return f"{val:.{digits}f}"
raise FormulaError(f"未知函数: {name}")
def evaluate(node: ASTNode, context: Dict[str, Any], field_name_map: Dict[str, str]) -> Any:
if isinstance(node, NumberLiteral):
return node.value
if isinstance(node, StringLiteral):
return node.value
if isinstance(node, BooleanLiteral):
return node.value
if isinstance(node, FieldReference):
field_id = field_name_map.get(node.name)
if field_id is None:
field_id = node.name
return context.get(field_id)
if isinstance(node, UnaryOp):
val = evaluate(node.operand, context, field_name_map)
if node.op == "-":
return -_to_number(val)
return val
if isinstance(node, BinaryOp):
left = evaluate(node.left, context, field_name_map)
right = evaluate(node.right, context, field_name_map)
op = node.op
if op == "+":
return _to_number(left) + _to_number(right)
if op == "-":
return _to_number(left) - _to_number(right)
if op == "*":
return _to_number(left) * _to_number(right)
if op == "/":
r = _to_number(right)
return _to_number(left) / r if r != 0 else None
if op == "%":
r = _to_number(right)
return _to_number(left) % r if r != 0 else None
if op == "&":
return _to_string(left) + _to_string(right)
if op in ("=", "=="):
return left == right
if op in ("!=", "<>"):
return left != right
if op == ">":
return _to_number(left) > _to_number(right)
if op == "<":
return _to_number(left) < _to_number(right)
if op == ">=":
return _to_number(left) >= _to_number(right)
if op == "<=":
return _to_number(left) <= _to_number(right)
if op == "&&":
return _to_bool(left) and _to_bool(right)
if op == "||":
return _to_bool(left) or _to_bool(right)
raise FormulaError(f"未知运算符: {op}")
if isinstance(node, FunctionCall):
evaluated_args = [evaluate(a, context, field_name_map) for a in node.args]
return _evaluate_function(node.name, evaluated_args)
raise FormulaError(f"未知 AST 节点: {type(node)}")
# ==================== Public API ====================
def compute_formula(
formula: str,
record_values: Dict[str, Any],
field_name_map: Dict[str, str],
) -> Any:
"""
计算公式。
formula: 公式字符串,如 "IF({状态}=\"完成\", {金额} * 1.1, {金额})"
record_values: {fieldId: value}
field_name_map: {fieldName: fieldId}
返回计算结果;出错时返回 '#ERROR'
"""
if not formula or not formula.strip():
return ""
try:
tokens = tokenize(formula)
parser = Parser(tokens)
ast = parser.parse()
result = evaluate(ast, record_values, field_name_map)
if isinstance(result, float):
if result == int(result) and abs(result) < 1e15:
return int(result)
return round(result, 10)
return result
except Exception:
return "#ERROR"