286 lines
12 KiB
Python
286 lines
12 KiB
Python
"""Bounded FASTEXPR syntax analysis and mixed-radix sampling, without execution.
|
||
|
||
This parser establishes syntax and identifier provenance, not full BRAIN semantics.
|
||
Unknown fields/operators must be resolved against snapshots before simulation.
|
||
"""
|
||
|
||
import math
|
||
import random
|
||
import re
|
||
from dataclasses import dataclass
|
||
|
||
PLACEHOLDER = re.compile(r"\{([A-Za-z_][A-Za-z0-9_]*)\}")
|
||
LEGACY_PLACEHOLDER = re.compile(r"<([A-Za-z_][A-Za-z0-9_]*)/>")
|
||
IDENTIFIER = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
|
||
GROUPS = {"sector", "industry", "subindustry", "market", "country", "exchange"}
|
||
CONSTANTS = {"true", "false", "nan", "NaN", "inf"}
|
||
TOKEN = re.compile(
|
||
r"""\s*(?:(\d+(?:\.\d*)?(?:[eE][+-]?\d+)?|\.\d+(?:[eE][+-]?\d+)?)|([A-Za-z_][A-Za-z0-9_]*)|("(?:[^"\\]|\\.)*"|'(?:[^'\\]|\\.)*')|(==|!=|<=|>=|&&|\|\||\*\*|[()+\-*/%^<>=!?:,;]))"""
|
||
)
|
||
PRECEDENCE = {
|
||
"||": 1,
|
||
"&&": 2,
|
||
"==": 3,
|
||
"!=": 3,
|
||
"<": 4,
|
||
">": 4,
|
||
"<=": 4,
|
||
">=": 4,
|
||
"+": 5,
|
||
"-": 5,
|
||
"*": 6,
|
||
"/": 6,
|
||
"%": 6,
|
||
"^": 7,
|
||
"**": 7,
|
||
}
|
||
|
||
|
||
@dataclass
|
||
class ExpressionError(ValueError):
|
||
message: str
|
||
position: int = 0
|
||
|
||
def __str__(self):
|
||
return f"{self.message}(位置 {self.position + 1})"
|
||
|
||
|
||
class Parser:
|
||
def __init__(self, expression):
|
||
if not expression.strip() or len(expression) > 20000:
|
||
raise ExpressionError("表达式为空或超过 20000 字符")
|
||
self.tokens = []
|
||
position = 0
|
||
while position < len(expression.rstrip()):
|
||
match = TOKEN.match(expression, position)
|
||
if not match:
|
||
raise ExpressionError("无法识别的字符", position)
|
||
self.tokens.append((match.lastindex, match.group(match.lastindex), match.start()))
|
||
position = match.end()
|
||
if len(self.tokens) > 5000:
|
||
raise ExpressionError("表达式过于复杂")
|
||
self.tokens.append((0, "EOF", len(expression)))
|
||
self.i = 0
|
||
self.locals = set()
|
||
self.fields = set()
|
||
self.operators = set()
|
||
|
||
def peek(self, offset=0):
|
||
return self.tokens[min(self.i + offset, len(self.tokens) - 1)][1]
|
||
|
||
def take(self, expected=None):
|
||
token = self.tokens[self.i]
|
||
if expected and token[1] != expected:
|
||
raise ExpressionError(f"需要 {expected},实际为 {token[1]}", token[2])
|
||
self.i += 1
|
||
return token
|
||
|
||
def expression(self, minimum=0, depth=0):
|
||
if depth > 64:
|
||
raise ExpressionError("嵌套层数超过 64")
|
||
kind, value, pos = self.take()
|
||
if value in ("+", "-", "!"):
|
||
left = {"kind": "unary", "value": value, "args": [self.expression(7, depth + 1)]}
|
||
elif value == "(":
|
||
left = self.expression(0, depth + 1)
|
||
self.take(")")
|
||
elif kind in (1, 3):
|
||
if kind == 1 and not math.isfinite(float(value)):
|
||
raise ExpressionError("数值必须有限", pos)
|
||
left = {"kind": "number" if kind == 1 else "string", "value": value}
|
||
elif kind == 2:
|
||
if self.peek() == "(":
|
||
self.operators.add(value)
|
||
self.take("(")
|
||
args, keywords = [], set()
|
||
if self.peek() != ")":
|
||
while True:
|
||
keyword = None
|
||
if self.tokens[self.i][0] == 2 and self.peek(1) == "=":
|
||
keyword = self.take()[1]
|
||
self.take("=")
|
||
if keyword in keywords:
|
||
raise ExpressionError("命名参数重复", pos)
|
||
keywords.add(keyword)
|
||
elif keywords:
|
||
raise ExpressionError("位置参数不能出现在命名参数后", pos)
|
||
argument = self.expression(0, depth + 1)
|
||
args.append(
|
||
{"kind": "keyword", "value": keyword, "args": [argument]} if keyword else argument
|
||
)
|
||
if self.peek() != ",":
|
||
break
|
||
self.take(",")
|
||
self.take(")")
|
||
left = {"kind": "call", "value": value, "args": args}
|
||
else:
|
||
if value not in self.locals and value not in CONSTANTS:
|
||
self.fields.add(value)
|
||
left = {"kind": "local" if value in self.locals else "field", "value": value}
|
||
else:
|
||
raise ExpressionError("需要字段、常量或算子调用", pos)
|
||
while self.peek() in PRECEDENCE and PRECEDENCE[self.peek()] >= minimum:
|
||
op = self.take()[1]
|
||
right = self.expression(PRECEDENCE[op] + (0 if op in ("^", "**") else 1), depth + 1)
|
||
left = {"kind": "binary", "value": op, "args": [left, right]}
|
||
if minimum == 0 and self.peek() == "?":
|
||
self.take("?")
|
||
yes = self.expression(0, depth + 1)
|
||
self.take(":")
|
||
left = {"kind": "conditional", "args": [left, yes, self.expression(0, depth + 1)]}
|
||
return left
|
||
|
||
def parse(self):
|
||
statements = []
|
||
final_is_assignment = False
|
||
while self.peek() != "EOF":
|
||
name = None
|
||
if self.tokens[self.i][0] == 2 and self.peek(1) == "=":
|
||
name = self.take()[1]
|
||
self.take("=")
|
||
node = self.expression()
|
||
if name:
|
||
self.locals.add(name)
|
||
node = {"kind": "assignment", "value": name, "args": [node]}
|
||
final_is_assignment = name is not None
|
||
statements.append(node)
|
||
if self.peek() != "EOF":
|
||
self.take(";")
|
||
if final_is_assignment:
|
||
raise ExpressionError("最后一项必须是返回表达式")
|
||
return {
|
||
"ast": statements,
|
||
"fields": sorted(self.fields),
|
||
"operators": sorted(self.operators),
|
||
"locals": sorted(self.locals),
|
||
}
|
||
|
||
|
||
def analyze(expression, fields=None, operators=None):
|
||
"""Return separate syntax, type and availability findings; unknown never means valid."""
|
||
try:
|
||
parsed = Parser(expression).parse()
|
||
except (ExpressionError, RecursionError) as exc:
|
||
return {
|
||
"status": "invalid",
|
||
"syntax": [str(exc)],
|
||
"types": [],
|
||
"availability": [],
|
||
"fields": [],
|
||
"operators": [],
|
||
"locals": [],
|
||
}
|
||
types, availability = [], []
|
||
known = {**{name: "GROUP" for name in GROUPS}, **(fields or {})}
|
||
for field in parsed["fields"]:
|
||
if field not in known and field not in CONSTANTS:
|
||
availability.append(f"字段 {field} 尚未在固定输入中核实")
|
||
elif field in known and known[field] not in ("MATRIX", "VECTOR", "GROUP"):
|
||
availability.append(f"字段 {field} 的类型尚不支持")
|
||
for operator in parsed["operators"]:
|
||
if operators is None or operator not in operators:
|
||
availability.append(f"算子 {operator} 尚未在算子目录中核实")
|
||
local_types = {}
|
||
|
||
def infer(node):
|
||
kind, value = node["kind"], node.get("value")
|
||
if kind == "field":
|
||
if value in CONSTANTS:
|
||
return "SCALAR"
|
||
return known.get(value, "UNKNOWN")
|
||
if kind in ("number", "string"):
|
||
return "SCALAR" if kind == "number" else "STRING"
|
||
if kind == "local":
|
||
return local_types.get(value, "UNKNOWN")
|
||
args = [infer(arg) for arg in node.get("args", [])]
|
||
if kind == "assignment":
|
||
local_types[value] = args[0]
|
||
if kind == "call" and value.startswith("vec_"):
|
||
if not args:
|
||
types.append(f"{value} 缺少 VECTOR 参数")
|
||
if args and args[0] not in ("VECTOR", "UNKNOWN"):
|
||
types.append(f"{value} 的首个参数必须是 VECTOR")
|
||
return "MATRIX"
|
||
if kind == "call" and "VECTOR" in args:
|
||
types.append(f"{value} 使用 VECTOR 前需要显式聚合")
|
||
if kind == "call" and value in {
|
||
"rank",
|
||
"ts_rank",
|
||
"ts_mean",
|
||
"ts_sum",
|
||
"ts_delta",
|
||
"ts_std_dev",
|
||
"zscore",
|
||
"group_rank",
|
||
"group_neutralize",
|
||
}:
|
||
minimum = 2 if value.startswith(("ts_", "group_")) else 1
|
||
if len(args) < minimum:
|
||
types.append(f"{value} 缺少必需参数")
|
||
if args and args[0] == "VECTOR":
|
||
types.append(f"{value} 不能直接使用 VECTOR,请显式选择聚合方法")
|
||
if kind == "call" and value in {"group_rank", "group_neutralize", "group_zscore"}:
|
||
if len(args) > 1 and args[1] not in ("GROUP", "UNKNOWN"):
|
||
types.append(f"{value} 的分组参数必须是 GROUP")
|
||
if kind == "binary" and "VECTOR" in args:
|
||
types.append("VECTOR 参与数值运算前需要显式聚合")
|
||
if "VECTOR" in args:
|
||
return "VECTOR"
|
||
return args[0] if kind in ("unary", "keyword", "assignment") and args else "MATRIX"
|
||
|
||
try:
|
||
result_type = None
|
||
for node in parsed.pop("ast"):
|
||
result_type = infer(node)
|
||
if result_type == "VECTOR":
|
||
types.append("最终 Alpha 输出不能直接是 VECTOR,请显式选择聚合方法")
|
||
except RecursionError:
|
||
types.append("表达式推导过于复杂,请拆分局部变量")
|
||
return {
|
||
**parsed,
|
||
"syntax": [],
|
||
"types": list(dict.fromkeys(types)),
|
||
"availability": availability,
|
||
"status": "invalid" if types else "needs_review" if availability else "valid",
|
||
"limitation": "仅验证支持的语法、字段归属及已知类型约束;平台语义与权限以实际模拟为准",
|
||
}
|
||
|
||
|
||
def normalize_template(expression):
|
||
return LEGACY_PLACEHOLDER.sub(lambda match: "{" + match[1] + "}", expression)
|
||
|
||
|
||
def expand(expression, variables, mode="all", limit=100, seed=0):
|
||
"""Sample integer indices in the Cartesian space without materializing that space."""
|
||
expression = normalize_template(expression)
|
||
names = list(dict.fromkeys(PLACEHOLDER.findall(expression)))
|
||
if set(names) != set(variables) or any(not values for values in variables.values()):
|
||
raise ValueError("占位符必须与非空变量候选逐一对应")
|
||
if "{" in PLACEHOLDER.sub("", expression) or "}" in PLACEHOLDER.sub("", expression):
|
||
raise ValueError("占位符格式应为 {name}")
|
||
total = math.prod(len(variables[name]) for name in names)
|
||
if not 1 <= limit <= 10000:
|
||
raise ValueError("生成上限必须在 1–10000 之间")
|
||
if mode == "all" and total > limit:
|
||
raise ValueError(f"组合数 {total} 超过上限 {limit},请缩小候选或使用随机采样")
|
||
count = min(total, limit)
|
||
if mode == "random":
|
||
# Floyd sampling supports arbitrary-size integers (random.sample(range(N)) does not).
|
||
rng, chosen = random.Random(seed), set()
|
||
for j in range(total - count, total):
|
||
candidate = rng.randrange(j + 1)
|
||
chosen.add(j if candidate in chosen else candidate)
|
||
indices = sorted(chosen)
|
||
else:
|
||
indices = range(count)
|
||
results = []
|
||
for index in indices:
|
||
bindings = {}
|
||
for name in reversed(names):
|
||
values = variables[name]
|
||
index, digit = divmod(index, len(values))
|
||
bindings[name] = values[digit]
|
||
text = PLACEHOLDER.sub(lambda match: str(bindings[match[1]]), expression)
|
||
results.append({"expression": text, "bindings": bindings})
|
||
return {"combination_count": str(total), "seed": seed if mode == "random" else None, "items": results}
|