Files

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