307 lines
14 KiB
Python
307 lines
14 KiB
Python
# sanguo_factor/factor_guard.py
|
|
"""分解器确定性护栏(C1/D1,spec §4.8 P4-2 折入)——纯函数零 IO.
|
|
|
|
白名单三源:算子(vnpy EXPRESSION_FUNCTIONS 干净子集,42 注册名裁掉
|
|
DataProxy/quesval 等杂项)+列名(三族 adapter 落地列,与 adapter 常量同源
|
|
=零漂移)+常量(数字字面量).全部违规结构化返回(中文 detail)供反馈循环
|
|
渲染(C3).上游范式: qlib OpsWrapper 命名空间封闭,修正为 ast.parse+Name
|
|
预校验(调研报告 §2).
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import ast
|
|
import re
|
|
from dataclasses import dataclass
|
|
|
|
from .fundamental_schema import FEATURE_COLUMNS
|
|
from .sentiment_adapter import SENTIMENT_COLUMNS
|
|
|
|
BAR_COLUMNS = ("open", "high", "low", "close", "volume", "turnover", "vwap")
|
|
COLUMN_DOMAINS: dict[str, str] = {
|
|
**{c: "bars_daily" for c in BAR_COLUMNS},
|
|
**{c: "fundamentals_pit" for c in FEATURE_COLUMNS},
|
|
**{c: "corpus_sentiment" for c in SENTIMENT_COLUMNS},
|
|
}
|
|
WIRED_DOMAINS = frozenset(COLUMN_DOMAINS.values())
|
|
|
|
ALLOWED_OPERATORS = frozenset({
|
|
"ts_delay", "ts_delta", "ts_sum", "ts_mean", "ts_std", "ts_min", "ts_max",
|
|
"ts_argmax", "ts_argmin", "ts_rank", "ts_quantile", "ts_slope",
|
|
"ts_rsquare", "ts_resi", "ts_corr", "ts_cov", "ts_decay_linear",
|
|
"ts_product", "ts_less", "ts_greater", "ts_log", "ts_abs",
|
|
"cs_rank", "cs_mean", "cs_std", "cs_sum", "cs_scale", "cs_neutralize",
|
|
"ta_atr", "ta_rsi", "log", "sign", "abs", "greater", "less",
|
|
})
|
|
|
|
# 算子 arity 表(必需位置参数个数,逐个对 vnpy 实现核实:
|
|
# vnpy alpha/dataset ts/cs/math/ta_function + 本仓 fundamental_neutralize;
|
|
# fast_ops 覆盖同名同签名;white list 无 *args/**kwargs 情形,keywords 已拦)
|
|
_OPERATOR_ARITY: dict[str, int] = {
|
|
"ts_delay": 2, "ts_delta": 2, "ts_sum": 2, "ts_mean": 2, "ts_std": 2,
|
|
"ts_min": 2, "ts_max": 2, "ts_argmax": 2, "ts_argmin": 2, "ts_rank": 2,
|
|
"ts_quantile": 3, "ts_slope": 2, "ts_rsquare": 2, "ts_resi": 2,
|
|
"ts_corr": 3, "ts_cov": 3, "ts_decay_linear": 2, "ts_product": 2,
|
|
"ts_less": 2, "ts_greater": 2, "ts_log": 1, "ts_abs": 1,
|
|
"cs_rank": 1, "cs_mean": 1, "cs_std": 1, "cs_sum": 1, "cs_scale": 1,
|
|
"cs_neutralize": 2,
|
|
"ta_atr": 4, "ta_rsi": 2,
|
|
"log": 1, "sign": 1, "abs": 1, "greater": 2, "less": 2,
|
|
}
|
|
|
|
SOURCE_CATEGORY = {"bars_daily": "technical", "fundamentals_pit": "fundamental",
|
|
"corpus_sentiment": "sentiment"}
|
|
|
|
# P3-13: \Z 不吃尾换行($ 会放过 "abc\n")
|
|
NAME_RE = re.compile(r"^[a-z][a-z0-9_]{2,39}\Z")
|
|
|
|
_ALLOWED_BINOPS = (ast.Add, ast.Sub, ast.Mult, ast.Div)
|
|
# P2-8 补 JoinedStr/FormattedValue/Await/Yield/YieldFrom——f-string/await/yield
|
|
# 曾零违规过门(不构成注入但非法方言混进注册表,eval 阶段才报错)
|
|
_BLOCKED_NODES = (ast.Attribute, ast.Subscript, ast.Compare, ast.BoolOp,
|
|
ast.IfExp, ast.Lambda, ast.ListComp, ast.SetComp, ast.DictComp,
|
|
ast.GeneratorExp, ast.Starred, ast.NamedExpr,
|
|
ast.Tuple, ast.List, ast.Set,
|
|
ast.JoinedStr, ast.FormattedValue, ast.Await,
|
|
ast.Yield, ast.YieldFrom)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class Violation:
|
|
code: str # syntax_error/unknown_operator/unknown_column/…
|
|
detail: str # 中文,含具体名字,直接渲染进反馈 prompt
|
|
|
|
|
|
def _column_hint(col: str, unwired: tuple[str, ...]) -> str:
|
|
msg = f"列 {col} 不在底座可用清单"
|
|
if unwired:
|
|
msg += (f"(卡片关联域 {', '.join(unwired)} 尚未接线表达式底座,"
|
|
"若该列属之需先由 data 域接入)")
|
|
return msg
|
|
|
|
|
|
def precheck_expression(expr: str,
|
|
unwired_domains: tuple[str, ...] = ()) -> list[Violation]:
|
|
"""ast.parse+Name 白名单预校验(C1).语法错→单条 syntax_error,其余逐节点."""
|
|
try:
|
|
tree = ast.parse(expr, mode="eval")
|
|
except (SyntaxError, ValueError) as e:
|
|
return [Violation("syntax_error", f"表达式无法解析: {e}")]
|
|
out: list[Violation] = []
|
|
func_ids = {id(n.func) for n in ast.walk(tree)
|
|
if isinstance(n, ast.Call) and isinstance(n.func, ast.Name)}
|
|
for node in ast.walk(tree):
|
|
if isinstance(node, ast.Name):
|
|
if id(node) not in func_ids and node.id not in COLUMN_DOMAINS:
|
|
out.append(Violation(
|
|
"unknown_column", _column_hint(node.id, unwired_domains)))
|
|
elif isinstance(node, ast.Call):
|
|
fname = getattr(node.func, "id", None)
|
|
if not isinstance(node.func, ast.Name) or fname not in ALLOWED_OPERATORS:
|
|
out.append(Violation("unknown_operator",
|
|
f"算子 {fname or '(复杂表达式)'} 不在白名单"))
|
|
else:
|
|
n_need = _OPERATOR_ARITY[fname]
|
|
if len(node.args) != n_need:
|
|
out.append(Violation(
|
|
"bad_arity",
|
|
f"算子 {fname} 需 {n_need} 个参数,给了 {len(node.args)}"))
|
|
if node.keywords:
|
|
out.append(Violation("syntax_not_allowed", "算子调用不允许关键字参数"))
|
|
elif isinstance(node, ast.BinOp) and not isinstance(node.op, _ALLOWED_BINOPS):
|
|
out.append(Violation(
|
|
"syntax_not_allowed",
|
|
f"二元运算 {type(node.op).__name__} 不允许(只支持 + - * /)"))
|
|
elif isinstance(node, ast.UnaryOp) and not isinstance(node.op, ast.USub):
|
|
out.append(Violation("syntax_not_allowed", "一元运算只允许取负 -"))
|
|
elif isinstance(node, _BLOCKED_NODES):
|
|
out.append(Violation("syntax_not_allowed",
|
|
f"语法 {type(node).__name__} 不允许"))
|
|
elif (isinstance(node, ast.Constant)
|
|
and not (isinstance(node.value, (int, float))
|
|
and not isinstance(node.value, bool))):
|
|
out.append(Violation("syntax_not_allowed", "只允许数字字面量"))
|
|
if not any(isinstance(n, ast.Name) and id(n) not in func_ids
|
|
for n in ast.walk(tree)):
|
|
out.append(Violation("no_data_column", "表达式未引用任何数据列"))
|
|
return list({(v.code, v.detail): v for v in out}.values())
|
|
|
|
|
|
def check_name(name: str, existing_names: set[str]) -> list[Violation]:
|
|
out: list[Violation] = []
|
|
if not NAME_RE.match(name):
|
|
out.append(Violation("bad_name", "因子名须 ^[a-z][a-z0-9_]{2,39}\\Z"))
|
|
if name in existing_names or name in ALLOWED_OPERATORS or name in COLUMN_DOMAINS:
|
|
out.append(Violation("name_taken",
|
|
f"名字 {name} 已被占用(在库因子/算子/列名)"))
|
|
return out
|
|
|
|
|
|
def derive_source(expr: str) -> str | Violation:
|
|
"""由表达式列族确定性推导 params.source(source 非LLM声明=防编数).
|
|
|
|
规则: 列至多 bars+单一特征族(跨 fund+sent 不支持——batch_eval 按单一
|
|
category join 特征,跨族列会全 NaN).
|
|
"""
|
|
tree = ast.parse(expr, mode="eval") # 调用方已过 precheck
|
|
func_ids = {id(n.func) for n in ast.walk(tree)
|
|
if isinstance(n, ast.Call) and isinstance(n.func, ast.Name)}
|
|
domains = {COLUMN_DOMAINS[n.id] for n in ast.walk(tree)
|
|
if isinstance(n, ast.Name) and id(n) not in func_ids}
|
|
feature = domains - {"bars_daily"}
|
|
if len(feature) > 1:
|
|
return Violation("cross_domain",
|
|
f"跨特征族 {sorted(feature)} 不支持(至多 bars+单一特征族)")
|
|
return next(iter(feature)) if feature else "bars_daily"
|
|
|
|
|
|
def category_for_source(source: str) -> str:
|
|
"""batch_eval 特征 join 的 category 映射(动态注册链消费,Task 6)."""
|
|
return SOURCE_CATEGORY.get(source, "technical")
|
|
|
|
|
|
# —— 复杂度五条件+AST 防换皮(QuantaAlpha regulator/factor_ast 适配闭列方言,
|
|
# 调研报告 §4:free_args/unique_vars 比率在全部叶=已知列的方言里退化为
|
|
# 引用列数,故五条件=符号长/引用列/节点数/字面量/嵌套深) ——
|
|
|
|
_SYMBOL_LEN_MAX = 300
|
|
_BASE_FEATURES_MAX = 6
|
|
_NODE_COUNT_MAX = 50
|
|
_NUMERIC_LIT_MAX = 10
|
|
_DEPTH_MAX = 12 # 引用列≤6 的合法和式(cs_rank(ts_mean(a+…+f,20)))结构深度≈9,留嵌套余量
|
|
_DUP_SUBTREE_MIN = 8
|
|
|
|
|
|
def _key(node: ast.AST, swap: bool = True) -> tuple:
|
|
"""结构规范化键(忽略位置;swap=True 时 + * 交换律子节点排序保对称相等).
|
|
|
|
swap=False=严格序键,仅供整树 dup_exact 比对(交换律重排是换皮非逐字
|
|
重发,须落 dup_subtree).
|
|
"""
|
|
if isinstance(node, ast.Constant) and not isinstance(node.value, bool):
|
|
return ("num", node.value)
|
|
if isinstance(node, ast.Name):
|
|
return ("name", node.id)
|
|
if isinstance(node, ast.Call) and isinstance(node.func, ast.Name):
|
|
return ("call", node.func.id,
|
|
tuple(_key(a, swap) for a in node.args))
|
|
if isinstance(node, ast.BinOp):
|
|
op = type(node.op).__name__
|
|
kids = [_key(node.left, swap), _key(node.right, swap)]
|
|
if swap and op in ("Add", "Mult"):
|
|
kids = sorted(kids, key=repr)
|
|
return ("bin", op, tuple(kids))
|
|
if isinstance(node, ast.UnaryOp):
|
|
return ("un", type(node.op).__name__, _key(node.operand, swap))
|
|
return ("other", type(node).__name__)
|
|
|
|
|
|
_CTX_NODES = (ast.Load, ast.Store, ast.Del)
|
|
|
|
|
|
def _size(node: ast.AST) -> int:
|
|
# ctx(Load/Store/Del)非表达式构件,剔掉防节点数虚高(Name 挂 Load 会 +1)
|
|
return sum(1 for n in ast.walk(node)
|
|
if not isinstance(n, _CTX_NODES))
|
|
|
|
|
|
def _depth(node: ast.AST) -> int:
|
|
kids = list(ast.iter_child_nodes(node))
|
|
return 1 + max((_depth(k) for k in kids), default=0)
|
|
|
|
|
|
def _parse_or_none(expr: str) -> ast.Expression | None:
|
|
try:
|
|
return ast.parse(expr, mode="eval")
|
|
except (SyntaxError, ValueError):
|
|
return None
|
|
|
|
|
|
def complexity_violations(expr: str) -> list[Violation]:
|
|
out: list[Violation] = []
|
|
compact = re.sub(r"\s+", "", expr)
|
|
if len(compact) > _SYMBOL_LEN_MAX:
|
|
out.append(Violation("complexity", f"符号长度 {len(compact)}>{_SYMBOL_LEN_MAX}"))
|
|
tree = _parse_or_none(expr)
|
|
if tree is None:
|
|
return out
|
|
func_ids = {id(n.func) for n in ast.walk(tree)
|
|
if isinstance(n, ast.Call) and isinstance(n.func, ast.Name)}
|
|
names = {n.id for n in ast.walk(tree)
|
|
if isinstance(n, ast.Name) and id(n) not in func_ids}
|
|
if len(names) > _BASE_FEATURES_MAX:
|
|
out.append(Violation("complexity", f"引用列数 {len(names)}>{_BASE_FEATURES_MAX}"))
|
|
nodes = _size(tree.body)
|
|
if nodes > _NODE_COUNT_MAX:
|
|
out.append(Violation("complexity", f"AST 节点数 {nodes}>{_NODE_COUNT_MAX}"))
|
|
lits = sum(1 for n in ast.walk(tree)
|
|
if isinstance(n, ast.Constant) and not isinstance(n.value, bool))
|
|
if lits > _NUMERIC_LIT_MAX:
|
|
out.append(Violation("complexity", f"数字字面量 {lits}>{_NUMERIC_LIT_MAX}"))
|
|
if _depth(tree.body) > _DEPTH_MAX:
|
|
out.append(Violation("complexity", f"嵌套深度 {_depth(tree.body)}>{_DEPTH_MAX}"))
|
|
return out
|
|
|
|
|
|
def _subtree_sizes(tree: ast.AST) -> dict[tuple, int]:
|
|
"""子树键→节点数(同键取最大;键相等⟹结构相等⟹节点数相等).
|
|
|
|
自 tree.body 起遍历并剔 ctx 节点:Expression 根/Load 是一切表达式共有
|
|
的退化"公共子树",不剔则任何 ≥8 节点候选对任何在库式必误报 dup_subtree.
|
|
"""
|
|
out: dict[tuple, int] = {}
|
|
for node in ast.walk(tree.body):
|
|
if isinstance(node, _CTX_NODES):
|
|
continue
|
|
k = _key(node)
|
|
s = _size(node)
|
|
if s > out.get(k, 0):
|
|
out[k] = s
|
|
return out
|
|
|
|
|
|
def dup_violations(expr: str, library_exprs: dict[str, str],
|
|
*, threshold: int = _DUP_SUBTREE_MIN) -> list[Violation]:
|
|
"""③语法级防换皮:整树同构(逐字重发/交换律重排)或 ≥threshold 节点公共子树.
|
|
|
|
与⑧IC 正交闸组成两级:语法挡重写法,值级挡换写法同值.
|
|
"""
|
|
tree = _parse_or_none(expr)
|
|
if tree is None:
|
|
return []
|
|
whole = _key(tree.body, swap=False) # 严格序:交换律重排≠逐字重发
|
|
whole_swap = _key(tree.body) # P3-9:交换律同构键(小因子换皮补检)
|
|
mine = _subtree_sizes(tree)
|
|
for lname, lexpr in library_exprs.items():
|
|
ltree = _parse_or_none(lexpr)
|
|
if ltree is None:
|
|
continue
|
|
if _key(ltree.body, swap=False) == whole:
|
|
return [Violation("dup_exact", f"与在库因子 {lname} 整树同构(仅换名重发)")]
|
|
if whole_swap == _key(ltree.body):
|
|
return [Violation(
|
|
"dup_subtree",
|
|
f"与在库因子 {lname} 整树交换律重排同构(疑似换皮)")]
|
|
theirs = _subtree_sizes(ltree)
|
|
if any(s >= threshold and k in theirs for k, s in mine.items()):
|
|
return [Violation(
|
|
"dup_subtree",
|
|
f"与在库因子 {lname} 存在 ≥{threshold} 节点公共子树(疑似换皮)")]
|
|
return []
|
|
|
|
|
|
def check_candidate(name: str, expr: str, *, existing_names: set[str],
|
|
library_exprs: dict[str, str],
|
|
unwired_domains: tuple[str, ...] = ()) -> list[Violation]:
|
|
"""分解候选总门:名字+白名单+复杂度+防换皮+来源推导(全确定性零LLM)."""
|
|
out = check_name(name, existing_names)
|
|
pre = precheck_expression(expr, unwired_domains)
|
|
out += pre
|
|
if not any(v.code == "syntax_error" for v in out):
|
|
out += complexity_violations(expr)
|
|
out += dup_violations(expr, library_exprs)
|
|
if not pre: # derive_source 契约=调用方先过干净 precheck(未知列会 KeyError)
|
|
src = derive_source(expr)
|
|
if isinstance(src, Violation):
|
|
out.append(src)
|
|
return out
|