""" 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"