121 lines
4.7 KiB
Python
121 lines
4.7 KiB
Python
# sanguo_factor/expression_match.py
|
|
"""表达式结构相似度(AST 防换皮,2026-10-10 spec §4.2).
|
|
|
|
QuantaAlpha factor_ast 算法思想(最大公共子树+交换律)的方言内实现:
|
|
不移植其 qlib 式解析器,直接用 Python ast——与 factor_guard 白名单
|
|
同方言,注册的新因子表达式必然可解析.判定永远人做:本模块只产提示.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import ast
|
|
|
|
_COMMUTATIVE = (ast.Add, ast.Mult)
|
|
|
|
|
|
def _is_meta(node: ast.AST) -> bool:
|
|
"""元数据子节点:expr_context(Load/Store)与算子(Div/Mult/USub...)——
|
|
前者是 Name 的语境标记,后者的类型已编码进 BinOp/UnaryOp 的 key.
|
|
不跳过则 ? 兜底签名全局撞车+虚增计数(Name 算 2/整树多 1)."""
|
|
return isinstance(node, (ast.expr_context, ast.operator, ast.unaryop))
|
|
|
|
|
|
def _node_size(node: ast.AST) -> int:
|
|
return 1 + sum(_node_size(c) for c in ast.iter_child_nodes(node)
|
|
if not _is_meta(c))
|
|
|
|
|
|
def _key(node: ast.AST) -> str:
|
|
"""规范化结构签名:交换律 binop 左右子树排序后拼接.
|
|
|
|
兜底分支递归编码全部非元数据子节点——同类型未知节点只有子树也
|
|
同构才撞签,不再 `?TypeName` 全局等价(codex review CRITICAL).
|
|
"""
|
|
if isinstance(node, ast.BinOp) and isinstance(node.op, _COMMUTATIVE):
|
|
lk, rk = _key(node.left), _key(node.right)
|
|
a, b = sorted((lk, rk))
|
|
return f"({a}|{b}|{type(node.op).__name__})"
|
|
if isinstance(node, ast.BinOp):
|
|
return f"({_key(node.left)}>{_key(node.right)}|{type(node.op).__name__})"
|
|
if isinstance(node, ast.UnaryOp):
|
|
return f"{type(node.op).__name__}({_key(node.operand)})"
|
|
if isinstance(node, ast.Call) and isinstance(node.func, ast.Name):
|
|
args = ",".join(_key(a) for a in node.args)
|
|
kws = ",".join(f"{k.arg}={_key(k.value)}"
|
|
for k in sorted(node.keywords, key=lambda k: k.arg or ""))
|
|
return f"{node.func.id}({args};{kws})"
|
|
if isinstance(node, ast.Name):
|
|
return f"#{node.id}"
|
|
if isinstance(node, ast.Constant):
|
|
return f"#{node.value!r}"
|
|
kids = ",".join(_key(c) for c in ast.iter_child_nodes(node)
|
|
if not _is_meta(c))
|
|
return f"?{type(node).__name__}<{kids}>"
|
|
|
|
|
|
def _call_func_ids(node: ast.AST) -> set[int]:
|
|
"""Call 的 func Name 节点 id 集:算子名已编码进 Call 签名,不再作为
|
|
独立子树参与匹配——否则仅共享算子名(ts_mean/f/cs_rank)也计公共
|
|
子树,产生 1/N 的地板噪声(codex review 修批实测)."""
|
|
return {id(c.func) for c in ast.walk(node) if isinstance(c, ast.Call)}
|
|
|
|
|
|
def _all_subtree_keys(node: ast.AST) -> dict[str, int]:
|
|
"""子树签名→节点数(同签名取最大)."""
|
|
skip = _call_func_ids(node)
|
|
out: dict[str, int] = {}
|
|
for sub in ast.walk(node):
|
|
if _is_meta(sub) or id(sub) in skip:
|
|
continue
|
|
k = _key(sub)
|
|
out[k] = max(out.get(k, 0), _node_size(sub))
|
|
return out
|
|
|
|
|
|
def similarity(expr_a: str, expr_b: str) -> float | None:
|
|
"""最大公共子树节点数 / 较小表达式节点数;解析失败或深树爆栈=None.
|
|
|
|
取 .body 剥掉 Expression 包装节点——否则其签名落 ? 兜底桶,
|
|
任意两表达式的根都会撞签(子树=全树,相似度恒 1.0).
|
|
RecursionError 防护:factor_guard _DEPTH_MAX=12 挡真实注册深树,
|
|
此处捕爆栈如实 None,不 raise(codex review LOW).
|
|
"""
|
|
try:
|
|
ta = ast.parse(expr_a, mode="eval").body
|
|
tb = ast.parse(expr_b, mode="eval").body
|
|
sa, sb = _node_size(ta), _node_size(tb)
|
|
if sa == 0 or sb == 0:
|
|
return None
|
|
keys_a = _all_subtree_keys(ta)
|
|
skip_b = _call_func_ids(tb)
|
|
best = 0
|
|
for sub in ast.walk(tb):
|
|
if _is_meta(sub) or id(sub) in skip_b:
|
|
continue
|
|
k = _key(sub)
|
|
if k in keys_a:
|
|
best = max(best, min(keys_a[k], _node_size(sub)))
|
|
return round(best / min(sa, sb), 4)
|
|
except (SyntaxError, ValueError, RecursionError):
|
|
return None
|
|
|
|
|
|
def top_similar(expression: str, candidates: dict[str, str],
|
|
floor: float = 0.6, limit: int = 5) -> list[dict]:
|
|
"""对候选池按相似度排序,过滤低于 floor 的(提示非硬拒).
|
|
|
|
limit=5 上限:注册路径 registry 每条只存前 5 提示,防大池刷屏
|
|
(codex review MEDIUM).
|
|
"""
|
|
out: list[dict] = []
|
|
for name, expr in candidates.items():
|
|
if not expression or not expr:
|
|
continue
|
|
try:
|
|
r = similarity(expression, expr)
|
|
except RecursionError:
|
|
continue
|
|
if r is not None and r >= floor:
|
|
out.append({"name": name, "ratio": r})
|
|
out.sort(key=lambda h: h["ratio"], reverse=True)
|
|
return out[:limit]
|