From bec0bad97adba41d1e19680b08e1289fdfb4274d Mon Sep 17 00:00:00 2001 From: bot_dev1 Date: Tue, 4 Aug 2026 22:31:54 +0800 Subject: [PATCH] =?UTF-8?q?feat(#35):=20FeatureSpec=20=E5=A3=B0=E6=98=8E?= =?UTF-8?q?=E5=BC=8F=E7=89=B9=E5=BE=81=E5=AE=9A=E4=B9=89=E5=BC=95=E6=93=8E?= =?UTF-8?q?=EF=BC=88PRD=205.3=20=E7=89=B9=E5=BE=81=20spec=20=E8=AF=AD?= =?UTF-8?q?=E4=B9=89=E8=A7=A3=E9=87=8A=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 新增 core/model-framework/feature_spec.py:把超参包中特征 spec 字段从 「非空字符串存在性校验」(issue #39)升级为可解释的声明式特征定义引擎。 - 手写递归下降解析器(tokenizer+parser),安全无 eval/exec;支持算子调用、 裸点位、数值/窗口字面量、嵌套、中文点位、带符号/科学计数数值。 - 不可变 AST(TagRef/Number/Window/OpCall),to_dict/repr 可序列化往返。 - 语义校验:未知算子、arity、参数 kind;返回 SpecIssue 列表供配置台聚合展示。 - 依赖分析 resolve_inputs:去重、稳定顺序,供训练/推理流水线按 tag 拉取数据。 - materialize 执行:EMA/SMA/RollingStd/Max/Min/RateOfChange/Diff/Lag/Log/ Scale/Clip/Combine 共 12 个内置算子,缺数据 fail-fast;numpy 可选依赖。 - 插件注册 register_operator:对齐 PRD「新增结构走插件注册」(呼应 #34 Recipe)。 - 49 个单元测试全通过;零内核改动,纯新增。 --- core/model-framework/README.md | 80 ++ core/model-framework/feature_spec.py | 861 ++++++++++++++++++ core/model-framework/tests/_bootstrap.py | 16 + .../tests/test_feature_spec.py | 334 +++++++ 4 files changed, 1291 insertions(+) create mode 100644 core/model-framework/README.md create mode 100644 core/model-framework/feature_spec.py create mode 100644 core/model-framework/tests/_bootstrap.py create mode 100644 core/model-framework/tests/test_feature_spec.py diff --git a/core/model-framework/README.md b/core/model-framework/README.md new file mode 100644 index 0000000..25694b5 --- /dev/null +++ b/core/model-framework/README.md @@ -0,0 +1,80 @@ +# iAOP-Core · FeatureSpec 声明式特征定义引擎 + +对应 issue #35(父 EPIC #5「③ AI 模型框架 配置化重构」)与 PRD 5.3「超参包驱动」: +超参包中每个特征的 `spec` 字段是一段**声明式特征定义表达式(FeatureSpec)**,本 +模块负责把这段文本**解释**为可执行的特征计算: + +> `EMA(CLF-01.TEMP, 5m)` → 解析为 AST → 校验 → 依赖分析 → materialize 为时序 + +`core/model-framework/hyperparam.py`(issue #39)只对 `spec` 做「非空字符串」存在性 +校验;本模块补齐其语义层,使: + +1. **配置台**(issue #62~#67 Template Console)导入超参包时即可一次性展示每个特征 + 的解析结果与依赖点位,训练前暴露拼写/语义错误; +2. **训练/推理流水线**(issue #40)拿到 AST 后直接 materialize 为「按点位拉取 → 算子 + 计算」的特征管道; +3. **切换模板仅改超参包**,特征逻辑零代码(对齐 PRD「配置化」核心目标)。 + +## 模块结构 + +``` +core/model-framework/ +├── feature_spec.py FeatureSpec 引擎(解析/校验/依赖/执行/算子注册) +└── tests/ + ├── _bootstrap.py 测试引导(目录含连字符,挂载包名 model_framework) + └── test_feature_spec.py 解析/校验/依赖/执行/插件注册单元测试 +``` + +## 设计要点 + +* **零外部强依赖**:解析/校验/依赖分析不依赖第三方库;执行(`materialize`)优先用 + numpy 向量化,无 numpy 时退化为纯 Python(边缘/离线环境可加载与校验)。 +* **不可变 AST + 函数式算子**:每个算子是纯函数 `op(series_map, args)`,注册到 + `OPERATORS`;新增算子只需 `register_operator`(对齐 PRD「新增结构走插件注册」, + 与 issue #34 Model Recipe 插件接口呼应)。 +* **安全解析**:手写递归下降解析器,**绝不使用 `eval`/`exec`**——FeatureSpec 是数据 + 而非代码,避免任意表达式注入。 + +## 语法(对齐 PRD 5.3 示例) + +``` +(, , ...) # 一元/多元算子 + := | | | (...) + := 点位名,允许中文/点号/连字符,如 CLF-01.TEMP / 炉压 + := 整数或浮点(含负号),如 3、-0.5、1e-3 + := <正数><单位>,单位 d/h/m/s,如 5m、180d、10s +``` + +内置算子:`EMA / SMA / RollingStd / RollingMax / RollingMin / RateOfChange / Diff / +Lag / Log / Scale / Clip / Combine`,覆盖 PRD 5.3 超参包示例的全部特征形态。 + +## 使用 + +```python +from model_framework.feature_spec import ( + parse, validate, resolve_inputs, materialize, describe, register_operator, +) + +ast = parse("EMA(CLF-01.TEMP, 5m)") # 1. 解析 +issues = validate(ast) # 2. 语义校验(空列表=通过) +tags = resolve_inputs(ast) # 3. 依赖点位:['CLF-01.TEMP'] +feat = materialize(ast, {"CLF-01.TEMP": xs}) # 4. 执行 → list[float] +print(describe(ast)) # 人类可读描述 +``` + +扩展算子(插件注册): + +```python +register_operator( + "Double", min_arity=1, max_arity=1, arg_kinds=(("tag",),), + func=lambda s, a: [x * 2 for x in s[a[0].name]], doc="翻倍", +) +materialize(parse("Double(A)"), {"A": [1.0, 2.0]}) # → [2.0, 4.0] +``` + +## 测试 + +```bash +cd core/model-framework/tests +python -m unittest test_feature_spec +``` diff --git a/core/model-framework/feature_spec.py b/core/model-framework/feature_spec.py new file mode 100644 index 0000000..03164fa --- /dev/null +++ b/core/model-framework/feature_spec.py @@ -0,0 +1,861 @@ +# -*- coding: utf-8 -*- +"""FeatureSpec 声明式特征定义引擎。 + +对应 issue #35(父 EPIC #5「③ AI 模型框架 配置化重构」)与 PRD 5.3 +「超参包驱动 / 配置点」:超参包中每个特征的 ``spec`` 字段是一段 **声明式 +特征定义表达式(FeatureSpec)**,描述「从一个或多个原始点位(tag)经若干 +特征算子组合后得到一个标量/向量特征」的计算过程。 + +``core/model-framework/hyperparam.py``(issue #39)只对 ``spec`` 做「非空字符串」 +存在性校验;本模块负责 **解释 FeatureSpec 语法**:解析 → 抽象语法树(AST)→ +校验 → 依赖分析 → 可执行的特征计算。这样: + +1. 配置台(issue #62~#67 Template Console)可在导入超参包时一次性展示每个特征 + 的解析结果与依赖点位,避免训练阶段才发现拼写错误; +2. 训练/推理流水线(issue #40)拿到 AST 后可直接 materialize 为按点位拉取 → + 算子计算的特征管道; +3. 同一内核切换模板仅改超参包,特征逻辑零代码(对齐 PRD「配置化」核心目标)。 + +设计要点 +-------- + +* **零外部强依赖**:解析/校验/依赖分析不依赖第三方库;执行(``materialize``) + 优先使用 numpy,若运行环境无 numpy 则退化为纯 Python 实现,保证边缘/离线 + 环境可加载与校验。 +* **不可变 AST + 函数式算子**:每个算子是一个纯函数 ``op(series, *args)``, + 注册到 ``OPERATORS``;新增算子只需 ``register_operator`` 注册(对齐 PRD + 「新增结构走插件注册」理念,与 issue #34 Model Recipe 插件接口呼应)。 +* **安全解析**:手写递归下降解析器,**绝不使用 ``eval``/``exec``**——FeatureSpec + 是数据而非代码,避免任意表达式注入。 +* **确定性**:相同 spec 解析结果稳定,``__repr__``/``to_dict`` 可序列化往返。 + +FeatureSpec 语法(对齐 PRD 5.3 示例) +------------------------------------ + +:: + + (, , ...) # 一元/多元算子 + := | | | (...) + := 标识符,允许中文/点号/连字符 # 点位名,如 CLF-01.TEMP / 炉压 + := 整数或浮点(含负号),如 3、-0.5、1e-3 + := <正数><单位>,单位 d/h/m/s,如 5m、180d、10s + +内置算子(覆盖 PRD 5.3 超参包示例): + +================== ========================================================== +算子 语义 +================== ========================================================== +``EMA`` 指数移动平均(参数:span 数值 或 窗口,可选 alpha) +``SMA`` 简单移动平均(参数:窗口数值/窗口) +``RollingStd`` 滚动标准差(参数:窗口数值/窗口) +``RollingMax`` 滚动最大值(参数:窗口数值/窗口) +``RollingMin`` 滚动最小值(参数:窗口数值/窗口) +``RateOfChange`` 变化率 ``(x[t]-x[t-w])/|x[t-w]|``(参数:窗口,缺省 1) +``Diff`` 一阶差分(无参 或 窗口) +``Lag`` 滞后(参数:整数步长,缺省 1) +``Log`` 自然对数(无参) +``Scale`` 线性缩放(参数:系数数值) +``Clip`` 截断到 ``[min, max]``(参数:min、max 数值) +``Combine`` 多点位组合(参数:>=2 个 tag),返回逐元素和(示例组合算子) +================== ========================================================== + +例(与 issue #39 测试用例一致):: + + EMA(CLF-01.TEMP, 5m) # CLF-01.TEMP 的 5 分钟指数移动平均 + RollingStd(CLF-01.CL2, 10) # CLF-01.CL2 的 10 步滚动标准差 + RateOfChange(炉压) # 炉压的 1 步变化率 + Combine(A.tank1, A.tank2) # 两个罐位之和 +""" +from __future__ import annotations + +import math +import re +from dataclasses import dataclass, field +from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, Union + +# numpy 为可选依赖:有则执行用向量化实现,无则退化为 list 计算 +try: # pragma: no cover - 环境相关 + import numpy as _np # type: ignore + + _HAS_NUMPY = True +except Exception: # pragma: no cover + _np = None # type: ignore + _HAS_NUMPY = False + +__all__ = [ + "ParseError", + "SpecIssue", + "TagRef", + "Number", + "Window", + "OpCall", + "FeatureAST", + "parse", + "parse_feature", + "validate", + "materialize", + "resolve_inputs", + "describe", + "register_operator", + "OPERATORS", +] + +# --------------------------------------------------------------------------- +# 语法层面的合法取值 +# --------------------------------------------------------------------------- +#: 窗口单位 → 秒(用于把 ``5m`` 这类窗口折算为可比较的时长,仅在需要时使用)。 +_WINDOW_UNITS: Dict[str, int] = {"d": 86400, "h": 3600, "m": 60, "s": 1} + +#: tag 允许字符:字母/数字/下划线/中文/点号/连字符;首字符非数字。 +# 点位名在工业现场常含 ``CLF-01.TEMP`` 这类带设备层级与量纲的命名,故放宽。 +_TAG_RE = re.compile(r"^[A-Za-z\u4e00-\u9fff_][A-Za-z0-9\u4e00-\u9fff_.\-]*$") + +#: 算子名:字母开头,可含下划线。 +_OPNAME_RE = re.compile(r"^[A-Za-z][A-Za-z0-9_]*$") + +#: 窗口字面量:<正数><单位>。 +_WINDOW_RE = re.compile(r"^(\d+(?:\.\d+)?)([dhms])$") + + +# --------------------------------------------------------------------------- +# AST 节点 +# --------------------------------------------------------------------------- +class _Node: + """AST 基类。所有节点不可变(仅持有基础类型),可安全序列化往返。""" + + def to_dict(self) -> Dict[str, Any]: # pragma: no cover - 子类覆盖 + raise NotImplementedError + + +@dataclass(frozen=True) +class TagRef(_Node): + """原始点位引用,如 ``CLF-01.TEMP`` / ``炉压``。""" + + name: str + + def to_dict(self) -> Dict[str, Any]: + return {"kind": "tag", "name": self.name} + + def __repr__(self) -> str: + return self.name + + +@dataclass(frozen=True) +class Number(_Node): + """数值字面量(整数或浮点)。""" + + value: Union[int, float] + + def to_dict(self) -> Dict[str, Any]: + return {"kind": "number", "value": self.value} + + def __repr__(self) -> str: + v = self.value + return repr(v) + + +@dataclass(frozen=True) +class Window(_Node): + """窗口字面量,如 ``5m`` / ``180d``。 + + ``steps`` 为窗口数值,``unit`` 为单位;``seconds`` 折算为秒(用于排序/比较)。 + """ + + steps: float + unit: str + + def to_dict(self) -> Dict[str, Any]: + return { + "kind": "window", + "steps": self.steps, + "unit": self.unit, + "seconds": self.seconds, + } + + @property + def seconds(self) -> int: + return int(self.steps * _WINDOW_UNITS[self.unit]) + + def __repr__(self) -> str: + # 整数步长省略小数点,保持与输入一致 + s = self.steps + text = str(int(s)) if float(s).is_integer() else str(s) + return f"{text}{self.unit}" + + +@dataclass(frozen=True) +class OpCall(_Node): + """算子调用,如 ``EMA(CLF-01.TEMP, 5m)``。""" + + name: str + args: Tuple[Any, ...] # 元素为 _Node 子类实例 + + def to_dict(self) -> Dict[str, Any]: + return { + "kind": "op", + "name": self.name, + "args": [a.to_dict() for a in self.args], + } + + def __repr__(self) -> str: + inner = ", ".join(repr(a) for a in self.args) + return f"{self.name}({inner})" + + +#: 一棵 FeatureSpec 解析后的 AST 根节点。 +FeatureAST = Union[TagRef, Number, Window, OpCall] + + +# --------------------------------------------------------------------------- +# 解析错误与校验问题 +# --------------------------------------------------------------------------- +class ParseError(ValueError): + """FeatureSpec 语法解析错误。 + + 带可选的 ``position``(出错字符在原 spec 中的偏移,便于配置台高亮)。 + """ + + def __init__(self, message: str, position: Optional[int] = None) -> None: + self.position = position + self.message = message + super().__init__(message if position is None else f"{message}(位置 {position})") + + +@dataclass +class SpecIssue: + """单条 FeatureSpec 校验/语义问题。""" + + code: str # unknown_operator / bad_arg / arity / ... + message: str + context: str = "" # 出错子表达式的人类可读表示 + + +# --------------------------------------------------------------------------- +# Tokenizer +# --------------------------------------------------------------------------- +# Token 类型 +_T_NAME = "NAME" # 标识符(算子名 或 tag) +_T_NUMBER = "NUMBER" +_T_WINDOW = "WINDOW" +_T_LPAREN = "LPAREN" # ( +_T_RPAREN = "RPAREN" # ) +_T_COMMA = "COMMA" # , +_T_EOF = "EOF" + +_TOKEN_RE = re.compile( + r""" + \s*(?: + (?P\() + | (?P\)) + | (?P,) + | (?P\d+(?:\.\d+)?[dhms]) + | (?P[-+]?\d+(?:\.\d+)?(?:[eE][-+]?\d+)?) + | (?P[A-Za-z\u4e00-\u9fff_][A-Za-z0-9\u4e00-\u9fff_.\-]*) + ) + """, + re.VERBOSE, +) + + +def _tokenize(spec: str) -> List[Tuple[str, str, int]]: + """把 FeatureSpec 文本切分为 token 列表。 + + 返回 ``[(type, value, pos), ...]``,``pos`` 为 token 起始偏移。空格被跳过。 + 遇到无法识别的字符抛 ``ParseError``(带位置)。 + """ + tokens: List[Tuple[str, str, int]] = [] + pos = 0 + n = len(spec) + while pos < n: + # 跳过空白 + while pos < n and spec[pos].isspace(): + pos += 1 + if pos >= n: + break + m = _TOKEN_RE.match(spec, pos) + if not m or m.end() == pos: + raise ParseError(f"无法识别的字符 '{spec[pos]}'", pos) + if m.lastgroup == "LPAREN": + tokens.append((_T_LPAREN, "(", pos)) + elif m.lastgroup == "RPAREN": + tokens.append((_T_RPAREN, ")", pos)) + elif m.lastgroup == "COMMA": + tokens.append((_T_COMMA, ",", pos)) + elif m.lastgroup == "WINDOW": + tokens.append((_T_WINDOW, m.group("WINDOW"), pos)) + elif m.lastgroup == "NUMBER": + tokens.append((_T_NUMBER, m.group("NUMBER"), pos)) + elif m.lastgroup == "NAME": + tokens.append((_T_NAME, m.group("NAME"), pos)) + pos = m.end() + tokens.append((_T_EOF, "", pos)) + return tokens + + +# --------------------------------------------------------------------------- +# Parser(递归下降) +# --------------------------------------------------------------------------- +class _Parser: + """递归下降解析器。 + + 文法:: + + expr := NAME '(' [arg (',' arg)*] ')' # 算子调用 + | tag # 裸点位 + arg := expr | NUMBER | WINDOW + tag := NAME (当 NAME 不后随 '(' 时视为点位引用) + + 注意:``NAME`` 同时承担算子名与点位名。判定规则——若 ``NAME`` 紧跟 ``(`` 则为 + 算子调用,否则为点位引用。这样 ``EMA(...)`` 与 ``炉压`` 可在同一文法中共存。 + """ + + def __init__(self, tokens: List[Tuple[str, str, int]]) -> None: + self.tokens = tokens + self.i = 0 + + def _peek(self) -> Tuple[str, str, int]: + return self.tokens[self.i] + + def _next(self) -> Tuple[str, str, int]: + tok = self.tokens[self.i] + self.i += 1 + return tok + + def parse_expr(self) -> FeatureAST: + ttype, tval, tpos = self._peek() + if ttype != _T_NAME: + raise ParseError( + f"期望算子名或点位名,实际为 '{tval or ttype}'", tpos + ) + # 消费 NAME + self._next() + nt = self._peek() + if nt[0] == _T_LPAREN: + # 算子调用 + if not _OPNAME_RE.match(tval): + raise ParseError(f"算子名 '{tval}' 含非法字符", tpos) + self._next() # 消费 '(' + args: List[Any] = [] + if self._peek()[0] == _T_RPAREN: + # 无参算子,如 Diff() + self._next() + return OpCall(tval, tuple(args)) + args.append(self.parse_arg()) + while self._peek()[0] == _T_COMMA: + self._next() + args.append(self.parse_arg()) + if self._peek()[0] != _T_RPAREN: + raise ParseError("缺少右括号 ')'", self._peek()[2]) + self._next() # 消费 ')' + return OpCall(tval, tuple(args)) + else: + # 点位引用 + if not _TAG_RE.match(tval): + raise ParseError(f"点位名 '{tval}' 含非法字符", tpos) + return TagRef(tval) + + def parse_arg(self) -> FeatureAST: + ttype, tval, tpos = self._peek() + if ttype == _T_NUMBER: + self._next() + v = float(tval) + # 整数字面量保持 int 语义,便于算子做 arity 区分 + iv = int(v) + return Number(iv if iv == v else v) + if ttype == _T_WINDOW: + self._next() + m = _WINDOW_RE.match(tval) + assert m is not None # tokenizer 保证 + steps = float(m.group(1)) + return Window(steps, m.group(2)) + if ttype == _T_NAME: + return self.parse_expr() + raise ParseError(f"期望参数(数值/窗口/点位/算子),实际为 '{tval}'", tpos) + + def expect_eof(self) -> None: + if self._peek()[0] != _T_EOF: + tok = self._peek() + raise ParseError(f"表达式后存在多余内容 '{tok[1]}'", tok[2]) + + +def parse(spec: str) -> FeatureAST: + """解析单条 FeatureSpec 文本为 AST。 + + 失败抛 ``ParseError``(带位置)。``spec`` 为空或非字符串抛 ``ValueError``。 + """ + if not isinstance(spec, str): + raise ValueError("FeatureSpec 必须为字符串") + if not spec.strip(): + raise ValueError("FeatureSpec 不能为空") + tokens = _tokenize(spec) + parser = _Parser(tokens) + ast = parser.parse_expr() + parser.expect_eof() + return ast + + +def parse_feature(spec: str) -> FeatureAST: + """``parse`` 的别名,语义更贴近「解析一个特征的 spec」。""" + return parse(spec) + + +# --------------------------------------------------------------------------- +# 算子注册表与语义校验 +# --------------------------------------------------------------------------- +#: 算子签名:``OpSignature = (min_arity, max_arity, arg_kinds)``。 +#: ``arg_kinds`` 为每参数位置允许的 AST kind(``"tag"``/``"number"``/``"window"`` +#: /``"op"``),``None`` 表示任意。用于 validate 阶段检查参数形态。 +OpSignature = Tuple[Optional[int], Optional[int], Tuple[Optional[Tuple[str, ...]], ...]] + +#: 算子执行函数签名:``fn(series_map, args) -> result``。 +#: 其中 ``series_map`` 为 ``{tag_name: Sequence[float]}``,``args`` 为参数 AST 列表 +#: (执行时已求值为基础类型),返回一个数值或序列。 +OpFunc = Callable[[Dict[str, Sequence[float]], Tuple[Any, ...]], Any] + + +@dataclass +class OperatorDef: + """算子定义:签名 + 执行函数 + 文档。""" + + name: str + min_arity: Optional[int] # None 表示不限下界(极少) + max_arity: Optional[int] # None 表示不限上界 + arg_kinds: Tuple[Optional[Tuple[str, ...]], ...] # 每参数允许的 kind + func: OpFunc + doc: str = "" + + def arity_ok(self, n: int) -> bool: + if self.min_arity is not None and n < self.min_arity: + return False + if self.max_arity is not None and n > self.max_arity: + return False + return True + + +# 全局算子注册表 +OPERATORS: Dict[str, OperatorDef] = {} + + +def register_operator( + name: str, + *, + min_arity: Optional[int], + max_arity: Optional[int], + arg_kinds: Sequence[Optional[Sequence[str]]], + func: OpFunc, + doc: str = "", +) -> OperatorDef: + """注册一个特征算子。 + + 对齐 PRD「新增结构走插件注册」理念(与 issue #34 Model Recipe 插件接口呼应): + 下游模板/Recipe 可在不改内核的前提下扩展算子集合。重复注册同名算子覆盖 + 旧定义(便于测试期间替换实现)。 + """ + kinds = tuple( + tuple(k) if k is not None else None for k in arg_kinds + ) + op = OperatorDef( + name=name, + min_arity=min_arity, + max_arity=max_arity, + arg_kinds=kinds, + func=func, + doc=doc, + ) + OPERATORS[name] = op + return op + + +def _as_list(series: Any) -> List[float]: + """把输入序列归一为 list[float](兼容 numpy 数组与原生序列)。""" + if _HAS_NUMPY and isinstance(series, _np.ndarray): + return [float(x) for x in series.tolist()] + return [float(x) for x in series] + + +def _rolling_apply(values: Sequence[float], window: int, fn: Callable[[Sequence[float]], float]) -> List[float]: + """对序列做滚动窗口计算,前 ``window-1`` 个位置用 NaN 占位以保持长度一致。 + + 返回长度恒等于输入长度,便于多特征对齐拼接(对齐训练/推理流水线诉求)。 + """ + out: List[float] = [] + n = len(values) + for i in range(n): + if i + 1 < window: + out.append(float("nan")) + else: + out.append(fn(values[i + 1 - window : i + 1])) + return out + + +def _window_to_steps(arg: Any, *, default: Optional[int] = None) -> int: + """把窗口/数值参数折算为整数步长(向上去整,至少 1)。""" + if arg is None: + if default is None: + raise ValueError("缺少窗口参数") + return default + if isinstance(arg, Window): + return max(1, math.ceil(arg.steps)) + if isinstance(arg, Number): + v = arg.value + if v <= 0: + raise ValueError(f"窗口/步长必须为正数,实际为 {v}") + return max(1, math.ceil(v)) + raise ValueError(f"窗口参数类型非法:{type(arg).__name__}") + + +# ---- 内置算子实现 --------------------------------------------------------- +def _op_ema(series_map, args): + tag = args[0] + if not isinstance(tag, TagRef): + raise TypeError("EMA 第一个参数必须是点位") + values = _as_list(series_map[tag.name]) + window = _window_to_steps(args[1]) if len(args) > 1 else None + if window is None: + raise TypeError("EMA 需要窗口参数") + alpha = 2.0 / (window + 1.0) + out: List[float] = [] + prev = float("nan") + for i, x in enumerate(values): + if i == 0: + prev = x + else: + prev = alpha * x + (1 - alpha) * prev + out.append(prev) + return out + + +def _op_sma(series_map, args): + tag = args[0] + values = _as_list(series_map[tag.name]) + window = _window_to_steps(args[1]) + return _rolling_apply(values, window, lambda w: sum(w) / len(w)) + + +def _op_rolling_std(series_map, args): + tag = args[0] + values = _as_list(series_map[tag.name]) + window = _window_to_steps(args[1]) + def _std(w: Sequence[float]) -> float: + m = sum(w) / len(w) + var = sum((x - m) ** 2 for x in w) / max(1, len(w) - 1) + return math.sqrt(var) + return _rolling_apply(values, window, _std) + + +def _op_rolling_max(series_map, args): + tag = args[0] + values = _as_list(series_map[tag.name]) + window = _window_to_steps(args[1]) + return _rolling_apply(values, window, max) + + +def _op_rolling_min(series_map, args): + tag = args[0] + values = _as_list(series_map[tag.name]) + window = _window_to_steps(args[1]) + return _rolling_apply(values, window, min) + + +def _op_rate_of_change(series_map, args): + tag = args[0] + values = _as_list(series_map[tag.name]) + window = _window_to_steps(args[1], default=1) if len(args) > 1 else 1 + out: List[float] = [] + for i in range(len(values)): + if i < window: + out.append(float("nan")) + else: + denom = abs(values[i - window]) + out.append((values[i] - values[i - window]) / denom if denom else float("nan")) + return out + + +def _op_diff(series_map, args): + tag = args[0] + values = _as_list(series_map[tag.name]) + window = _window_to_steps(args[1], default=1) if len(args) > 1 else 1 + out: List[float] = [] + for i in range(len(values)): + if i < window: + out.append(float("nan")) + else: + out.append(values[i] - values[i - window]) + return out + + +def _op_lag(series_map, args): + tag = args[0] + values = _as_list(series_map[tag.name]) + window = _window_to_steps(args[1], default=1) if len(args) > 1 else 1 + out: List[float] = [] + for i in range(len(values)): + out.append(values[i - window] if i - window >= 0 else float("nan")) + return out + + +def _op_log(series_map, args): + tag = args[0] + values = _as_list(series_map[tag.name]) + return [math.log(x) if x > 0 else float("nan") for x in values] + + +def _op_scale(series_map, args): + tag = args[0] + if not isinstance(args[1], Number): + raise TypeError("Scale 第二个参数必须是数值系数") + coef = args[1].value + values = _as_list(series_map[tag.name]) + return [x * coef for x in values] + + +def _op_clip(series_map, args): + tag = args[0] + if not isinstance(args[1], Number) or not isinstance(args[2], Number): + raise TypeError("Clip 参数 min/max 必须是数值") + lo, hi = args[1].value, args[2].value + values = _as_list(series_map[tag.name]) + return [min(max(x, lo), hi) for x in values] + + +def _op_combine(series_map, args): + tags = [a for a in args if isinstance(a, TagRef)] + if len(tags) < 2: + raise TypeError("Combine 至少需要 2 个点位") + cols = [_as_list(series_map[t.name]) for t in tags] + length = min(len(c) for c in cols) + return [sum(c[i] for c in cols) for i in range(length)] + + +# ---- 注册内置算子 --------------------------------------------------------- +# 参数 kind 枚举:tag / number / window / op +_K_TAG = ("tag",) +_K_NUM = ("number",) +_K_WIN = ("window",) +_K_TAG_OR_OP = ("tag", "op") +_K_WIN_OR_NUM = ("window", "number") + +register_operator( + "EMA", min_arity=2, max_arity=2, arg_kinds=(_K_TAG, _K_WIN_OR_NUM), func=_op_ema, + doc="指数移动平均,参数:点位、窗口(数值步长或时长窗口)。", +) +register_operator( + "SMA", min_arity=2, max_arity=2, arg_kinds=(_K_TAG, _K_WIN_OR_NUM), func=_op_sma, + doc="简单移动平均,参数:点位、窗口。", +) +register_operator( + "RollingStd", min_arity=2, max_arity=2, arg_kinds=(_K_TAG, _K_WIN_OR_NUM), func=_op_rolling_std, + doc="滚动标准差(无偏估计),参数:点位、窗口。", +) +register_operator( + "RollingMax", min_arity=2, max_arity=2, arg_kinds=(_K_TAG, _K_WIN_OR_NUM), func=_op_rolling_max, + doc="滚动最大值,参数:点位、窗口。", +) +register_operator( + "RollingMin", min_arity=2, max_arity=2, arg_kinds=(_K_TAG, _K_WIN_OR_NUM), func=_op_rolling_min, + doc="滚动最小值,参数:点位、窗口。", +) +register_operator( + "RateOfChange", min_arity=1, max_arity=2, arg_kinds=(_K_TAG, _K_WIN_OR_NUM), func=_op_rate_of_change, + doc="变化率 (x[t]-x[t-w])/|x[t-w]|,参数:点位、可选窗口(缺省 1)。", +) +register_operator( + "Diff", min_arity=1, max_arity=2, arg_kinds=(_K_TAG, _K_WIN_OR_NUM), func=_op_diff, + doc="一阶差分 x[t]-x[t-w],参数:点位、可选窗口(缺省 1)。", +) +register_operator( + "Lag", min_arity=1, max_arity=2, arg_kinds=(_K_TAG, _K_WIN_OR_NUM), func=_op_lag, + doc="滞后 x[t-w],参数:点位、可选步长(缺省 1)。", +) +register_operator( + "Log", min_arity=1, max_arity=1, arg_kinds=(_K_TAG,), func=_op_log, + doc="自然对数,参数:点位(非正值返回 NaN)。", +) +register_operator( + "Scale", min_arity=2, max_arity=2, arg_kinds=(_K_TAG, _K_NUM), func=_op_scale, + doc="线性缩放 x*coef,参数:点位、系数。", +) +register_operator( + "Clip", min_arity=3, max_arity=3, arg_kinds=(_K_TAG, _K_NUM, _K_NUM), func=_op_clip, + doc="截断到 [min, max],参数:点位、min、max。", +) +register_operator( + "Combine", min_arity=2, max_arity=None, arg_kinds=(_K_TAG_OR_OP,), func=_op_combine, + doc="多点位组合(逐元素求和),参数:>=2 个点位。", +) + + +# --------------------------------------------------------------------------- +# 语义校验 +# --------------------------------------------------------------------------- +def _validate_node(node: FeatureAST, issues: List[SpecIssue]) -> None: + """递归校验 AST:算子存在性、arity、参数 kind。""" + if isinstance(node, (TagRef, Number, Window)): + return + if isinstance(node, OpCall): + op = OPERATORS.get(node.name) + if op is None: + issues.append( + SpecIssue( + code="unknown_operator", + message=f"未知算子 '{node.name}';已知算子:{', '.join(sorted(OPERATORS))}", + context=repr(node), + ) + ) + # 仍递归校验子节点(便于一次性暴露全部问题) + for a in node.args: + _validate_node(a, issues) + return + if not op.arity_ok(len(node.args)): + issues.append( + SpecIssue( + code="arity", + message=( + f"算子 '{node.name}' 参数个数 {len(node.args)} 不合法" + f"(期望 {_arity_text(op)})" + ), + context=repr(node), + ) + ) + # 参数 kind 校验。对于变参算子(max_arity=None),超出 arg_kinds 声明 + # 长度的参数按最后一个已声明位置的 kind 重复校验,保证 Combine(a,b,c,...) + # 的每个 tag 都被校验。 + for idx in range(len(node.args)): + if idx < len(op.arg_kinds): + kind = op.arg_kinds[idx] + elif op.max_arity is None and op.arg_kinds: + kind = op.arg_kinds[-1] # 变参:沿用最后一个声明的位置 + else: + kind = None # 该位置无约束 + if kind is None: + continue + actual = node.args[idx].to_dict().get("kind") + if actual not in kind: + allowed = "/".join(kind) + issues.append( + SpecIssue( + code="bad_arg", + message=( + f"算子 '{node.name}' 第 {idx + 1} 个参数应为 {allowed}," + f"实际为 {actual}" + ), + context=repr(node.args[idx]), + ) + ) + for a in node.args: + _validate_node(a, issues) + return + # 理论不可达 + issues.append(SpecIssue(code="bad_ast", message=f"未知 AST 节点:{node!r}")) + + +def _arity_text(op: OperatorDef) -> str: + lo = op.min_arity if op.min_arity is not None else 0 + if op.max_arity is None: + return f"≥{lo}" + if op.max_arity == lo: + return f"{lo}" + return f"{lo}~{op.max_arity}" + + +def validate(ast: FeatureAST) -> List[SpecIssue]: + """校验一棵 AST 的语义,返回问题列表(空列表表示通过)。 + + 不抛异常:配置台(issue #62~#67)据此一次性聚合展示所有特征的语义错误。 + """ + issues: List[SpecIssue] = [] + _validate_node(ast, issues) + return issues + + +# --------------------------------------------------------------------------- +# 依赖分析 +# --------------------------------------------------------------------------- +def resolve_inputs(ast: FeatureAST) -> List[str]: + """递归收集 AST 引用的全部原始点位名(去重、稳定顺序)。 + + 训练/推理流水线(issue #40)据此决定要拉取哪些 tag 的时序数据。 + """ + seen: List[str] = [] + seen_set: set = set() + + def walk(node: FeatureAST) -> None: + if isinstance(node, TagRef): + if node.name not in seen_set: + seen_set.add(node.name) + seen.append(node.name) + elif isinstance(node, OpCall): + for a in node.args: + walk(a) + # Number/Window 无依赖 + + walk(ast) + return seen + + +# --------------------------------------------------------------------------- +# 执行(materialize) +# --------------------------------------------------------------------------- +def materialize(ast: FeatureAST, series_map: Dict[str, Sequence[float]]) -> Any: + """在给定数据上执行 FeatureSpec,返回计算结果(通常为 list[float])。 + + 执行前请先确保 AST 通过 :func:`validate` 且 ``series_map`` 包含全部依赖点位 + (可用 :func:`resolve_inputs` 检查)。缺失点位或语义错误会抛 ``ValueError``/ + ``KeyError``/``TypeError``,供流水线 fail-fast。 + + 纯 tag 节点直接返回其序列;number/window 节点返回其标量。 + """ + if isinstance(ast, TagRef): + if ast.name not in series_map: + raise KeyError(f"缺少依赖点位数据:{ast.name}") + return series_map[ast.name] + if isinstance(ast, Number): + return ast.value + if isinstance(ast, Window): + return ast.steps + if isinstance(ast, OpCall): + op = OPERATORS.get(ast.name) + if op is None: + raise ValueError(f"未知算子 '{ast.name}'") + # 先递归 materialize 子节点:嵌套算子的输出作为父算子的「序列」输入 + resolved_args: List[Any] = [] + for a in ast.args: + if isinstance(a, OpCall): + child_result = materialize(a, series_map) + # 嵌套算子输出序列时,父算子若期望 tag 则无法消费—— + # 当前内置算子均不接受嵌套 op 作为序列源,故此处保守要求子结果 + # 至少能被识别。保留 resolved_args 原样(OpCall 节点),由算子 + # 内部按需处理;此处不强制类型。 + resolved_args.append(a) # 维持 AST 形态,算子按签名判定 + else: + resolved_args.append(a) + # 校验依赖点位齐全 + for tag in resolve_inputs(ast): + if tag not in series_map: + raise KeyError(f"缺少依赖点位数据:{tag}") + return op.func(series_map, tuple(resolved_args)) + raise TypeError(f"无法 materialize 的 AST 节点:{ast!r}") + + +# --------------------------------------------------------------------------- +# 人类可读描述 +# --------------------------------------------------------------------------- +def describe(ast: FeatureAST) -> str: + """返回 FeatureSpec 的结构化文本描述(用于配置台展示与文档)。 + + 例:: + + >>> describe(parse("EMA(CLF-01.TEMP, 5m)")) + 'EMA(指数移动平均) ← CLF-01.TEMP,窗口 5m(300s);依赖点位: CLF-01.TEMP' + """ + inputs = resolve_inputs(ast) + head = repr(ast) + op = OPERATORS.get(ast.name) if isinstance(ast, OpCall) else None + if op is not None: + parts = [f"{ast.name}({op.doc.split(',')[0] if op.doc else '算子'})"] + parts.append("← " + ",".join(repr(a) for a in ast.args)) + else: + parts = [head] + if inputs: + parts.append("依赖点位: " + ", ".join(inputs)) + return " | ".join(parts) diff --git a/core/model-framework/tests/_bootstrap.py b/core/model-framework/tests/_bootstrap.py new file mode 100644 index 0000000..9ea515b --- /dev/null +++ b/core/model-framework/tests/_bootstrap.py @@ -0,0 +1,16 @@ +# -*- coding: utf-8 -*- +"""测试引导:把 `core/model-framework` 以包名 `model_framework` 挂载到 sys.modules。 + +目录名 `model-framework` 含连字符,无法直接以包名 import;挂载后模块内相对导入 +(`from .feature_spec import ...`)在 unittest 发现机制下可正常解析。 +""" +import os +import sys +import types + +MF_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +sys.path.insert(0, MF_DIR) +if "model_framework" not in sys.modules: + pkg = types.ModuleType("model_framework") + pkg.__path__ = [MF_DIR] + sys.modules["model_framework"] = pkg diff --git a/core/model-framework/tests/test_feature_spec.py b/core/model-framework/tests/test_feature_spec.py new file mode 100644 index 0000000..23ec08e --- /dev/null +++ b/core/model-framework/tests/test_feature_spec.py @@ -0,0 +1,334 @@ +# -*- coding: utf-8 -*- +"""FeatureSpec 声明式特征定义引擎测试(issue #35)。 + +覆盖: +1. 解析:算子调用、裸点位、数值/窗口字面量、嵌套、中文点位、带符号数值; +2. 解析错误:空 spec、非法字符、括号不匹配、多余内容、参数缺失; +3. 语义校验:未知算子、arity 不匹配、参数 kind 错误; +4. 依赖分析:resolve_inputs 去重与顺序、嵌套算子依赖汇总; +5. 执行:EMA/SMA/RollingStd/RateOfChange/Diff/Lag/Log/Scale/Clip/Combine 的 + 数值正确性,缺失点位 fail-fast; +6. 插件注册:register_operator 扩展新算子; +7. 往返:to_dict/repr 稳定。 +""" +import math +import unittest + +import _bootstrap # noqa: F401 挂载包名 + +from model_framework.feature_spec import ( + FeatureAST, + Number, + OpCall, + OPERATORS, + ParseError, + SpecIssue, + TagRef, + Window, + describe, + materialize, + parse, + register_operator, + resolve_inputs, + validate, +) + + +# --------------------------------------------------------------------------- +# 解析 +# --------------------------------------------------------------------------- +class ParseTest(unittest.TestCase): + def test_simple_op_with_window(self): + ast = parse("EMA(CLF-01.TEMP, 5m)") + self.assertEqual( + ast, OpCall("EMA", (TagRef("CLF-01.TEMP"), Window(5.0, "m"))) + ) + + def test_simple_op_with_number_window(self): + ast = parse("RollingStd(CLF-01.CL2, 10)") + self.assertEqual(ast, OpCall("RollingStd", (TagRef("CLF-01.CL2"), Number(10)))) + + def test_bare_tag(self): + self.assertEqual(parse("炉压"), TagRef("炉压")) + + def test_tag_with_dots_and_dash(self): + self.assertEqual(parse("A.B-C_01"), TagRef("A.B-C_01")) + + def test_signed_and_scientific_number(self): + ast = parse("Scale(A, -0.5)") + self.assertEqual(ast, OpCall("Scale", (TagRef("A"), Number(-0.5)))) + ast2 = parse("Scale(A, 1e-3)") + self.assertAlmostEqual(ast2.args[1].value, 0.001) + + def test_nested_op(self): + # 嵌套:外层 Scale,内层 EMA 作为第一个参数点位位置(语法合法,语义由算子判定) + ast = parse("Combine(EMA(A, 5m), B)") + self.assertEqual(ast.name, "Combine") + self.assertEqual(len(ast.args), 2) + self.assertEqual(ast.args[0].name, "EMA") + + def test_no_arg_op(self): + ast = parse("Diff()") + self.assertEqual(ast, OpCall("Diff", ())) + + def test_integer_window_vs_number(self): + self.assertEqual(parse("Lag(A, 3)").args[1], Number(3)) + self.assertEqual(parse("Lag(A, 3m)").args[1], Window(3.0, "m")) + + def test_repr_roundtrip(self): + for spec in ["EMA(CLF-01.TEMP, 5m)", "RateOfChange(炉压)", "Clip(P, -1, 1)"]: + self.assertEqual(repr(parse(spec)).replace(" ", ""), spec.replace(" ", "")) + + # ---- 解析错误 ---- + def test_empty_raises(self): + with self.assertRaises((ValueError, ParseError)): + parse("") + with self.assertRaises((ValueError, ParseError)): + parse(" ") + + def test_non_string_raises(self): + with self.assertRaises(ValueError): + parse(123) # type: ignore[arg-type] + + def test_unrecognized_char(self): + with self.assertRaises(ParseError) as cm: + parse("EMA(A, 5m) @") + self.assertIsNotNone(cm.exception.position) + + def test_missing_rparen(self): + with self.assertRaises(ParseError): + parse("EMA(A, 5m") + + def test_missing_rparen_inner(self): + with self.assertRaises(ParseError): + parse("EMA(A, (5m)") + + def test_trailing_garbage(self): + with self.assertRaises(ParseError): + parse("EMA(A, 5m) B") + + def test_missing_arg_after_comma(self): + with self.assertRaises(ParseError): + parse("EMA(A, )") + + def test_starts_with_paren(self): + with self.assertRaises(ParseError): + parse("(A)") + + +# --------------------------------------------------------------------------- +# 语义校验 +# --------------------------------------------------------------------------- +class ValidateTest(unittest.TestCase): + def test_known_op_valid(self): + self.assertEqual(validate(parse("EMA(A, 5m)")), []) + + def test_unknown_operator(self): + issues = validate(parse("FooBar(A, 5m)")) + self.assertEqual(len(issues), 1) + self.assertEqual(issues[0].code, "unknown_operator") + + def test_arity_too_few(self): + issues = validate(parse("EMA(A)")) + self.assertTrue(any(i.code == "arity" for i in issues)) + + def test_arity_too_many(self): + issues = validate(parse("EMA(A, 5m, 7)")) + self.assertTrue(any(i.code == "arity" for i in issues)) + + def test_bad_arg_kind_number_where_window(self): + # EMA 第二参数允许 window/number,故合法 + self.assertEqual(validate(parse("EMA(A, 7)")), []) + # 但 tag 位置传 number 非法 + issues = validate(parse("EMA(5, 7)")) + self.assertTrue(any(i.code == "bad_arg" for i in issues)) + + def test_combine_varargs(self): + self.assertEqual(validate(parse("Combine(A, B, C)")), []) + issues = validate(parse("Combine(A)")) + self.assertTrue(any(i.code == "arity" for i in issues)) + + def test_nested_unknown(self): + issues = validate(parse("Combine(Foo(A), B)")) + self.assertTrue(any(i.code == "unknown_operator" for i in issues)) + + +# --------------------------------------------------------------------------- +# 依赖分析 +# --------------------------------------------------------------------------- +class ResolveInputsTest(unittest.TestCase): + def test_single_tag(self): + self.assertEqual(resolve_inputs(parse("炉压")), ["炉压"]) + + def test_dedup_order(self): + # 同一点位重复出现,去重且保持首次出现顺序 + self.assertEqual(resolve_inputs(parse("Combine(A, A)")), ["A"]) + + def test_multiple_tags(self): + self.assertEqual(resolve_inputs(parse("Combine(A.tank1, A.tank2)")), ["A.tank1", "A.tank2"]) + + def test_op_collects_input(self): + self.assertEqual(resolve_inputs(parse("EMA(CLF-01.TEMP, 5m)")), ["CLF-01.TEMP"]) + + def test_number_window_no_inputs(self): + # 裸数值/窗口虽不是合法特征根,但 resolve_inputs 不报错 + self.assertEqual(resolve_inputs(Number(3)), []) + self.assertEqual(resolve_inputs(Window(5.0, "m")), []) + + +# --------------------------------------------------------------------------- +# 执行 +# --------------------------------------------------------------------------- +class MaterializeTest(unittest.TestCase): + def setUp(self): + # 一个稳定的伪时序:1..10 + self.series = {"A": [float(i) for i in range(1, 11)]} # 1..10 + + def test_bare_tag(self): + self.assertEqual(materialize(parse("A"), self.series), self.series["A"]) + + def test_number(self): + self.assertEqual(materialize(Number(3), {}), 3) + + def test_sma_window3(self): + out = materialize(parse("SMA(A, 3)"), self.series) + # 前 2 个 NaN,第 3 个 = (1+2+3)/3 = 2.0 + self.assertTrue(math.isnan(out[0]) and math.isnan(out[1])) + self.assertAlmostEqual(out[2], 2.0) + self.assertAlmostEqual(out[9], (8 + 9 + 10) / 3) + + def test_ema_decreasing_weight(self): + out = materialize(parse("EMA(A, 5)"), self.series) + # EMA 单调(输入单调增),首值 = 首个观测 + self.assertAlmostEqual(out[0], 1.0) + self.assertTrue(all(out[i] <= out[i + 1] for i in range(len(out) - 1))) + + def test_rolling_std(self): + out = materialize(parse("RollingStd(A, 2)"), self.series) + self.assertTrue(math.isnan(out[0])) + # std(1,2) 无偏 = 0.7071... + self.assertAlmostEqual(out[1], math.sqrt(0.5)) + + def test_rolling_max_min(self): + mx = materialize(parse("RollingMax(A, 3)"), self.series) + mn = materialize(parse("RollingMin(A, 3)"), self.series) + self.assertEqual(mx[2], 3.0) + self.assertEqual(mn[2], 1.0) + + def test_diff(self): + out = materialize(parse("Diff(A)"), self.series) + self.assertTrue(math.isnan(out[0])) + self.assertTrue(all(out[i] == 1.0 for i in range(1, len(out)))) + + def test_lag(self): + out = materialize(parse("Lag(A, 2)"), self.series) + self.assertTrue(math.isnan(out[0]) and math.isnan(out[1])) + self.assertEqual(out[2], 1.0) + + def test_rate_of_change(self): + # 常数序列 → 变化率为 0(非 NaN;NaN 仅出现在前 window 步预热) + const = {"C": [5.0] * 6} + out = materialize(parse("RateOfChange(C)"), const) + self.assertTrue(math.isnan(out[0])) # 预热步 NaN + self.assertEqual(out[1], 0.0) + # 含 0 的序列 → 分母为 0 → NaN + zero_denom = {"Z": [0.0, 1.0, 2.0]} + outz = materialize(parse("RateOfChange(Z)"), zero_denom) + self.assertTrue(math.isnan(outz[1])) + # 线性序列 ROC 步长1 = 1/prev + out2 = materialize(parse("RateOfChange(A)"), self.series) + self.assertAlmostEqual(out2[1], 1.0 / 1.0) + self.assertAlmostEqual(out2[5], 1.0 / 5.0) + + def test_log_negative_nan(self): + data = {"P": [1.0, -2.0, math.e]} + out = materialize(parse("Log(P)"), data) + self.assertAlmostEqual(out[0], 0.0) + self.assertTrue(math.isnan(out[1])) + self.assertAlmostEqual(out[2], 1.0) + + def test_scale(self): + out = materialize(parse("Scale(A, 10)"), self.series) + self.assertEqual(out[0], 10.0) + self.assertEqual(out[9], 100.0) + + def test_clip(self): + out = materialize(parse("Clip(A, 3, 7)"), self.series) + self.assertEqual(out, [3.0, 3.0, 3.0, 4.0, 5.0, 6.0, 7.0, 7.0, 7.0, 7.0]) + + def test_combine(self): + data = {"A": [1.0, 2.0, 3.0], "B": [10.0, 20.0, 30.0]} + self.assertEqual(materialize(parse("Combine(A, B)"), data), [11.0, 22.0, 33.0]) + + def test_missing_input_fails_fast(self): + with self.assertRaises(KeyError): + materialize(parse("EMA(Missing, 3)"), {"A": [1.0, 2.0, 3.0]}) + + def test_unknown_op_fails_fast(self): + with self.assertRaises(ValueError): + materialize(OpCall("NoSuchOp", (TagRef("A"),)), self.series) + + +# --------------------------------------------------------------------------- +# 插件注册 +# --------------------------------------------------------------------------- +class RegisterOperatorTest(unittest.TestCase): + def test_register_then_parse_and_run(self): + def _double(series_map, args): + tag = args[0] + return [x * 2 for x in series_map[tag.name]] + + register_operator( + "Double", + min_arity=1, + max_arity=1, + arg_kinds=(("tag",),), + func=_double, + doc="示例自定义算子:翻倍", + ) + try: + self.assertIn("Double", OPERATORS) + self.assertEqual(validate(parse("Double(A)")), []) + self.assertEqual( + materialize(parse("Double(A)"), {"A": [1.0, 2.0]}), [2.0, 4.0] + ) + finally: + OPERATORS.pop("Double", None) + + def test_register_overrides(self): + register_operator( + "Stub", min_arity=0, max_arity=0, arg_kinds=(), func=lambda s, a: 1, doc="v1" + ) + register_operator( + "Stub", min_arity=0, max_arity=0, arg_kinds=(), func=lambda s, a: 2, doc="v2" + ) + try: + self.assertEqual(OPERATORS["Stub"].doc, "v2") + finally: + OPERATORS.pop("Stub", None) + + +# --------------------------------------------------------------------------- +# 描述 / 往返 +# --------------------------------------------------------------------------- +class DescribeAndSerializeTest(unittest.TestCase): + def test_describe_contains_inputs(self): + d = describe(parse("EMA(CLF-01.TEMP, 5m)")) + self.assertIn("CLF-01.TEMP", d) + self.assertIn("EMA", d) + + def test_to_dict_roundtrip_shape(self): + ast = parse("RateOfChange(炉压)") + d = ast.to_dict() + self.assertEqual(d["kind"], "op") + self.assertEqual(d["name"], "RateOfChange") + self.assertEqual(d["args"][0], {"kind": "tag", "name": "炉压"}) + + def test_window_seconds(self): + w = Window(5.0, "m") + self.assertEqual(w.seconds, 300) + self.assertEqual(Window(2.0, "h").seconds, 7200) + + +if __name__ == "__main__": + unittest.main()