From 917295edafb1d6fee96744a9e35775b2d95d4fb4 Mon Sep 17 00:00:00 2001 From: GitHub Date: Tue, 1 Sep 2026 22:17:57 +0800 Subject: [PATCH] =?UTF-8?q?release:=20v1.27.0=20=E2=80=94=20=E9=80=9A?= =?UTF-8?q?=E8=BE=BE=E4=BF=A1=E5=85=AC=E5=BC=8F=E8=A7=A3=E6=9E=90=E5=99=A8?= =?UTF-8?q?=E4=B8=89=E9=80=9A=E9=81=93=20+=20=E8=BD=AE=E5=8A=A8=E7=BB=84?= =?UTF-8?q?=E5=90=88=E5=BC=95=E6=93=8E=20+=20=E5=9B=9E=E6=B5=8B=E9=A1=B5WF?= =?UTF-8?q?/=E8=AF=84=E4=BC=B0=E5=BC=80=E5=85=B3=20+=20Docker=20=E9=83=A8?= =?UTF-8?q?=E7=BD=B2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 升级计划 P3 + P4(部分)。全量 1252 单测、ruff/mypy strict、前端 vue-tsc+vite build 全绿。 - 通达信公式解析器(formula.py):自建 tokenizer + 递归下降 AST + 30+ 函数白名单求值 (不走 Python eval);命名布尔输出=信号列、数值输出=排名列;除零→NaN、预热期不出信号 - 公式三通道:CLI easy-tdx formula compute|screen|backtest;REST /formula/validate|compute| backtest|screen(run/async);Python API run_formula_backtest(买/卖列自动挑选) - 轮动组合引擎(rotation.py):排名定期换仓(打分只用截至当日数据)、槽位等额、 跌出排名自动补位、日/周/月刷新、槽内止盈止损;momentum_score/formula_score 打分; REST /backtest/rotation/run/async - 回测页附加分析开关(Web UI):勾选后随回测并行跑 WF(逐窗红涨绿跌柱状图+汇总卡, 窗口数 2~12)与一条龙评估(评分分项条/高适配徽标/买入持有对比/8 项适配检查); 新增 WalkForwardPanel/EvaluatePanel 组件与 store runWalkforward/runEvaluate; WF 端点 ?n_windows= 透传;修复报告 numpy 标量 REST 400(源头清洗) - Docker 部署(Dockerfile + docker-compose.yml,/data 卷 + 健康检查)与 scripts/verify_ci.sh 一键门禁 - 升级计划文档 docs/upgrade-plan-2026H2.md(四阶段全部完成 + 诚实实测数据) - 未做(独立排期):Playwright E2E、WebSocket 实时联动、引擎逐 bar 向量化 --- CHANGELOG.md | 14 + Dockerfile | 16 + docker-compose.yml | 35 ++ docs/upgrade-plan-2026H2.md | 135 ++++++ pyproject.toml | 2 +- scripts/verify_ci.sh | 43 ++ src/easy_tdx/backtest/__init__.py | 23 + src/easy_tdx/backtest/formula_strategy.py | 187 ++++++++ src/easy_tdx/backtest/rotation.py | 429 +++++++++++++++++ src/easy_tdx/cli/__init__.py | 2 + src/easy_tdx/cli/cmd_formula.py | 208 +++++++++ src/easy_tdx/formula.py | 519 +++++++++++++++++++++ src/easy_tdx/web/app.py | 2 + src/easy_tdx/web/routers/formula.py | 241 ++++++++++ tests/unit/test_backtest_rotation.py | 159 +++++++ tests/unit/test_formula.py | 214 +++++++++ tests/unit/test_formula_integration.py | 193 ++++++++ web-ui/src/api.ts | 27 ++ web-ui/src/components/EvaluatePanel.vue | 255 ++++++++++ web-ui/src/components/WalkForwardPanel.vue | 168 +++++++ web-ui/src/stores/backtest.ts | 76 +++ web-ui/src/types.ts | 93 ++++ web-ui/src/views/BacktestView.vue | 99 +++- 23 files changed, 3135 insertions(+), 5 deletions(-) create mode 100644 Dockerfile create mode 100644 docker-compose.yml create mode 100644 docs/upgrade-plan-2026H2.md create mode 100644 scripts/verify_ci.sh create mode 100644 src/easy_tdx/backtest/formula_strategy.py create mode 100644 src/easy_tdx/backtest/rotation.py create mode 100644 src/easy_tdx/cli/cmd_formula.py create mode 100644 src/easy_tdx/formula.py create mode 100644 src/easy_tdx/web/routers/formula.py create mode 100644 tests/unit/test_backtest_rotation.py create mode 100644 tests/unit/test_formula.py create mode 100644 tests/unit/test_formula_integration.py create mode 100644 web-ui/src/components/EvaluatePanel.vue create mode 100644 web-ui/src/components/WalkForwardPanel.vue diff --git a/CHANGELOG.md b/CHANGELOG.md index d131084..a819711 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,20 @@ 本文件记录 easy-tdx 的版本变更。格式遵循 [Keep a Changelog](https://keepachangelog.com/zh-CN/)。 +## [1.27.0] — 2026-09-01 + +**公式与轮动版本**——升级计划 P3 + P4(部分)落地:通达信公式解析器让写惯公式的用户零 Python 进入筛选/回测,轮动组合引擎补齐「排名换仓」组合形态,附 Docker 部署与一键门禁脚本。 + +### 新增 + +- **通达信公式解析器**(`formula.py`)——自建 tokenizer + 递归下降 AST + 白名单求值(**不走 Python eval**,无注入面):支持 `:=` 中间变量 / `名称:` 命名输出、`+ - * /`(除零→NaN)、比较、`AND OR NOT`(兼容 `&& || !`)、花括号注释、中文标识符;序列别名 C/O/H/L/V/AMOUNT;函数白名单 30+(MA/EMA/SMA/HHV/LLV/REF/CROSS/LONGCROSS/IF/MACD/KDJ/RSI/BOLL/ATR…,全部后视函数,**无未来数据**);命名布尔输出自动归类为**信号列**、数值输出归类为**排名列**;未知函数/变量报带位置的 `FormulaError`。 +- **公式回测适配器**(`backtest/formula_strategy.py`)——信号列注入 K 线 + `ColumnSignalStrategy` 逐 bar 交易;买/卖列自动挑选(「买/卖」与 BUY/SELL 名称提示优先,其次声明顺序);信号下一根开盘成交;结果附 S-D 评级与综合评分。 +- **公式三通道**——CLI `easy-tdx formula compute|screen|backtest`(`--formula` 或 `--file`,screen 支持逗号分隔/@文件标的列表);REST `POST /formula/validate`(语法+归类校验,无需数据)、`/formula/compute`(内联 ohlcv 或 symbol)、`/formula/backtest/run/async`、`/formula/screen/run/async`(后台任务);Python API `run_formula_backtest()`。 +- **轮动组合引擎**(`backtest/rotation.py`)——排名定期换仓:打分函数只喂截至当日收盘的前缀数据(无未来泄漏);固定槽位**等额**(预算 = 净值/槽数,杜绝首买全仓单票);跌出前 `keep_rank` 名自动卖出、空槽自动补位;`daily/weekly/monthly` 刷新;可选槽内止盈止损(收盘触发、次开成交);复用主引擎 19 项绩效 + 组合评级。内置 `momentum_score(period)` 与 `formula_score(公式)` 打分(与公式模块联动)。REST `POST /backtest/rotation/run/async`。 +- **回测页附加分析开关(Web UI)**——回测页新增「附加分析」区:勾选「Walk-Forward 样本外验证」随回测自动附加 WF 任务(窗口数可调 2~12,逐窗收益红涨绿跌柱状图 + 盈利窗占比/连乘收益/最差窗汇总卡);勾选「一条龙评估」附加评估任务(综合评分 0-100 分项条 + 高适配徽标 + 买入持有基准对比与「跑输买入持有」警示 + 8 项适配性检查清单 + 评级复用本地口径)。两任务与主回测共用同一份内联行情、并行互不阻塞、独立错误提示;新增 `WalkForwardPanel.vue` / `EvaluatePanel.vue` 组件与 store 的 `runWalkforward`/`runEvaluate` action(统一 `pollTask` 轮询助手);WF 端点支持 `?n_windows=` 查询参数。附带修复:WF/fitness/evaluate 报告的 numpy 标量在 REST 序列化时 400 的问题(`types.to_json_native` 源头清洗,各结果 `to_dict` 统一接入)。 +- **Docker 部署**(`Dockerfile` + `docker-compose.yml`)——python:3.12-slim,装 `[web,warehouse]` 可选依赖,`/data` 卷持久化自选/策略库/任务库/K 线仓库,带健康检查。 +- **一键门禁脚本**(`scripts/verify_ci.sh`)——ruff + ruff format + mypy strict + 全量 pytest 一条命令(`--fast` 跳过测试),可挂 git pre-push hook。 + ## [1.26.0] — 2026-09-01 **本地数据仓库版本**——把碎片化缓存升级为统一数据底座(升级计划 P2 阶段;P2-2 评级后端化已随 1.25.0 提前交付)。此前下游项目(indicator-lab 的 DuckDB 仓库、backtest-system 的 cache/ 目录)都在自建数据层,现在 easy-tdx 原生提供。 diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..e17563c --- /dev/null +++ b/Dockerfile @@ -0,0 +1,16 @@ +# easy-tdx 容器镜像(配合 docker-compose.yml 使用) +FROM python:3.12-slim + +WORKDIR /app + +# 先装依赖层(利用构建缓存) +COPY pyproject.toml README.md ./ +RUN pip install --no-cache-dir --upgrade pip \ + && pip install --no-cache-dir .[web,warehouse] + +# 装本地源码(开发机构建;发布镜像可直接从 PyPI 装 easy-tdx) +COPY . . +RUN pip install --no-cache-dir --no-deps . + +EXPOSE 8000 +CMD ["easy-tdx", "serve", "--host", "0.0.0.0", "--port", "8000", "--no-open-browser"] diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..1ebd71d --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,35 @@ +# easy-tdx 一键部署:API + Web UI(Docker Compose,v1.27 新增) +# +# 用法: +# docker compose up -d # 构建并启动(http://localhost:8000) +# docker compose logs -f easy-tdx +# docker compose down +# +# 说明: +# - 镜像内安装 [web,warehouse] 可选依赖(FastAPI + DuckDB); +# - 数据目录(自选/策略库/任务库/K线仓库)挂载到宿主机 ./data,重建容器不丢; +# - serve 默认 0.0.0.0:8000,自动托管 Web UI;纯 API 用法把 command 换成 +# ["easy-tdx", "serve", "--host", "0.0.0.0", "--no-ui"]。 + +services: + easy-tdx: + image: easy-tdx:latest + build: + context: . + dockerfile: Dockerfile + container_name: easy-tdx + command: ["easy-tdx", "serve", "--host", "0.0.0.0", "--port", "8000", "--no-open-browser"] + ports: + - "8000:8000" + environment: + EASY_TDX_CONFIG_DIR: /data + TZ: Asia/Shanghai + volumes: + - ./data:/data + restart: unless-stopped + healthcheck: + test: ["CMD", "python", "-c", "import urllib.request;urllib.request.urlopen('http://localhost:8000/openapi.json', timeout=5)"] + interval: 30s + timeout: 10s + retries: 3 + start_period: 15s diff --git a/docs/upgrade-plan-2026H2.md b/docs/upgrade-plan-2026H2.md new file mode 100644 index 0000000..b284e58 --- /dev/null +++ b/docs/upgrade-plan-2026H2.md @@ -0,0 +1,135 @@ +# easy-tdx 升级开发计划(2026 H2) + +> **执行进度**(2026-09-01 最终更新:全部阶段完成): +> - ✅ **P0 v1.24.0**——QFQ 对拍验证体系、回测任务 SQLite 持久化 + 导出、品种感知费率 +> - ✅ **P1 v1.25.0**——Walk-Forward 引擎、适配性评估、一条龙评估、综合评分、评级后端化、多 seed 验证 + 晋级门槛、寻优两段式加速(指标缓存 + 进程并行) +> - ✅ **P2 v1.26.0**——DuckDB K 线仓库 + provisional 状态机 + 增量同步 + 健康自检 + CLI `warehouse` 命令组(P2-2 评级后端化已随 1.25.0 提前交付) +> - ✅ **P3 v1.27.0**——通达信公式解析器(tokenizer + AST + 白名单求值,无 Python eval)+ 公式三通道(CLI `formula` / REST `/formula/*` / Python API)+ 轮动组合引擎(排名换仓 + 槽位等额 + 自动补位 + 止盈止损,REST `/backtest/rotation/run/async`) +> - ✅ **P1 前端补全(随 1.27.0)**——回测页「附加分析」开关:WF 逐窗柱状图 + 一条龙评估报告卡(评分分项/高适配徽标/买入持有基准对比/8 项适配检查),复用同一份内联行情并行执行 +> - ✅ **P4(部分)v1.27.0**——Docker Compose 部署 + `scripts/verify_ci.sh` 一键门禁。**未做**:Playwright E2E(需前端基建,独立排期)、WebSocket 实时推送联动 EventBus(README 既有 TODO,与本计划无耦合)、引擎逐 bar 循环向量化(寻优加速的下一步方向)——三项已整理为独立排期提示词 +> +> 实测备注(诚实数据): +> - 寻优加速——指标缓存命中率 41.7%(36 点网格)但墙钟 ~1.01x(本引擎指标层非瓶颈,逐 bar Python 循环才是);进程并行 4 workers 约 2x。后续更大加速的方向是引擎循环向量化。 +> - 最终回归:1252 个单元测试全过(基线 1078),ruff / ruff format / mypy strict(245 文件)全绿。 +> +> 依据:对两个基于 easy-tdx 的下游项目的逆向调研 +> - [mvpbaggio/backtest-system](https://github.com/mvpbaggio/backtest-system)(v1.4,MIT)— 回测框架,运行时依赖 easy-tdx(行情拉取、PerformanceAnalyzer、MyTT、Param 注册表思路) +> - [kendev93/indicator-lab](https://github.com/kendev93/indicator-lab)(AGPL-3.0)— 指标实验室,参考改写了 easy-tdx 的 TDX 协议层与 `.day` 格式(见其 THIRD_PARTY_NOTICES.md),运行时不依赖 +> 调研日期:2026-09-01,基线版本 v1.23.3 + +--- + +## 一、核心洞察 + +两个下游项目**独立地**在补 easy-tdx 的同一类空白,这比任何单个功能都更有信号价值: + +1. **防过拟合验证是最大空白**。backtest-system 自研了严格 7 窗 Walk-Forward;indicator-lab 自研了「策略适配性评估」(train/val/test 三段 + 8 项检查)。两条路殊途同归——easy-tdx 的回测引擎很全(滑点/执行仿真/归因/寻优/组合),但**没有任何样本外验证工具**,下游只能自己造。 +2. **统一 K 线落盘是第二空白**。easy-tdx 的缓存是碎片化的(股票列表缓存、best_host、进程内 XDXR 字典、扫描 JSON 增量缓存),没有统一磁盘 K 线层;indicator-lab 为此建了 DuckDB 仓库,backtest-system 为此自建 cache/ 目录 + 7 天更新 + 数据自检。 +3. **QFQ 质量收到了直接差评**。backtest-system README 明确弃用 easy-tdx 的 QFQ,理由是「茅台出现负价、浦发除权方向算反」,随后自研了板块感知阈值的跳空检测前复权。本地兜底 `mac/adjust.py` 已存在,但缺少对拍验证体系,无法自证可靠。 +4. **降低使用门槛有巨大空间**。indicator-lab 的通达信公式解析器(粘贴公式自动识别参数/命名信号)直接命中中国最大的量化用户群体——写惯通达信公式的股民,而 easy-tdx 目前要求写 Python。 + +--- + +## 二、借鉴点清单(按价值排序) + +| # | 借鉴点 | 来源 | easy-tdx 现状 | 价值 | +|---|--------|------|--------------|------| +| 1 | Walk-Forward 样本外验证(7 窗、每窗独立开仓防跨窗重复计收益) | backtest-system `walkforward.py` | ❌ 无 | ★★★★★ | +| 2 | 策略适配性评估(60/20/20 三段独立回测 + 8 个可解释检查项 + 滚动适配过滤防未来泄漏) | indicator-lab strategy-fitness | ❌ 无 | ★★★★★ | +| 3 | 综合评分 + 一条龙评估(score:收益50/夏普15/回撤10/Sortino5/WF20;evaluate:拉数/对齐/选模式/对比基准引擎) | backtest-system `benchmark.py` | ❌ 无 | ★★★★☆ | +| 4 | 通达信公式解析器(`名称:=数值` 参数识别、命名布尔输出作信号、命名数值用于排序/卖出) | indicator-lab | ❌ 无(34 指标需改源码新增) | ★★★★☆ | +| 5 | 本地数据仓库(增量导入、源只读、在线补缺不覆盖、provisional/completed 状态机) | indicator-lab(DuckDB) | ❌ 无统一 K 线磁盘缓存 | ★★★★☆ | +| 6 | 动态组合轮动回测(按指标排序 + 固定槽位等额 + 卖出自动补位 + 日/周/月刷新 + 槽内止盈止损) | indicator-lab portfolio-backtest | ⚠️ 有组合回测/再平衡,但无「排名轮动」模式 | ★★★★☆ | +| 7 | 两段式引擎协议(指标计算缓存 与 信号组合 解耦,迭代快 10 倍) | backtest-system `register_two_stage` | ❌ 优化器每组参数全量重算 | ★★★☆☆ | +| 8 | 多 seed 验证 + 晋级门槛(正收益比例/夏普/WF/交易数四门槛) | backtest-system `engine_iter.py` | ❌ 无 | ★★★☆☆ | +| 9 | QFQ 互检:NONE 原始价 + 向下跳空检测(主板10%/双创20%/北交所30% 阈值) | backtest-system `data_source.py` | ⚠️ 有 XDXR 公式法本地兜底,无对拍 | ★★★☆☆(质量修复) | +| 10 | 真实出场模型:吊灯 ATR14×3 + 保本 BE + 移动止盈 TP、跳空按更差开盘成交 | backtest-system | ⚠️ 有订单/滑点/执行仿真,出场模式较简单 | ★★★☆☆ | +| 11 | 品种感知费率(股票 vs ETF/B 股:最低佣金、印花税差异) | indicator-lab | ⚠️ 有费率参数,非品种感知 | ★★☆☆☆ | +| 12 | 任务体验:进度、失败/跳过摘要、JSON/CSV 导出 | indicator-lab | ❌ 任务仅内存、重启即清、无导出 | ★★★☆☆ | +| 13 | Playwright E2E(mock API、不依赖真实数据)+ verify_ci.sh + git hooks | indicator-lab | ❌ 前端无 E2E | ★★☆☆☆ | +| 14 | Docker Compose 部署 | indicator-lab | ❌ 无 | ★★☆☆☆ | +| 15 | 后端数据评级(S-D 五档六维加权,现仅存 web-ui/src/grading 前端 TS) | (自补,非下游首创) | ⚠️ 前端有、Python/CLI/REST 无 | ★★★☆☆ | + +**不借鉴**:两个项目的回测执行内核(easy-tdx 引擎更全:TWAP/VWAP、Brinson 归因、DSL、缠论桥接);indicator-lab 的 `at_least` 条件组合(`combo.py` 已有 MAJORITY);图表截 220 根方案(前端已有自己方案)。 + +--- + +## 三、升级开发计划 + +### P0 — 信任与持久化(v1.24.0,约 1~2 周) + +**目标:先修质量口碑,再谈新功能。** + +| 任务 | 内容 | 验收标准 | +|------|------|---------| +| P0-1 QFQ 对拍验证体系 | ① 用 `adjust.py` 公式法与 backtest-system 跳空检测法做双引擎互检,不一致即告警;② 建立「已知除权案例」回归集(茅台/浦发等重度除权股);③ `has_bad_prices` 从兜底升级为所有 QFQ 出口(CLI/Web/unified)的强制门禁,失败自动降级本地重算并标记 | 案例集全过;任何出口不再可能出现负价/方向反转;新增 `tests/test_qfq_crosscheck.py` | +| P0-2 回测任务持久化 | SQLite `~/.easy_tdx/tasks.db`(复用 watchlist.db/strategies.db 模式),任务状态/结果落盘,serve 重启不丢;REST 增加 JSON/CSV 导出端点 | 重启 serve 后 /compare 仍能看历史任务;可下载结果文件 | +| P0-3 真实平均持仓天数 | 去掉 `performance.py` 中 `avg_holding_days = 5.0` 的固定值,从成交记录真实统计 | 单测覆盖多笔开平仓场景 | +| P0-4 品种感知费率 | 费率模型按品种区分(股票:佣金万2.5~万3 最低5元+印花税卖出千1;ETF:佣金更低、免印花税;B 股单独口径) | 回测引擎按 symbol 自动套用,可覆盖 | + +### P1 — 防过拟合验证链(v1.25.0,约 2~3 周)⭐ 主打版本 + +**目标:补上两个下游都在自己造的最大空白,让「回测好」升级为「样本外也好」。** + +| 任务 | 内容 | 验收标准 | +|------|------|---------| +| P1-1 Walk-Forward 引擎 | 新增 `backtest/walkforward.py`:后 70% 切 N 窗(默认 7)严格样本外,**每窗独立开仓**(防跨窗重复计收益,backtest-system v1.2.1 的教训);输出逐窗收益曲线 + 窗间稳定性指标;接入 CLI `easy-tdx backtest ... --wf` 与 REST `/backtest/wf/run/async` | 与 backtest-system 对齐的窗独立语义;Web 前端展示逐窗柱状图 | +| P1-2 策略综合评分 | `backtest/scoring.py`:score_strategy() 0-100 加权(收益 50 / 夏普 15 / 回撤 10 / Sortino 5 / WF 稳定性 20),与前端现有 S-D 评级打通(评级后端化一并完成,见 P2-2 可提前) | CLI/Web 均输出评分与分项 | +| P1-3 适配性评估 | `backtest/fitness.py`:train/valid/test(默认 60/20/20)三段独立回测 + 可解释检查项(收益一致性、回撤一致性、交易数充分性、胜率区间、参数敏感性等 8 项),≥75% 通过且样本达标 → 「高适配」标记;支持**滚动适配过滤**(仅用早于当天的已平仓数据,杜绝未来数据泄漏) | Web 策略库/对比页显示适配徽章;可解释报告 | +| P1-4 一条龙评估 | `evaluate_strategy()`:默认随机抽样股票池(固定 seed 可复现)→ 对齐 → 自动选出场模式 → 与基准引擎(买入持有 + MyTT MACD)同规则对比 | 一条命令出完整对比报告 | +| P1-5 两段式寻优加速 | 优化器支持指标缓存复用:参数只影响信号组合层时,指标层计算一次(backtest-system 实测快 10 倍) | 网格寻优基准测试提速 ≥3 倍 | +| P1-6 多 seed 验证 + 晋级门槛 | `run-all` / 优化器输出增加多随机种子组合验证;晋级门槛四项(正收益比例/夏普/WF/交易数)可配置 | 报告含跨 seed 稳定性列 | + +### P2 — 本地数据仓库(v1.26.0,约 2~3 周) + +**目标:把碎片化缓存升级为统一数据底座,服务全市场扫描/因子/回测的提速。** + +| 任务 | 内容 | 验收标准 | +|------|------|---------| +| P2-1 K 线仓库 | 新增 `warehouse/` 模块(存储引擎选 DuckDB,零服务、列存、SQL 友好):`easy-tdx warehouse sync`(全量/增量)、源只读、在线补缺不覆盖、按品种价格/成交量系数;**provisional 状态机**(15:05 前的今日数据标记临时,筛选/回测默认忽略,收盘后转 completed) | 二次 sync 增量;断网可用;screen/factor/pfactor 可切 `--source warehouse` | +| P2-2 评级后端化 | `web-ui/src/grading/`(engine.ts/thresholds.ts)移植为 Python `backtest/grading.py`,CLI `--grade`、REST 返回 grade 字段;前端改为消费后端结果(保留前端兜底) | API/CLI 输出 S-D;与前端旧实现结果一致率 100%(对拍单测) | +| P2-3 仓库健康自检 | 数据自检命令:缺口检测、异常跳变检测(复用 P0-1 跳空检测)、最新度报告 | `easy-tdx warehouse check` | + +### P3 — 公式与轮动(v1.27.0,约 3~4 周) + +**目标:把「写通达信公式」的庞大用户群接进来。** + +| 任务 | 内容 | 验收标准 | +|------|------|---------| +| P3-1 通达信公式解析器 | `indicator/formula.py`:解析通达信/麦语言公式,自动识别 `名称:=数值` 参数、命名布尔输出→信号、命名数值→排序/排序卖出字段;内置函数映射到 MyTT;安全除零、无未来数据、数据不足跳过 | 一批典型公式(含用户常见主力/洗盘类指标)解析通过并可直接回测 | +| P3-2 公式三通道接入 | CLI `easy-tdx formula screen/backtest`、REST `/formula/compute`、Web 新页面(粘贴公式 → 选股/回测一体) | 三通道行为一致 | +| P3-3 轮动组合引擎 | `backtest/rotation.py`:按指标排序选股 + 固定槽位等额 + 卖出自动补位 + 日/周/月刷新 + 槽内止盈止损/指标阈值卖出/指标比较卖出 | CLI/REST/Web 均可跑;与单标的回测同一套绩效输出 | + +### P4 — 工程化(滚动进行) + +| 任务 | 内容 | +|------|------| +| P4-1 Playwright E2E | mock API 的前端 E2E(不依赖真实行情/网络),纳入 CI | +| P4-2 WebSocket 实时联动 | `/ws/realtime/{symbol}` 接通 `realtime/EventBus`(README 已自认未联动;两个下游都没碰实时,这是 easy-tdx 的独有优势区,应当做实) | +| P4-3 Docker Compose | 一键起 serve + Web UI 的部署方案 | +| P4-4 verify_ci 风格脚本 | 一条命令跑完 ruff/mypy/pytest/前端 typecheck+build/E2E,可安装 git hooks | + +--- + +## 四、版本节奏与依赖关系 + +``` +v1.24.0 (P0) ──→ v1.25.0 (P1 防过拟合链,依赖 P0-4 品种费率) + └→ v1.26.0 (P2 数据仓库,与 P1 可并行启动,P2-2 建议提前到 P1 一起做) + └→ v1.27.0 (P3 公式+轮动,依赖 P2-1 仓库提速全市场公式选股) +P4 工程化滚动穿插。 +``` + +**成功指标**(对外可宣传): +- v1.25 后:官方提供 WF + 适配性双验证,下游不再需要自造防过拟合轮子; +- v1.26 后:下游不再需要自建数据层(backtest-system 的 cache/、indicator-lab 的 DuckDB 均可换成 easy-tdx 仓库); +- v1.27 后:通达信公式用户零代码进入回测。 + +--- + +## 五、风险与注意 + +1. **P3 公式解析器工作量大**:通达信公式方言庞杂(函数集、隐式循环语义),建议首版只支持「日期序列 + 常用函数白名单」,明确不支持清单,渐进扩充。indicator-lab 的实现可作参考(注意其 AGPL 许可——**只看思路不抄代码**,避免传染)。 +2. **DuckDB 引入新增运行时依赖**:当前核心依赖仅 3 个是卖点。建议放入 optional-dependencies `[warehouse]` 组,import 惰性加载。 +3. **WF 每窗独立开仓**是 backtest-system 踩过的坑(v1.2.1 修复),实现时直接采用正确语义,勿重蹈覆辙。 +4. **向后兼容**:新增能力全部走可选参数/可选依赖,默认行为不变;v1.24 的 QFQ 门禁若触发降级,需在输出中显式标记(grade 字段),不静默。 diff --git a/pyproject.toml b/pyproject.toml index d2d754d..9dc6232 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "hatchling.build" [project] name = "easy-tdx" -version = "1.26.0" +version = "1.27.0" description = "通达信 TCP 协议行情数据客户端,支持在线行情、离线数据读取与写入同步" readme = "README.md" requires-python = ">=3.10" diff --git a/scripts/verify_ci.sh b/scripts/verify_ci.sh new file mode 100644 index 0000000..7367812 --- /dev/null +++ b/scripts/verify_ci.sh @@ -0,0 +1,43 @@ +#!/usr/bin/env bash +# easy-tdx 一键本地门禁(等价 CI 的质量检查,v1.27 新增)。 +# +# 用法:bash scripts/verify_ci.sh [--fast] +# --fast 跳过全量测试(只跑 ruff + mypy + 格式检查) +# +# 可选安装为 git hook(pre-push): +# ln -s ../../scripts/verify_ci.sh .git/hooks/pre-push + +set -euo pipefail +cd "$(dirname "$0")/.." + +PY="${PYTHON:-.venv/Scripts/python.exe}" +if [ ! -f "$PY" ]; then + PY="${PYTHON:-.venv/bin/python}" +fi +if [ ! -f "$PY" ]; then + PY="python" +fi + +FAST=0 +[ "${1:-}" = "--fast" ] && FAST=1 + +echo "── ruff check ──────────────────────────────────────────" +"$PY" -m ruff check src/ tests/ + +echo "── ruff format --check ─────────────────────────────────" +"$PY" -m ruff format --check src/ tests/ + +echo "── mypy --strict ───────────────────────────────────────" +"$PY" -m mypy src/easy_tdx/ + +if [ "$FAST" = "1" ]; then + echo "── 跳过测试(--fast)───────────────────────────────────" + echo "✓ verify_ci (fast) 全部通过" + exit 0 +fi + +echo "── pytest(全量单元测试)────────────────────────────────" +"$PY" -m pytest tests/ -q --ignore=tests/integration + +echo "" +echo "✓ verify_ci 全部通过" diff --git a/src/easy_tdx/backtest/__init__.py b/src/easy_tdx/backtest/__init__.py index c8df2da..24f79b1 100644 --- a/src/easy_tdx/backtest/__init__.py +++ b/src/easy_tdx/backtest/__init__.py @@ -20,10 +20,20 @@ print(result.performance) """ +from easy_tdx.backtest.benchmark import evaluate_strategy, run_buy_hold_benchmark # noqa: F401 from easy_tdx.backtest.combo import CombinationRunner, ComboResult, FactorSignals # noqa: F401 from easy_tdx.backtest.engine import BacktestEngine # noqa: F401 +from easy_tdx.backtest.fitness import FitnessEngine, FitnessReport # noqa: F401 +from easy_tdx.backtest.formula_strategy import run_formula_backtest # noqa: F401 +from easy_tdx.backtest.grading import GradeResult, grade_performance # noqa: F401 +from easy_tdx.backtest.rotation import RotationEngine, RotationResult # noqa: F401 +from easy_tdx.backtest.scoring import StrategyScore, score_strategy # noqa: F401 from easy_tdx.backtest.strategy import Strategy, StrategyDataProxy, crossover # noqa: F401 from easy_tdx.backtest.types import BacktestResult, Position, Signal, Trade # noqa: F401 +from easy_tdx.backtest.walkforward import ( # noqa: F401 + WalkForwardEngine, + WalkForwardResult, +) __all__ = [ "BacktestEngine", @@ -31,10 +41,23 @@ __all__ = [ "CombinationRunner", "ComboResult", "FactorSignals", + "FitnessEngine", + "FitnessReport", + "GradeResult", "Strategy", "StrategyDataProxy", + "StrategyScore", "Signal", "Trade", "Position", + "WalkForwardEngine", + "WalkForwardResult", "crossover", + "evaluate_strategy", + "grade_performance", + "RotationEngine", + "RotationResult", + "run_buy_hold_benchmark", + "run_formula_backtest", + "score_strategy", ] diff --git a/src/easy_tdx/backtest/formula_strategy.py b/src/easy_tdx/backtest/formula_strategy.py new file mode 100644 index 0000000..e757f74 --- /dev/null +++ b/src/easy_tdx/backtest/formula_strategy.py @@ -0,0 +1,187 @@ +"""通达信公式 → 回测引擎适配器(v1.27 新增)。 + +把 :mod:`easy_tdx.formula` 的信号列注入 K 线,再用轻量策略在逐 bar 循环 +里读取信号——公式用户零 Python 即可回测(三通道口径一致:CLI +``easy-tdx formula backtest``、REST ``/formula/backtest/run/async``、 +Python API :func:`run_formula_backtest`)。 + +信号约定(与 indicator-lab 一致的语义): + +- 公式的**命名布尔输出**即信号列。默认买入列 = 第一个信号列(或名字含 + 「买」/``B`` 的信号列),默认卖出列 = 第二个信号列(或名字含「卖」/ + ``S`` 的信号列),也可显式指定; +- 信号在**下一根 K 线开盘**成交(与引擎 ``next_open`` 默认一致,无未来 + 数据);预热期 NaN 视为无信号。 +""" + +from __future__ import annotations + +import re +from typing import Any + +import pandas as pd + +from easy_tdx.backtest.engine import BacktestEngine +from easy_tdx.backtest.strategy import Strategy +from easy_tdx.formula import CompiledFormula, FormulaResult, compile_formula + +__all__ = [ + "ColumnSignalStrategy", + "FormulaStrategyError", + "attach_formula_columns", + "run_formula_backtest", +] + +# 名称提示只认「买/卖」与多字母 BUY/SELL——单字母 A/B/S 是常见中间变量名, +# 按词边界匹配会误判,故不做单字母提示 +_BUY_HINT = re.compile(r"买|buy", re.IGNORECASE) +_SELL_HINT = re.compile(r"卖|sell", re.IGNORECASE) + + +class FormulaStrategyError(ValueError): + """公式回测配置错误(无可用信号列等)。""" + + +def attach_formula_columns( + df: pd.DataFrame, + compiled: CompiledFormula, +) -> tuple[pd.DataFrame, FormulaResult]: + """把公式的全部输出列注入 df(副本),返回 (新 df, 公式结果)。""" + result = compiled.compute(df) + if not result.columns: + raise FormulaStrategyError("公式没有命名输出(用 `名称: 表达式;` 声明输出)") + out = df.copy() + for name, arr in result.columns.items(): + out[name] = arr + return out, result + + +class ColumnSignalStrategy(Strategy): + """按已注入的信号列交易:买入列=1 全仓买,卖出列=1 全仓卖。""" + + def __init__(self, buy_col: str, sell_col: str | None = None) -> None: + super().__init__() + self._buy_col = buy_col + self._sell_col = sell_col + self._holding = False + + def init(self) -> None: + self._holding = False + + def next(self) -> None: + buy_v = getattr(self.data, self._buy_col)[0] + buy_on = buy_v == buy_v and buy_v >= 1.0 # NaN 安全 + if buy_on and not self._holding: + self.buy() + self._holding = True + return + if self._holding and self._sell_col is not None: + sell_v = getattr(self.data, self._sell_col)[0] + if sell_v == sell_v and sell_v >= 1.0: + self.sell() + self._holding = False + + +def pick_signal_columns( + result: FormulaResult, + buy_col: str | None = None, + sell_col: str | None = None, +) -> tuple[str, str | None]: + """解析买/卖信号列:显式指定优先,否则按声明顺序 + 名称提示自动挑选。""" + signals = result.signals + if not signals: + raise FormulaStrategyError( + f"公式没有布尔信号输出(现有数值输出: {result.values});" + "信号需为比较/逻辑表达式,如 `买入: CROSS(MA(C,5), MA(C,20));`" + ) + if buy_col is not None: + if buy_col not in result.columns: + raise FormulaStrategyError(f"指定的买入列 {buy_col!r} 不在公式输出中") + else: + hinted = [s for s in signals if _BUY_HINT.search(s)] + buy_col = hinted[0] if hinted else signals[0] + if sell_col is not None: + if sell_col not in result.columns: + raise FormulaStrategyError(f"指定的卖出列 {sell_col!r} 不在公式输出中") + else: + hinted = [s for s in signals if _SELL_HINT.search(s) and s != buy_col] + rest = [s for s in signals if s != buy_col] + sell_col = hinted[0] if hinted else (rest[0] if rest else None) + return buy_col, sell_col + + +def _clean_json(obj: Any) -> Any: + """递归清洗 numpy 标量/Timestamp/NaN → JSON 原生(REST 任务可序列化)。""" + import numpy as _np + + if isinstance(obj, dict): + return {str(k): _clean_json(v) for k, v in obj.items()} + if isinstance(obj, (list, tuple)): + return [_clean_json(v) for v in obj] + if isinstance(obj, _np.integer): + return int(obj) + if isinstance(obj, _np.floating): + f = float(obj) + return f if _np.isfinite(f) else None + if isinstance(obj, _np.bool_): + return bool(obj) + if obj is None or isinstance(obj, (str, int, bool)): + return obj + if isinstance(obj, float): + return float(obj) if _np.isfinite(obj) else None + if hasattr(obj, "isoformat"): + return obj.isoformat() + return str(obj) + + +def run_formula_backtest( + df: pd.DataFrame, + formula_text: str | CompiledFormula, + buy_col: str | None = None, + sell_col: str | None = None, + cash: float = 100000.0, + commission: float = 0.0003, + min_commission: float = 5.0, + stamp_tax: float = 0.001, + slippage: float = 0.0, + execution: str = "next_open", + symbol: str | None = None, + auto_fees: bool = False, +) -> dict[str, Any]: + """公式一条龙回测:注入信号列 → 挑买/卖列 → 引擎回测 → 附公式元信息。 + + Returns: + ``{"performance", "trades", "equity_curve", "config", "grade", "score", + "formula": {"signals", "values", "buy_col", "sell_col"}}`` + """ + from easy_tdx.backtest.grading import grade_performance + from easy_tdx.backtest.scoring import score_strategy + + compiled = ( + formula_text if isinstance(formula_text, CompiledFormula) else compile_formula(formula_text) + ) + enriched, result = attach_formula_columns(df, compiled) + b_col, s_col = pick_signal_columns(result, buy_col, sell_col) + + engine = BacktestEngine( + strategy=ColumnSignalStrategy(buy_col=b_col, sell_col=s_col), + cash=cash, + commission=commission, + min_commission=min_commission, + stamp_tax=stamp_tax, + slippage=slippage, + execution=execution, + symbol=symbol, + auto_fees=auto_fees, + ) + bt = engine.run(enriched) + out: dict[str, Any] = dict(_clean_json(bt.to_dict())) + out["grade"] = grade_performance(dict(bt.performance)).to_dict() + out["score"] = score_strategy(dict(bt.performance)).to_dict() + out["formula"] = { + "signals": result.signals, + "values": result.values, + "buy_col": b_col, + "sell_col": s_col, + } + return out diff --git a/src/easy_tdx/backtest/rotation.py b/src/easy_tdx/backtest/rotation.py new file mode 100644 index 0000000..b9f0716 --- /dev/null +++ b/src/easy_tdx/backtest/rotation.py @@ -0,0 +1,429 @@ +"""轮动组合回测引擎(v1.27 新增)。 + +按排名定期换仓的组合策略回测(借鉴 indicator-lab 的动态组合语义): + +- **排名**:每个调仓日用 ``score_fn`` 对股票池逐标的打分(只用截至当日 + 收盘的数据,无未来泄漏),分数可来自动量、因子或**通达信公式的数值 + 输出**(:func:`formula_score`); +- **固定槽位等额**:资金分为 ``slots`` 个槽位,每槽 = 当前净值 / 槽数; +- **卖出自动补位**:持仓跌出前 ``keep_rank`` 名(默认 = 槽数,可加缓冲 + 池降低换手)→ 次日开盘卖出,空出的槽位买入新的前排名(次日开盘); +- **刷新频率**:``daily`` / ``weekly``(每周首个交易日)/ ``monthly`` + (每月首个交易日); +- **槽内止盈止损**:收盘价较成本跌破 ``stop_loss`` 或涨破 ``take_profit`` + → 次日开盘卖出(不做盘中路径假设,全部次开成交,口径与主引擎一致)。 + +执行语义:调仓信号在 T 日收盘产生、T+1 开盘成交(next_open),复用 +主引擎的绩效分析器(19 项指标)与组合评级。 +""" + +from __future__ import annotations + +import math +from collections.abc import Callable +from dataclasses import dataclass, field +from typing import Any + +import numpy as np +import pandas as pd + +from easy_tdx.backtest.performance import PerformanceAnalyzer + +__all__ = ["RotationEngine", "RotationResult", "momentum_score", "formula_score"] + +ScoreFn = Callable[[pd.DataFrame], float] + + +def momentum_score(period: int = 20) -> ScoreFn: + """动量打分:最近 ``period`` 根涨幅(越高排名越靠前)。""" + + def _score(df: pd.DataFrame) -> float: + close = pd.to_numeric(df["close"], errors="coerce").to_numpy(dtype=float) + if len(close) < period + 1: + return 0.0 + prev = close[-period - 1] + return float(close[-1] / prev - 1.0) if prev > 0 else 0.0 + + return _score + + +def formula_score(formula_text_or_compiled: Any, value_col: str | None = None) -> ScoreFn: + """用通达信公式的数值输出做打分(与公式模块无缝联动)。""" + from easy_tdx.formula import CompiledFormula, compile_formula + + compiled = ( + formula_text_or_compiled + if isinstance(formula_text_or_compiled, CompiledFormula) + else compile_formula(formula_text_or_compiled) + ) + + def _score(df: pd.DataFrame) -> float: + result = compiled.compute(df) + if not result.values: + return 0.0 + col = value_col or result.values[-1] + arr = np.asarray(result.columns.get(col, [0.0]), dtype=float) + v = arr[-1] if len(arr) else 0.0 + return float(v) if math.isfinite(v) else 0.0 + + return _score + + +@dataclass +class _Position: + symbol: str + shares: float = 0.0 + cost: float = 0.0 # 平均成本(含费用近似) + + @property + def value_hint(self) -> float: + return self.shares * self.cost + + +@dataclass +class RotationResult: + """轮动回测结果。""" + + performance: dict[str, Any] = field(default_factory=dict) + equity_curve: list[dict[str, Any]] = field(default_factory=list) + trades: list[dict[str, Any]] = field(default_factory=list) + final_holdings: dict[str, dict[str, float]] = field(default_factory=dict) + rebalance_dates: list[str] = field(default_factory=list) + config: dict[str, Any] = field(default_factory=dict) + + def to_dict(self) -> dict[str, Any]: + return { + "performance": _clean(self.performance), + "equity_curve": _clean(self.equity_curve), + "trades": _clean(self.trades), + "final_holdings": _clean(self.final_holdings), + "n_rebalances": len(self.rebalance_dates), + "rebalance_dates": self.rebalance_dates, + "config": _clean(self.config), + } + + +def _clean(obj: Any) -> Any: + """numpy/Timestamp/NaN → JSON 原生(递归)。""" + if isinstance(obj, dict): + return {str(k): _clean(v) for k, v in obj.items()} + if isinstance(obj, (list, tuple)): + return [_clean(v) for v in obj] + if isinstance(obj, np.integer): + return int(obj) + if isinstance(obj, np.floating): + f = float(obj) + return f if math.isfinite(f) else None + if isinstance(obj, np.bool_): + return bool(obj) + if isinstance(obj, float) and not math.isfinite(obj): + return None + if obj is None or isinstance(obj, (str, int, bool)): + return obj + if hasattr(obj, "isoformat"): + return obj.isoformat() + return str(obj) + + +class RotationEngine: + """排名轮动组合回测。 + + Example:: + engine = RotationEngine( + stock_dfs={"SH:600519": df1, "SZ:000858": df2, ...}, + score_fn=momentum_score(20), + slots=3, + refresh="weekly", + ) + result = engine.run() + result.performance["total_return"] + """ + + def __init__( + self, + stock_dfs: dict[str, pd.DataFrame], + score_fn: ScoreFn, + slots: int = 5, + refresh: str = "weekly", + keep_rank: int | None = None, + cash: float = 1_000_000.0, + commission: float = 0.0003, + min_commission: float = 5.0, + stamp_tax: float = 0.001, + stop_loss: float | None = None, + take_profit: float | None = None, + max_score_history: int = 250, + ) -> None: + """Initialize. + + Args: + stock_dfs: 股票池(symbol → K 线 DataFrame,时间升序)。 + score_fn: 打分函数 ``f(df_prefix) -> float``(只喂截至当日的数据)。 + slots: 持仓槽位数(等额分配)。 + refresh: 调仓频率 ``daily`` / ``weekly`` / ``monthly``。 + keep_rank: 跌出前 N 名才卖出(默认 = slots,缓冲池可设更大)。 + cash: 初始资金。 + commission / min_commission / stamp_tax: 费率(卖出收印花税)。 + stop_loss / take_profit: 槽内止损/止盈(比例,如 0.1 = ±10%)。 + max_score_history: 预计算的打分滚动窗口上限(性能保护)。 + """ + if not stock_dfs: + raise ValueError("stock_dfs 不能为空") + if slots < 1: + raise ValueError("slots 必须 ≥ 1") + if refresh not in ("daily", "weekly", "monthly"): + raise ValueError(f"refresh 只支持 daily/weekly/monthly,当前 {refresh}") + self._dfs = {sym: self._normalize(df) for sym, df in stock_dfs.items() if len(df) >= 2} + if len(self._dfs) < 2: + raise ValueError("股票池有效标的不足 2 只(至少 2 根 K 线)") + self._score_fn = score_fn + self._slots = int(slots) + self._refresh = refresh + self._keep_rank = keep_rank or self._slots + self._cash = float(cash) + self._commission = commission + self._min_commission = min_commission + self._stamp_tax = stamp_tax + self._stop_loss = stop_loss + self._take_profit = take_profit + self._max_history = max_score_history + + # ── 主流程 ─────────────────────────────────────────────────────────────── + + def run(self) -> RotationResult: + result = RotationResult( + config={ + "slots": self._slots, + "refresh": self._refresh, + "keep_rank": self._keep_rank, + "cash": self._cash, + "commission": self._commission, + "stop_loss": self._stop_loss, + "take_profit": self._take_profit, + } + ) + calendar = self._common_calendar() + if len(calendar) < 10: + return result + + cash = self._cash + positions: dict[str, _Position] = {} + pending: list[tuple[str, str, str]] = [] # (symbol, direction, reason) 次开执行 + equity_records: list[dict[str, Any]] = [] + trades: list[dict[str, Any]] = [] + rebalances: list[str] = [] + + # 每标的有效 bar 指针:date -> 各标的最近一根 ≤ d 的 bar + pointers = {sym: -1 for sym in self._dfs} + + prev_key: tuple[int, ...] | None = None + peak = self._cash + equity_total = self._cash # 最近一日净值(槽位预算基准) + + for day_i, d in enumerate(calendar): + d_str = d.strftime("%Y-%m-%d") + + # 1. 推进各标的指针到 ≤ d 的最新一根 + bar_today: dict[str, pd.Series] = {} + for sym, df in self._dfs.items(): + dts = self._dt_index(sym) + while pointers[sym] + 1 < len(dts) and dts[pointers[sym] + 1] <= d: + pointers[sym] += 1 + if pointers[sym] >= 0: + bar_today[sym] = df.iloc[pointers[sym]] + + # 2. 次开执行昨日信号(用当日开盘价) + for sym, direction, reason in pending: + if sym not in bar_today: + continue + price = float(bar_today[sym]["open"]) + if not math.isfinite(price) or price <= 0: + continue + if direction == "SELL" and sym in positions: + pos = positions.pop(sym) + gross = pos.shares * price + fee = self._fee(gross, is_sell=True) + cash += gross - fee + pnl = (price - pos.cost) * pos.shares - fee + trades.append( + self._trade_row( + day_i, d_str, sym, "SELL", pos.shares, price, fee, pnl, reason=reason + ) + ) + elif direction == "BUY" and sym not in positions and cash > 0: + budget = self._slot_budget(cash, equity_total, len(positions)) + if budget <= price * 100: + continue + shares = math.floor(budget / (price * (1 + self._commission)) / 100) * 100 + if shares <= 0: + continue + gross = shares * price + fee = self._fee(gross, is_sell=False) + cash -= gross + fee + positions[sym] = _Position( + symbol=sym, shares=shares, cost=(gross + fee) / shares + ) + trades.append( + self._trade_row( + day_i, d_str, sym, "BUY", shares, price, fee, 0.0, reason=reason + ) + ) + pending = [] + + # 3. 止盈止损检查(收盘口径,次日执行) + for sym in list(positions): + if sym not in bar_today: + continue + close = float(bar_today[sym]["close"]) + cost = positions[sym].cost + if self._stop_loss is not None and close <= cost * (1 - self._stop_loss): + pending.append((sym, "SELL", "stop_loss")) + elif self._take_profit is not None and close >= cost * (1 + self._take_profit): + pending.append((sym, "SELL", "take_profit")) + + # 4. 调仓判定 + key = ( + (d.isocalendar()[0], d.isocalendar()[1]) + if self._refresh == "weekly" + else ((d.year, d.month) if self._refresh == "monthly" else (day_i,)) + ) + is_rebalance = key != prev_key + prev_key = key + if is_rebalance and day_i >= 1: + rebalances.append(d_str) + ranked = self._rank_all(pointers, d) + top_keep = [s for s, _ in ranked[: self._keep_rank]] + top_slots = [s for s, _ in ranked[: self._slots]] + # 卖出:跌出 keep_rank 的持仓 + for sym in list(positions): + if sym not in top_keep and all(p[0] != sym or p[1] != "SELL" for p in pending): + pending.append((sym, "SELL", "rank_exit")) + # 买入:前 slots 中未持有的(等空出的槽位) + free = self._slots - len(positions) + sum(1 for p in pending if p[1] == "SELL") + for sym in top_slots: + if free <= 0: + break + if sym not in positions and all(p[0] != sym or p[1] != "BUY" for p in pending): + pending.append((sym, "BUY", "rotation")) + free -= 1 + + # 5. 收盘估值(净值曲线点) + position_value = 0.0 + for sym, pos in positions.items(): + if sym in bar_today: + position_value += pos.shares * float(bar_today[sym]["close"]) + total = cash + position_value + equity_total = total + peak = max(peak, total) + dd_pct = (peak - total) / peak if peak > 0 else 0.0 + equity_records.append( + { + "datetime": d_str, + "cash": cash, + "position_value": position_value, + "total": total, + "drawdown": peak - total, + "drawdown_pct": dd_pct, + } + ) + + result.equity_curve = equity_records + result.trades = trades + result.rebalance_dates = rebalances + result.final_holdings = { + sym: {"shares": pos.shares, "cost": pos.cost} for sym, pos in positions.items() + } + result.performance = self._analyze(equity_records, trades) + return result + + # ── 辅助 ───────────────────────────────────────────────────────────────── + + @staticmethod + def _normalize(df: pd.DataFrame) -> pd.DataFrame: + out = df.copy() + dt_col = "datetime" if "datetime" in out.columns else "date" + out["_ts"] = pd.to_datetime(out[dt_col]) + return out.sort_values("_ts").reset_index(drop=True) + + def _common_calendar(self) -> pd.DatetimeIndex: + all_dts = pd.DatetimeIndex([]) + for sym in self._dfs: + all_dts = all_dts.union(self._dt_index(sym)) + return all_dts.sort_values() + + def _dt_index(self, sym: str) -> pd.DatetimeIndex: + return pd.DatetimeIndex(self._dfs[sym]["_ts"]) + + def _rank_all(self, pointers: dict[str, int], d: Any) -> list[tuple[str, float]]: + """对全部标的按截至 d 的前缀数据打分并降序排名。""" + scored: list[tuple[str, float]] = [] + for sym, df in self._dfs.items(): + idx = pointers[sym] + if idx < 5: + scored.append((sym, 0.0)) + continue + start = max(0, idx - self._max_history) + prefix = df.iloc[start : idx + 1].drop(columns=["_ts"], errors="ignore") + try: + s = float(self._score_fn(prefix)) + except Exception: # noqa: BLE001 — 单标的打分失败按 0 处理 + s = 0.0 + scored.append((sym, s if math.isfinite(s) else 0.0)) + scored.sort(key=lambda x: x[1], reverse=True) + return scored + + def _slot_budget(self, cash: float, equity_total: float, held: int) -> float: + """空槽预算 = 当前净值 / 槽数(等额口径),受剩余现金约束。""" + target = equity_total / self._slots + return max(0.0, min(target, cash)) + + def _fee(self, gross: float, *, is_sell: bool) -> float: + fee = max(gross * self._commission, self._min_commission) + if is_sell: + fee += gross * self._stamp_tax + return fee + + @staticmethod + def _trade_row( + day: int, + d_str: str, + sym: str, + direction: str, + size: float, + price: float, + fee: float, + pnl: float, + reason: str, + ) -> dict[str, Any]: + return { + "datetime": d_str, + "symbol": sym, + "direction": direction, + "size": size, + "price": price, + "commission": fee, + "slippage": 0.0, + "pnl": pnl, + "cost_basis": 0.0 if direction == "BUY" else price * size, + "rejected": False, + "reason": reason, + } + + @staticmethod + def _analyze( + equity_records: list[dict[str, Any]], trades: list[dict[str, Any]] + ) -> dict[str, Any]: + """复用主引擎绩效分析器(19 项指标)。""" + if len(equity_records) < 2: + return {} + equity = pd.DataFrame(equity_records) + trades_df = ( + pd.DataFrame(trades) + if trades + else pd.DataFrame( + {"datetime": [], "direction": [], "pnl": [], "rejected": [], "size": []} + ) + ) + trades_df["rejected"] = trades_df.get("rejected", False) + analyzer = PerformanceAnalyzer(equity, trades_df, risk_free_rate=0.03) + return dict(analyzer.compute()) diff --git a/src/easy_tdx/cli/__init__.py b/src/easy_tdx/cli/__init__.py index f7bd32e..7d16ffb 100644 --- a/src/easy_tdx/cli/__init__.py +++ b/src/easy_tdx/cli/__init__.py @@ -23,6 +23,7 @@ from .cmd_company import company_info, company_info_content, finance_info from .cmd_ex import ex from .cmd_factor import factor from .cmd_finance import f10, fund_flow +from .cmd_formula import formula from .cmd_indicator import indicator, indicator_list from .cmd_info import server_info, symbol_info from .cmd_kline import kline @@ -98,3 +99,4 @@ cli.add_command(run_all) cli.add_command(screen) cli.add_command(serve) cli.add_command(warehouse) +cli.add_command(formula) diff --git a/src/easy_tdx/cli/cmd_formula.py b/src/easy_tdx/cli/cmd_formula.py new file mode 100644 index 0000000..274a1b2 --- /dev/null +++ b/src/easy_tdx/cli/cmd_formula.py @@ -0,0 +1,208 @@ +"""``easy-tdx formula`` 命令组:通达信公式的计算 / 选股 / 回测。 + +公式方言见 :mod:`easy_tdx.formula`(命名布尔输出 = 信号列)。 + +示例:: + + easy-tdx formula compute SH 600519 --formula "金叉: CROSS(MA(C,5), MA(C,20));" + easy-tdx formula screen --symbols SH:600519,SZ:000001 --formula "..." --signal 金叉 + easy-tdx formula backtest SH 600519 --file my_formula.txt +""" + +from __future__ import annotations + +import json +from typing import Any + +import click +import pandas as pd + + +@click.group("formula") +def formula() -> None: + """通达信公式:计算 / 选股 / 回测(命名布尔输出即信号)。""" + + +def _load_formula(text: str | None, file: str | None) -> str: + from easy_tdx.formula import compile_formula + + if file: + from pathlib import Path + + source = Path(file).read_text(encoding="utf-8") + elif text: + source = text + else: + click.echo("错误: 必须提供 --formula 或 --file", err=True) + raise SystemExit(1) + compile_formula(source) # 提前暴露语法错误 + return source + + +def _fetch(market: str, code: str, count: int, adjust: str) -> pd.DataFrame: + from ..cli.conn import get_mac_client + from ..cli.parsers import parse_adjust, parse_market, parse_period + + with get_mac_client() as client: + return client.get_stock_kline( + parse_market(market), + code, + period=parse_period("DAILY"), + start=0, + count=count, + adjust=parse_adjust(adjust), + ) + + +@formula.command("compute") +@click.argument("market", type=click.Choice(["SH", "SZ", "BJ"])) +@click.argument("code") +@click.option("--formula", "formula_text", default=None, help="公式文本") +@click.option("--file", "formula_file", default=None, help="公式文件路径(.txt)") +@click.option("--count", default=120, type=int, help="K 线根数(默认 120)") +@click.option( + "--adjust", default="QFQ", type=click.Choice(["NONE", "QFQ", "HFQ"]), help="复权(默认 QFQ)" +) +@click.option("--tail", default=10, type=int, help="同时输出最近 N 根的信号明细(默认 10)") +def formula_compute( + market: str, + code: str, + formula_text: str | None, + formula_file: str | None, + count: int, + adjust: str, + tail: int, +) -> None: + """在指定标的上计算公式,输出最后一根的各列值 + 最近信号明细。""" + from easy_tdx.formula import compile_formula + + source = _load_formula(formula_text, formula_file) + df = _fetch(market, code, count, adjust) + if df is None or len(df) == 0: + click.echo(f"错误: {market}:{code} 未取到 K 线", err=True) + raise SystemExit(1) + result = compile_formula(source).compute(df) + payload = { + "symbol": f"{market}:{code}", + "signals": result.signals, + "values": result.values, + "last_row": result.last_row(), + } + if tail > 0 and result.columns: + dt_col = "datetime" if "datetime" in df.columns else "date" + recent_idx = df[dt_col].iloc[-tail:] + recent = result.to_frame().iloc[-tail:] + recent.insert(0, "date", [str(pd.Timestamp(d).date()) for d in recent_idx]) + payload["recent"] = json.loads(recent.to_json(orient="records", force_ascii=False)) + click.echo(json.dumps(payload, ensure_ascii=False, default=str)) + + +@formula.command("screen") +@click.option("--symbols", required=True, help="标的列表:逗号分隔或 @文件(SH:600519,SZ:000001)") +@click.option("--formula", "formula_text", default=None, help="公式文本") +@click.option("--file", "formula_file", default=None, help="公式文件路径") +@click.option( + "--signal", "signal_col", default=None, help="筛选信号列(默认第一个信号列,最后一根=1 通过)" +) +@click.option("--count", default=120, type=int, help="每标的 K 线根数") +@click.option( + "--adjust", default="QFQ", type=click.Choice(["NONE", "QFQ", "HFQ"]), help="复权(默认 QFQ)" +) +def formula_screen( + symbols: str, + formula_text: str | None, + formula_file: str | None, + signal_col: str | None, + count: int, + adjust: str, +) -> None: + """批量选股:公式信号列在最后一根 = 1 的标的(附各数值列,供排序)。""" + from easy_tdx.formula import compile_formula + + source = _load_formula(formula_text, formula_file) + if symbols.startswith("@"): + from pathlib import Path + + symbol_list = [ + ln.strip() + for ln in Path(symbols[1:]).read_text(encoding="utf-8").splitlines() + if ln.strip() and not ln.strip().startswith("#") + ] + else: + symbol_list = [s.strip() for s in symbols.split(",") if s.strip()] + + compiled = compile_formula(source) + hits: list[dict[str, Any]] = [] + errors: list[dict[str, str]] = [] + for sym in symbol_list: + market, code = sym.split(":", 1) + try: + df = _fetch(market, code, count, adjust) + if df is None or len(df) == 0: + raise ValueError("未取到 K 线") + result = compiled.compute(df) + col = signal_col or (result.signals[0] if result.signals else None) + if col is None: + raise ValueError("公式无布尔信号输出") + if result.last_row().get(col, 0.0) >= 1.0: + row: dict[str, Any] = {"symbol": sym} + row.update({k: v for k, v in result.last_row().items() if k != col}) + hits.append(row) + except Exception as exc: # noqa: BLE001 — 单标的失败跳过 + errors.append({"symbol": sym, "error": str(exc)}) + click.echo( + json.dumps( + {"total": len(symbol_list), "hits": hits, "errors": errors}, + ensure_ascii=False, + default=str, + ) + ) + + +@formula.command("backtest") +@click.argument("market", type=click.Choice(["SH", "SZ", "BJ"])) +@click.argument("code") +@click.option("--formula", "formula_text", default=None, help="公式文本") +@click.option("--file", "formula_file", default=None, help="公式文件路径") +@click.option("--buy", "buy_col", default=None, help="买入信号列(默认自动挑选)") +@click.option("--sell", "sell_col", default=None, help="卖出信号列(默认自动挑选)") +@click.option("--count", default=500, type=int, help="K 线根数(默认 500)") +@click.option("--cash", default=100000.0, type=float, help="初始资金") +@click.option( + "--adjust", default="QFQ", type=click.Choice(["NONE", "QFQ", "HFQ"]), help="复权(默认 QFQ)" +) +@click.option("--auto-fees", "auto_fees", is_flag=True, help="按品种自动费率(ETF 免印花税等)") +def formula_backtest( + market: str, + code: str, + formula_text: str | None, + formula_file: str | None, + buy_col: str | None, + sell_col: str | None, + count: int, + cash: float, + adjust: str, + auto_fees: bool, +) -> None: + """公式回测:信号列在下一根开盘成交(next_open),输出绩效 + 评级 + 评分。""" + from easy_tdx.backtest.formula_strategy import run_formula_backtest + + source = _load_formula(formula_text, formula_file) + df = _fetch(market, code, count, adjust) + if df is None or len(df) == 0: + click.echo(f"错误: {market}:{code} 未取到 K 线", err=True) + raise SystemExit(1) + try: + out = run_formula_backtest( + df, + source, + buy_col=buy_col, + sell_col=sell_col, + cash=cash, + symbol=f"{market}:{code}", + auto_fees=auto_fees, + ) + except ValueError as exc: + click.echo(f"错误: {exc}", err=True) + raise SystemExit(1) + click.echo(json.dumps(out, ensure_ascii=False, default=str)) diff --git a/src/easy_tdx/formula.py b/src/easy_tdx/formula.py new file mode 100644 index 0000000..2254976 --- /dev/null +++ b/src/easy_tdx/formula.py @@ -0,0 +1,519 @@ +"""通达信公式解析器(v1.27 新增)。 + +把通达信/麦语言风格的技术指标公式翻译成 numpy 向量计算,让写惯公式的 +用户零 Python 进入 easy-tdx 的筛选/回测体系(借鉴 indicator-lab 的公式 +解析思路,实现为独立子集方言)。 + +支持的方言子集:: + + {注释花括号} + N := 9; { 中间变量(参数) } + RSV := (C - LLV(L, N)) / (HHV(H, N) - LLV(L, N)) * 100; + K := SMA(RSV, 3, 1); + 金叉: CROSS(K, D); { 命名布尔输出 → 信号列 } + 强度: K - D; { 命名数值输出 → 排名/卖出参考列 } + +语法规则: + +- 语句以 ``;`` 结尾;``NAME := expr`` 为中间变量、``NAME: expr`` 为输出; + 裸表达式作为匿名输出 ``OUTPUT_1``; +- 运算符:``+ - * /``(除零安全,分母 0 → NaN)、比较 ``> < >= <= =``、 + 逻辑 ``AND OR NOT``(兼容 ``&& || !``)、括号、一元负号; +- 序列名:``C/CLOSE, O/OPEN, H/HIGH, L/LOW, V/VOL/VOL, AMOUNT/AMT``; +- 函数白名单(全部后视函数,**无未来数据**):MA/EMA/SMA/WMA/DMA/HHV/LLV/ + REF/SUM/COUNT/CROSS/LONGCROSS/EXIST/EVERY/BARSLAST/IF/MAX/MIN/ABS/POW/ + SQRT/LN/LOG/EXP/STD/AVEDEV/MACD/KDJ/RSI/BOLL/CCI/ATR/OBM/DMI 等 + (映射到 MyTT 与 numpy,见 :data:`_FUNCTIONS`); +- 输出归类:布尔表达式(比较/逻辑/CROSS 等)的命名输出 → **信号列** + (``signals``);数值表达式 → **数值列**(``values``,用于排名/阈值)。 + +安全:自建 tokenizer + AST 求值,**不走 Python eval**;未知函数/变量报 +:class:`FormulaError`(带位置)。 +""" + +from __future__ import annotations + +import re +from collections.abc import Callable +from dataclasses import dataclass, field +from typing import Any + +import numpy as np +import pandas as pd + +__all__ = ["FormulaError", "FormulaResult", "CompiledFormula", "compile_formula"] + +# ── Token ───────────────────────────────────────────────────────────────────── + +_TOKEN_RE = re.compile( + r""" + (?P\s+) + | (?P\{[^}]*\}) + | (?P\d+\.\d+|\.\d+|\d+) + | (?P[A-Za-z_\u4e00-\u9fff][A-Za-z0-9_\u4e00-\u9fff]*) + | (?P:=|>=|<=|==|&&|\|\||[-+*/(),:;> list[_Token]: + tokens: list[_Token] = [] + i = 0 + while i < len(text): + m = _TOKEN_RE.match(text, i) + if m is None: + raise FormulaError(f"无法识别的字符 {text[i]!r}(位置 {i})", pos=i) + i = m.end() + if m.lastgroup in ("ws", "comment"): + continue + tokens.append(_Token(kind=m.lastgroup or "op", value=m.group(), pos=m.start())) + tokens.append(_Token(kind="eof", value="", pos=len(text))) + return tokens + + +# ── AST ─────────────────────────────────────────────────────────────────────── + + +@dataclass +class _Node: + """表达式节点(用元数据极简表示,求值器按 kind 分派)。""" + + kind: str # num / name / call / bin / un / cmp / logic + value: int | float | str | None = None + children: list[_Node] = field(default_factory=list) + + +@dataclass +class _Statement: + """一条语句:中间赋值(is_output=False)或命名输出。""" + + name: str | None + expr: _Node + is_output: bool + pos: int + + +# ── Parser(递归下降)──────────────────────────────────────────────────────── + + +class _Parser: + def __init__(self, tokens: list[_Token]) -> None: + self._tokens = tokens + self._i = 0 + + def _peek(self) -> _Token: + return self._tokens[self._i] + + def _next(self) -> _Token: + tok = self._tokens[self._i] + self._i += 1 + return tok + + def _expect_op(self, op: str) -> _Token: + tok = self._peek() + if tok.kind == "op" and tok.value == op: + return self._next() + raise FormulaError(f"期望 {op!r},得到 {tok.value!r}(位置 {tok.pos})", pos=tok.pos) + + def _match_op(self, *ops: str) -> _Token | None: + tok = self._peek() + if tok.kind == "op" and tok.value in ops: + return self._next() + # 关键字运算符(AND/OR/NOT 分词为 name,按大写匹配) + if tok.kind == "name" and tok.value.upper() in ops: + return self._next() + return None + + def parse_statements(self) -> list[_Statement]: + stmts: list[_Statement] = [] + anonymous = 0 + while self._peek().kind != "eof": + tok = self._peek() + if tok.kind == "op" and tok.value == ";": # 空语句 + self._next() + continue + if tok.kind != "name": + raise FormulaError( + f"期望变量名开头,得到 {tok.value!r}(位置 {tok.pos})", pos=tok.pos + ) + # NAME := expr | NAME : expr | 裸表达式 + if ( + self._tokens[self._i + 1].kind == "op" + and self._tokens[self._i + 1].value in (":=", ":") + and not ( + self._tokens[self._i + 1].value == ":" + and self._tokens[self._i + 2].kind == "op" + and self._tokens[self._i + 2].value == "=" + ) + ): + name = self._next().value + assign = self._next() # := 或 : + expr = self.parse_expression() + self._expect_op(";") + stmts.append( + _Statement(name=name, expr=expr, is_output=(assign.value == ":"), pos=tok.pos) + ) + else: + anonymous += 1 + expr = self.parse_expression() + self._expect_op(";") + stmts.append( + _Statement(name=f"OUTPUT_{anonymous}", expr=expr, is_output=True, pos=tok.pos) + ) + return stmts + + # 表达式优先级:OR < AND < 比较 < 加减 < 乘除 < 一元 < 原子 + def parse_expression(self) -> _Node: + return self._parse_or() + + def _parse_or(self) -> _Node: + left = self._parse_and() + while tok := self._match_op("OR", "||"): + right = self._parse_and() + left = _Node(kind="logic", value="or", children=[left, right]) + left.pos_hint = tok.pos # type: ignore[attr-defined] + return left + + def _parse_and(self) -> _Node: + left = self._parse_cmp() + while tok := self._match_op("AND", "&&"): + right = self._parse_cmp() + left = _Node(kind="logic", value="and", children=[left, right]) + left.pos_hint = tok.pos # type: ignore[attr-defined] + return left + + def _parse_cmp(self) -> _Node: + left = self._parse_add() + while tok := self._match_op(">", "<", ">=", "<=", "=", "=="): + right = self._parse_add() + op = "==" if tok.value in ("=", "==") else tok.value + left = _Node(kind="cmp", value=op, children=[left, right]) + return left + + def _parse_add(self) -> _Node: + left = self._parse_mul() + while tok := self._match_op("+", "-"): + right = self._parse_mul() + left = _Node(kind="bin", value=tok.value, children=[left, right]) + return left + + def _parse_mul(self) -> _Node: + left = self._parse_unary() + while tok := self._match_op("*", "/"): + right = self._parse_unary() + left = _Node(kind="bin", value=tok.value, children=[left, right]) + return left + + def _parse_unary(self) -> _Node: + if tok := self._match_op("-", "+"): + child = self._parse_unary() + if tok.value == "-": + return _Node(kind="un", value="neg", children=[child]) + return child + if tok := self._match_op("!", "NOT"): + child = self._parse_unary() + return _Node(kind="un", value="not", children=[child]) + return self._parse_primary() + + def _parse_primary(self) -> _Node: + tok = self._peek() + if tok.kind == "num": + self._next() + v = float(tok.value) + # 整数字面量保持 int(MyTT 窗口/周期参数要求 int) + if v.is_integer() and abs(v) < 1e15: + v = int(v) + return _Node(kind="num", value=v) + if tok.kind == "op" and tok.value == "(": + self._next() + node = self.parse_expression() + self._expect_op(")") + return node + if tok.kind == "name": + self._next() + # 函数调用 + if self._peek().kind == "op" and self._peek().value == "(": + self._next() + args: list[_Node] = [] + if not (self._peek().kind == "op" and self._peek().value == ")"): + args.append(self.parse_expression()) + while self._match_op(","): + args.append(self.parse_expression()) + self._expect_op(")") + return _Node(kind="call", value=tok.value.upper(), children=args) + return _Node(kind="name", value=tok.value) + raise FormulaError(f"意外的记号 {tok.value!r}(位置 {tok.pos})", pos=tok.pos) + + +# ── 序列与函数环境 ───────────────────────────────────────────────────────────── + +_SERIES_ALIASES: dict[str, str] = { + "C": "close", + "CLOSE": "close", + "收盘价": "close", + "O": "open", + "OPEN": "open", + "开盘价": "open", + "H": "high", + "HIGH": "high", + "最高价": "high", + "L": "low", + "LOW": "low", + "最低价": "low", + "V": "vol", + "VOL": "vol", + "VOLUME": "vol", + "成交量": "vol", + "AMOUNT": "amount", + "AMT": "amount", + "成交额": "amount", +} + +_BOOL_FUNCS = {"CROSS", "LONGCROSS", "EXIST", "EVERY"} # 返回布尔的函数 + + +def _build_functions() -> dict[str, Callable[..., Any]]: + """函数白名单:MyTT 后视函数 + numpy 补齐(不透传任意 Python)。""" + import easy_tdx.MyTT as mytt + + fns: dict[str, Callable[..., Any]] = {} + for name in ( + "MA", + "EMA", + "SMA", + "WMA", + "DMA", + "HHV", + "LLV", + "REF", + "SUM", + "COUNT", + "CROSS", + "LONGCROSS", + "EXIST", + "EVERY", + "BARSLAST", + "IF", + "MAX", + "MIN", + "ABS", + "STD", + "AVEDEV", + "MACD", + "KDJ", + "RSI", + "BOLL", + "CCI", + "ATR", + "OBV", + "DMI", + "FILTER", + ): + if hasattr(mytt, name): + fns[name] = getattr(mytt, name) + # numpy 补齐(TDX 语义) + fns["POW"] = np.power + fns["SQRT"] = np.sqrt + fns["LN"] = np.log + fns["LOG"] = np.log10 + fns["EXP"] = np.exp + fns["NOT"] = np.logical_not + return fns + + +_FUNCTIONS: dict[str, Callable[..., Any]] | None = None + + +def _functions() -> dict[str, Callable[..., Any]]: + global _FUNCTIONS # noqa: PLW0603 — 模块级缓存 + if _FUNCTIONS is None: + _FUNCTIONS = _build_functions() + return _FUNCTIONS + + +# ── 求值器 ──────────────────────────────────────────────────────────────────── + + +class _Evaluator: + def __init__(self, df: pd.DataFrame) -> None: + self._arrays: dict[str, np.ndarray] = {} + for col in df.columns: + if col in ("datetime", "date"): + continue + try: + arr = pd.to_numeric(df[col], errors="coerce").to_numpy(dtype=float) + except (TypeError, ValueError): + continue # 非数值列(如文本)跳过 + self._arrays[str(col).lower()] = arr + self._vars: dict[str, Any] = {} + self._n = len(df) + + def eval_statements(self, stmts: list[_Statement]) -> FormulaResult: + result = FormulaResult(n=self._n) + for stmt in stmts: + val = self.eval(stmt.expr) + if stmt.name is not None: + self._vars[stmt.name.upper()] = val + if stmt.is_output and stmt.name is not None: + arr = np.asarray(val, dtype=float) + result.columns[stmt.name] = arr + if self._is_boolean(stmt.expr, val): + result.signals.append(stmt.name) + else: + result.values.append(stmt.name) + return result + + @staticmethod + def _is_boolean(expr: _Node, val: Any) -> bool: + """输出归类:比较/逻辑/CROSS 节点或 0/1 值域 → 信号列。""" + if expr.kind in ("cmp", "logic"): + return True + if expr.kind == "call" and expr.value in _BOOL_FUNCS: + return True + arr = np.asarray(val, dtype=float) + finite = arr[np.isfinite(arr)] + if finite.size == 0: + return False + return bool(finite.min() >= 0.0 and finite.max() <= 1.0) + + def eval(self, node: _Node) -> Any: + if node.kind == "num": + # 保持解析期类型(int 窗口参数 / float 数值) + return node.value + if node.kind == "name": + key = str(node.value) + upper = key.upper() + if upper in _SERIES_ALIASES: + col = _SERIES_ALIASES[upper] + if col not in self._arrays: + raise FormulaError(f"K 线数据缺少列 {col!r}(公式引用了 {key})") + return self._arrays[col] + if key in self._vars: + return self._vars[key] + if upper in self._vars: + return self._vars[upper] + raise FormulaError(f"未知变量 {key!r}(未定义且不是序列名/函数)") + if node.kind == "call": + fname = str(node.value) + fns = _functions() + if fname not in fns: + raise FormulaError(f"未知或不支持的函数 {fname}(白名单外)") + args = [self.eval(c) for c in node.children] + try: + with np.errstate(divide="ignore", invalid="ignore", over="ignore"): + out = fns[fname](*args) + except Exception as exc: # noqa: BLE001 — 包装带函数名 + raise FormulaError(f"函数 {fname} 求值失败:{exc}") from exc + return out + if node.kind == "bin": + a = np.asarray(self.eval(node.children[0]), dtype=float) + b = np.asarray(self.eval(node.children[1]), dtype=float) + a, b = np.broadcast_arrays(a, b) + if node.value == "+": + return a + b + if node.value == "-": + return a - b + if node.value == "*": + return a * b + if node.value == "/": + # 除零安全:分母 0 → NaN(不炸、不 inf) + with np.errstate(divide="ignore", invalid="ignore"): + out = np.divide(a, b, out=np.full(a.shape, np.nan), where=b != 0) + return out + raise FormulaError(f"未知运算符 {node.value}") + if node.kind == "un": + child = np.asarray(self.eval(node.children[0]), dtype=float) + return -child if node.value == "neg" else np.logical_not(child != 0).astype(float) + if node.kind == "cmp": + a = np.asarray(self.eval(node.children[0]), dtype=float) + b = np.asarray(self.eval(node.children[1]), dtype=float) + a, b = np.broadcast_arrays(a, b) + op = node.value + with np.errstate(invalid="ignore"): + if op == ">": + out = a > b + elif op == "<": + out = a < b + elif op == ">=": + out = a >= b + elif op == "<=": + out = a <= b + else: # == + out = np.isclose(a, b) + return out.astype(float) # NaN 参与比较 → False(0) + if node.kind == "logic": + a = np.asarray(self.eval(node.children[0]), dtype=float) + b = np.asarray(self.eval(node.children[1]), dtype=float) + a, b = np.broadcast_arrays(a, b) + if node.value == "and": + return ((a != 0) & (b != 0)).astype(float) + return ((a != 0) | (b != 0)).astype(float) + raise FormulaError(f"未知节点类型 {node.kind}") + + +# ── 公共 API ────────────────────────────────────────────────────────────────── + + +class FormulaError(ValueError): + """公式语法/求值错误(附位置信息)。""" + + def __init__(self, message: str, pos: int | None = None) -> None: + super().__init__(message if pos is None else f"{message} @col {pos}") + self.pos = pos + + +@dataclass +class FormulaResult: + """公式计算结果:命名输出列 + 信号/数值归类。""" + + columns: dict[str, np.ndarray] = field(default_factory=dict) + signals: list[str] = field(default_factory=list) # 布尔输出名(信号列) + values: list[str] = field(default_factory=list) # 数值输出名(排名列) + n: int = 0 + + def to_frame(self) -> pd.DataFrame: + """输出列拼成 DataFrame(保留声明顺序)。""" + if not self.columns: + return pd.DataFrame() + return pd.DataFrame(dict(self.columns)) + + def last_row(self) -> dict[str, float]: + """各输出列最后一根 bar 的值(选股扫描口径)。""" + out: dict[str, float] = {} + for name, arr in self.columns.items(): + arr = np.asarray(arr, dtype=float) + out[name] = float(arr[-1]) if len(arr) and np.isfinite(arr[-1]) else 0.0 + return out + + +class CompiledFormula: + """已编译的公式(解析一次,多处计算)。""" + + def __init__(self, text: str) -> None: + self._text = text + self._statements = _Parser(_tokenize(text)).parse_statements() + if not self._statements: + raise FormulaError("公式为空或只有注释") + + @property + def text(self) -> str: + return self._text + + def compute(self, df: pd.DataFrame) -> FormulaResult: + """在 K 线上计算公式(数据不足预热期自动为 NaN/0,不抛错)。""" + if df is None or len(df) == 0: + raise FormulaError("K 线数据为空") + return _Evaluator(df).eval_statements(self._statements) + + +def compile_formula(text: str) -> CompiledFormula: + """编译通达信公式文本(语法错误抛 :class:`FormulaError`)。""" + return CompiledFormula(text) diff --git a/src/easy_tdx/web/app.py b/src/easy_tdx/web/app.py index e0a997b..12082d5 100644 --- a/src/easy_tdx/web/app.py +++ b/src/easy_tdx/web/app.py @@ -239,6 +239,7 @@ def _create_app( from easy_tdx.web.routers.chanlun import router as chanlun_router from easy_tdx.web.routers.ex_market import router as ex_market_router from easy_tdx.web.routers.finance import router as finance_router + from easy_tdx.web.routers.formula import router as formula_router from easy_tdx.web.routers.indicator import router as indicator_router from easy_tdx.web.routers.mac_data import router as mac_data_router from easy_tdx.web.routers.mac_quotes import router as mac_quotes_router @@ -253,6 +254,7 @@ def _create_app( app.include_router(market_router, prefix="/api/v1") app.include_router(bars_router, prefix="/api/v1") app.include_router(finance_router, prefix="/api/v1") + app.include_router(formula_router, prefix="/api/v1") app.include_router(block_router, prefix="/api/v1") app.include_router(chanlun_router, prefix="/api/v1") app.include_router(realtime_router, prefix="/api/v1") diff --git a/src/easy_tdx/web/routers/formula.py b/src/easy_tdx/web/routers/formula.py new file mode 100644 index 0000000..4e9e805 --- /dev/null +++ b/src/easy_tdx/web/routers/formula.py @@ -0,0 +1,241 @@ +"""通达信公式路由:校验 / 计算 / 选股 / 回测(v1.27 新增)。 + +与 CLI(``easy-tdx formula ...``)和 Python API(:mod:`easy_tdx.formula`) +同口径——公式方言与信号归类见该模块文档。 +""" + +from __future__ import annotations + +from typing import Any + +from fastapi import APIRouter, Depends +from pydantic import BaseModel, Field + +from easy_tdx.web.deps import get_client +from easy_tdx.web.task_runner import get_runner + +router = APIRouter(tags=["formula"]) + + +class FormulaValidateRequest(BaseModel): + """公式校验请求(无需行情数据)。""" + + text: str = Field(..., min_length=1, max_length=8000) + + +class FormulaValidateResponse(BaseModel): + ok: bool + signals: list[str] = Field(default_factory=list) + values: list[str] = Field(default_factory=list) + error: str | None = None + + +class FormulaComputeRequest(FormulaValidateRequest): + """公式计算请求:内联 ohlcv 或按 symbol 取行情(二选一)。""" + + ohlcv: list[dict[str, Any]] | None = Field(default=None, max_length=2000) + symbol: str | None = Field(default=None, pattern=r"^(SZ|SH|BJ):\d{6}$") + category: str = "DAY" + count: int = Field(default=250, ge=20, le=2000) + tail: int = Field(default=10, ge=0, le=100, description="附带最近 N 根输出明细") + + +class FormulaBacktestRequest(FormulaComputeRequest): + """公式回测请求(信号列下一根开盘成交)。""" + + buy_col: str | None = Field(default=None, description="买入信号列(默认自动挑选)") + sell_col: str | None = Field(default=None, description="卖出信号列(默认自动挑选)") + cash: float = Field(default=1_000_000.0, gt=0) + commission: float = Field(default=0.0003, ge=0, le=0.01) + auto_fees: bool = Field(default=False) + + +class FormulaScreenRequest(FormulaValidateRequest): + """公式选股请求(后台任务,逐标的取行情)。""" + + symbols: list[str] = Field(..., min_length=1, max_length=50) + signal_col: str | None = None + category: str = "DAY" + count: int = Field(default=250, ge=60, le=2000) + + +@router.post("/formula/validate", response_model=FormulaValidateResponse) +async def validate_formula(req: FormulaValidateRequest) -> FormulaValidateResponse: + """校验公式语法并给出信号/数值输出归类(不取行情)。""" + from easy_tdx.formula import FormulaError, compile_formula + + try: + compiled = compile_formula(req.text) + # 无数据时的静态归类:检查输出表达式的 AST 类型 + signals: list[str] = [] + values: list[str] = [] + bool_kinds = {"cmp", "logic"} + bool_funcs = {"CROSS", "LONGCROSS", "EXIST", "EVERY"} + for stmt in compiled._statements: # noqa: SLF001 — 同包内部协定 + if stmt.is_output and stmt.name: + if stmt.expr.kind in bool_kinds or ( + stmt.expr.kind == "call" and stmt.expr.value in bool_funcs + ): + signals.append(stmt.name) + else: + values.append(stmt.name) + return FormulaValidateResponse(ok=True, signals=signals, values=values) + except FormulaError as exc: + return FormulaValidateResponse(ok=False, error=str(exc)) + + +@router.post("/formula/compute") +async def compute_formula( + req: FormulaComputeRequest, client: Any = Depends(get_client) +) -> dict[str, Any]: + """在 K 线上计算公式:最后一根各列值 + 可选最近 N 根明细。""" + import pandas as pd + + from easy_tdx.formula import FormulaError, compile_formula + + try: + compiled = compile_formula(req.text) + except FormulaError as exc: + raise ValueError(str(exc)) from exc + + df = await _resolve_df(client, req) + try: + result = compiled.compute(df) + except FormulaError as exc: + raise ValueError(str(exc)) from exc + + payload: dict[str, Any] = { + "signals": result.signals, + "values": result.values, + "last_row": result.last_row(), + } + if req.tail > 0 and result.columns: + dt_col = "datetime" if "datetime" in df.columns else "date" + recent = result.to_frame().iloc[-req.tail :] + recent.insert( + 0, "date", [str(pd.Timestamp(d).date()) for d in df[dt_col].iloc[-req.tail :]] + ) + payload["recent"] = recent.to_dict(orient="records") + return payload + + +@router.post("/formula/backtest/run/async", status_code=202) +async def run_formula_backtest_async( + req: FormulaBacktestRequest, client: Any = Depends(get_client) +) -> dict[str, str]: + """提交公式回测后台任务(信号列下一根开盘成交),轮询 /backtest/tasks/{id}。""" + df = await _resolve_df(client, req) + snapshot = req.model_copy() + runner = get_runner() + task_id = runner.submit( + lambda: _run_formula_backtest(df, snapshot), + description=f"公式回测 | {snapshot.symbol or '内联数据'}", + ) + return {"task_id": task_id, "status": "running"} + + +@router.post("/formula/screen/run/async", status_code=202) +async def run_formula_screen_async( + req: FormulaScreenRequest, client: Any = Depends(get_client) +) -> dict[str, str]: + """提交公式选股后台任务:信号列最后一根 = 1 的标的 + 数值列(供排序)。""" + from easy_tdx.formula import FormulaError, compile_formula + from easy_tdx.web.convert import category_from_str, market_from_str + + try: + compiled = compile_formula(req.text) + except FormulaError as exc: + raise ValueError(str(exc)) from exc + + bars: dict[str, Any] = {} + for symbol in req.symbols: + market_str, code = symbol.split(":", 1) + try: + page = await client.get_security_bars( + market_from_str(market_str), code, category_from_str(req.category), 0, req.count + ) + except Exception: # noqa: BLE001 — 单标的失败跳过 + continue + if page is not None and len(page) >= 30: + bars[symbol] = page + if not bars: + raise ValueError("所有标的均未取到有效 K 线") + + snapshot = req.model_copy() + runner = get_runner() + task_id = runner.submit( + lambda: _run_formula_screen(bars, compiled, snapshot.signal_col), + description=f"公式选股 | {len(bars)}只标的", + ) + return {"task_id": task_id, "status": "running"} + + +# ── 内部实现 ─────────────────────────────────────────────────────────────────── + + +async def _resolve_df(client: Any, req: FormulaComputeRequest) -> Any: + """内联 ohlcv 或按 symbol 取行情。""" + import pandas as pd + + if req.ohlcv is not None: + df = pd.DataFrame(req.ohlcv) + required = {"datetime", "open", "high", "low", "close", "vol"} + missing = required - set(df.columns) + if missing: + raise ValueError(f"ohlcv 缺少必需列: {sorted(missing)}") + if not pd.api.types.is_datetime64_any_dtype(df["datetime"]): + df["datetime"] = pd.to_datetime(df["datetime"], errors="coerce") + return df + if req.symbol is not None: + from easy_tdx.web.convert import category_from_str, market_from_str + + market_str, code = req.symbol.split(":", 1) + df = await client.get_security_bars( + market_from_str(market_str), code, category_from_str(req.category), 0, req.count + ) + if df is None or len(df) == 0: + raise ValueError(f"标的 {req.symbol} 未取到 K 线数据") + return df + raise ValueError("必须提供 ohlcv 或 symbol") + + +def _run_formula_backtest(df: Any, req: FormulaBacktestRequest) -> dict[str, Any]: + """执行公式回测(后台线程内调用)。""" + from easy_tdx.backtest.formula_strategy import run_formula_backtest as _run + + try: + return _run( + df, + req.text, + buy_col=req.buy_col, + sell_col=req.sell_col, + cash=req.cash, + commission=req.commission, + symbol=req.symbol, + auto_fees=req.auto_fees, + ) + except ValueError as exc: + raise ValueError(str(exc)) from exc + + +def _run_formula_screen( + bars: dict[str, Any], compiled: Any, signal_col: str | None +) -> dict[str, Any]: + """执行公式选股(后台线程内调用)。""" + from easy_tdx.formula import FormulaError + + hits: list[dict[str, Any]] = [] + errors: list[dict[str, str]] = [] + for symbol, df in bars.items(): + try: + result = compiled.compute(df) + col = signal_col or (result.signals[0] if result.signals else None) + if col is None: + raise ValueError("公式无布尔信号输出") + if result.last_row().get(col, 0.0) >= 1.0: + row: dict[str, Any] = {"symbol": symbol} + row.update({k: v for k, v in result.last_row().items() if k != col}) + hits.append(row) + except (FormulaError, ValueError) as exc: + errors.append({"symbol": symbol, "error": str(exc)}) + return {"total": len(bars), "hits": hits, "errors": errors} diff --git a/tests/unit/test_backtest_rotation.py b/tests/unit/test_backtest_rotation.py new file mode 100644 index 0000000..c89f18b --- /dev/null +++ b/tests/unit/test_backtest_rotation.py @@ -0,0 +1,159 @@ +"""轮动组合引擎测试(排名换仓 / 槽位等额 / 止盈止损 / 刷新频率 / 绩效)。""" + +from __future__ import annotations + +import json + +import numpy as np +import pandas as pd +import pytest + +from easy_tdx.backtest.rotation import RotationEngine, RotationResult, formula_score, momentum_score + + +def _stock( + n: int = 250, seed: int = 1, drift: float = 0.001, start: str = "2024-01-01" +) -> pd.DataFrame: + rng = np.random.default_rng(seed) + close = 10.0 * np.cumprod(1.0 + drift + rng.normal(0, 0.01, n)) + return pd.DataFrame( + { + "datetime": pd.date_range(start, periods=n, freq="B"), + "open": close * 0.999, + "high": close * 1.02, + "low": close * 0.98, + "close": close, + "vol": 1e6, + "amount": close * 1e6, + } + ) + + +def _pool(drifts: dict[str, float], n: int = 250) -> dict[str, pd.DataFrame]: + return {sym: _stock(n, seed=i, drift=drift) for i, (sym, drift) in enumerate(drifts.items())} + + +# ── 基础结构 ───────────────────────────────────────────────────────────────── + + +def test_rotation_basic_run_and_structure(): + pool = _pool({"SH:600519": 0.002, "SZ:000001": 0.001, "SZ:000858": 0.0005, "SH:601318": 0.0}) + engine = RotationEngine(pool, momentum_score(20), slots=2, refresh="weekly") + result = engine.run() + assert isinstance(result, RotationResult) + assert len(result.equity_curve) >= 200 + assert result.performance.get("total_return") is not None + assert result.config["slots"] == 2 + # 净值曲线字段完整(可喂组合评级) + first = result.equity_curve[0] + assert {"datetime", "cash", "position_value", "total", "drawdown_pct"} <= set(first) + + +def test_rotation_strong_pool_makes_money(): + """普涨池 + 动量排名 → 正收益。""" + pool = _pool({f"SH:60000{i}": 0.004 for i in range(5)}) + result = RotationEngine(pool, momentum_score(20), slots=3, refresh="monthly").run() + assert result.performance["total_return"] > 0 + + +def test_rotation_weak_pool_loses_less_than_buyhold(): + """普跌池 → 负收益(动量轮动不做空)。""" + pool = _pool({f"SH:60000{i}": -0.004 for i in range(5)}) + result = RotationEngine(pool, momentum_score(20), slots=2).run() + assert result.performance["total_return"] < 0 + + +def test_rotation_trades_have_reasons(): + pool = _pool({f"SH:60000{i}": 0.002 if i % 2 else -0.001 for i in range(6)}) + result = RotationEngine(pool, momentum_score(10), slots=2, refresh="weekly").run() + reasons = {t["reason"] for t in result.trades} + assert "rotation" in reasons # 买入 + assert "rank_exit" in reasons # 跌出排名的卖出 + + +def test_rotation_respects_slots(): + """持仓数永远 ≤ slots。""" + pool = _pool({f"SH:60000{i}": 0.001 + 0.0005 * i for i in range(8)}) + engine = RotationEngine(pool, momentum_score(10), slots=3, refresh="weekly") + # 用逐日持仓推断:trades 序列重放 + holdings = 0 + peak_holdings = 0 + for t in result_trades_sorted(engine): + if t["direction"] == "BUY": + holdings += 1 + peak_holdings = max(peak_holdings, holdings) + else: + holdings -= 1 + assert peak_holdings <= 3 + + +def result_trades_sorted(engine: RotationEngine) -> list[dict]: + result = engine.run() + return result.trades + + +def test_rotation_stop_loss_triggers(): + """深跌池 + 10% 止损 → 出现 stop_loss 卖出。""" + pool = _pool({f"SH:60000{i}": -0.006 for i in range(4)}) + result = RotationEngine( + pool, momentum_score(5), slots=2, refresh="monthly", stop_loss=0.05 + ).run() + reasons = {t["reason"] for t in result.trades} + assert "stop_loss" in reasons + + +def test_rotation_refresh_frequencies(): + pool = _pool({f"SH:60000{i}": 0.001 * (i + 1) for i in range(4)}) + r_daily = RotationEngine(pool, momentum_score(10), slots=2, refresh="daily").run() + r_monthly = RotationEngine(pool, momentum_score(10), slots=2, refresh="monthly").run() + # 月调仓的调仓日数 ≤ 日调仓 + assert len(r_monthly.rebalance_dates) <= len(r_daily.rebalance_dates) + # 月调仓约 12 次/年(250 交易日) + assert 3 <= len(r_monthly.rebalance_dates) <= 15 + + +def test_rotation_formula_score_synergy(): + """公式打分与轮动联动:数值输出作为排名分。""" + pool = _pool({f"SH:60000{i}": 0.001 * (i + 1) for i in range(4)}) + score = formula_score("动量分: C / REF(C, 20) * 100;") + result = RotationEngine(pool, score, slots=2, refresh="monthly").run() + assert result.performance["total_return"] is not None + + +def test_rotation_result_serializable(): + pool = _pool({f"SH:60000{i}": 0.001 * (i + 1) for i in range(4)}) + result = RotationEngine(pool, momentum_score(10), slots=2).run() + d = result.to_dict() + text = json.dumps(d, ensure_ascii=False) + assert "equity_curve" in text + assert d["n_rebalances"] >= 1 + + +def test_rotation_rejects_bad_config(): + pool = _pool({"SH:600519": 0.001, "SZ:000001": 0.001}) + with pytest.raises(ValueError, match="refresh"): + RotationEngine(pool, momentum_score(5), refresh="yearly") + with pytest.raises(ValueError, match="stock_dfs"): + RotationEngine({}, momentum_score(5)) + with pytest.raises(ValueError, match="slots"): + RotationEngine(pool, momentum_score(5), slots=0) + + +def test_rotation_equal_weight_no_allin_single_stock(): + """首日建仓是等额分批,不是一把全买一只(槽位预算 = 净值/槽数)。""" + pool = _pool({f"SH:60000{i}": 0.001 * (i + 1) for i in range(6)}) + result = RotationEngine(pool, momentum_score(10), slots=3, refresh="monthly").run() + first_day_buys = [ + t for t in result.trades if t["direction"] == "BUY" and t["reason"] == "rotation" + ][:3] + if len(first_day_buys) >= 2: + values = [t["size"] * t["price"] for t in first_day_buys] + # 同日买入的各笔金额接近(等额),差异 < 25%(价格整百取整的摩擦) + assert max(values) / max(min(values), 1) < 1.25 + + +def test_momentum_score_helper(): + df = _stock(30, seed=1, drift=0.01) + score = momentum_score(10)(df) + assert score > 0 + assert momentum_score(10)(_stock(5)) == 0.0 # 数据不足 → 0 diff --git a/tests/unit/test_formula.py b/tests/unit/test_formula.py new file mode 100644 index 0000000..98e17cf --- /dev/null +++ b/tests/unit/test_formula.py @@ -0,0 +1,214 @@ +"""通达信公式解析器测试(tokenizer / AST / 白名单求值 / 信号归类 / 安全性)。""" + +from __future__ import annotations + +import numpy as np +import pandas as pd +import pytest + +from easy_tdx.formula import FormulaError, compile_formula + + +def _df(n: int = 60, seed: int = 3) -> pd.DataFrame: + rng = np.random.default_rng(seed) + dates = pd.date_range("2024-01-01", periods=n, freq="B") + close = 10.0 * np.cumprod(1.0 + 0.002 + rng.normal(0, 0.015, n)) + return pd.DataFrame( + { + "datetime": dates, + "open": close * 0.999, + "high": close * 1.02, + "low": close * 0.98, + "close": close, + "vol": rng.uniform(1e6, 5e6, n), + "amount": close * rng.uniform(1e6, 5e6, n), + } + ) + + +# ── 编译与语法 ─────────────────────────────────────────────────────────────── + + +def test_compile_and_outputs(): + formula = compile_formula( + """ + N := 9; + RSV := (C - LLV(L, N)) / (HHV(H, N) - LLV(L, N)) * 100; + K := SMA(RSV, 3, 1); + 金叉: CROSS(K, 20); + 强度: K; + """ + ) + res = formula.compute(_df()) + assert "金叉" in res.signals + assert "强度" in res.values + frame = res.to_frame() + assert list(frame.columns) == ["金叉", "强度"] + assert len(frame) == 60 + + +def test_syntax_error_has_position(): + with pytest.raises(FormulaError): + compile_formula("A := ;") + with pytest.raises(FormulaError): + compile_formula("A := UNKNOWN_FUNC(C)") + with pytest.raises(FormulaError): + compile_formula("A := B + ") # 引用未定义变量且语法断裂 + + +def test_unknown_variable_rejected(): + with pytest.raises(FormulaError, match="未知变量"): + compile_formula("A: X1;").compute(_df()) + + +def test_unknown_function_rejected(): + with pytest.raises(FormulaError, match="白名单"): + compile_formula("A: EVAL(C);").compute(_df()) + + +def test_empty_formula_rejected(): + with pytest.raises(FormulaError, match="为空"): + compile_formula("{只有注释}") + + +def test_no_python_eval_injection(): + """公式层不走 Python eval:危险标识符按未知变量/函数拒绝。""" + with pytest.raises(FormulaError): + compile_formula("__import__('os'): 1;").compute(_df()) + with pytest.raises(FormulaError): + compile_formula("A: OPEN(C);").compute(_df()) + + +# ── 语义正确性 ─────────────────────────────────────────────────────────────── + + +def test_series_aliases(): + """C/O/H/L/V/AMOUNT 别名与底层列一致。""" + df = _df(50) + res = compile_formula("高价: H; 低价: L; 收盘: CLOSE; 量: VOL; 额: AMOUNT;").compute(df) + assert np.allclose(res.columns["高价"], df["high"]) + assert np.allclose(res.columns["收盘"], df["close"]) + assert np.allclose(res.columns["量"], df["vol"]) + + +def test_ma_matches_mytt(): + from easy_tdx.MyTT import MA + + df = _df(50) + res = compile_formula("均线: MA(C, 5);").compute(df) + assert np.allclose(res.columns["均线"], MA(df["close"].to_numpy(), 5), equal_nan=True) + + +def test_cross_semantics(): + """CROSS(A,B):A 上穿 B 的那一根为 1,其余 0。""" + df = _df(50) + res = compile_formula( + """ + 快: MA(C, 3); + 慢: MA(C, 10); + 金叉: CROSS(快, 慢); + """ + ).compute(df) + golden = res.columns["金叉"] + assert set(np.unique(golden[np.isfinite(golden)])).issubset({0.0, 1.0}) + assert golden.sum() >= 0 # 结构完整(趋势数据至少存在或为 0) + # CROSS 手工复算对拍 + from easy_tdx.MyTT import CROSS, MA + + fast = MA(df["close"].to_numpy(), 3) + slow = MA(df["close"].to_numpy(), 10) + assert np.allclose(golden, CROSS(fast, slow), equal_nan=True) + + +def test_safe_division_zero_denominator_nan(): + """除零 → NaN(不炸、不 inf)。""" + df = _df(30) + res = compile_formula("比值: C / (C - C);").compute(df) # 分母全 0 + assert np.isnan(res.columns["比值"]).all() + + +def test_logic_operators(): + df = _df(40) + res = compile_formula( + """ + 条件1: C > MA(C, 5); + 条件2: C > MA(C, 20); + 同时: 条件1 AND 条件2; + 任一: 条件1 OR 条件2; + 取反: NOT(条件1); + """ + ).compute(df) + c1 = res.columns["条件1"] > 0.5 + c2 = res.columns["条件2"] > 0.5 + assert np.allclose(res.columns["同时"] > 0.5, c1 & c2) + assert np.allclose(res.columns["任一"] > 0.5, c1 | c2) + assert np.allclose(res.columns["取反"] > 0.5, ~c1) + + +def test_comparison_and_unary(): + df = _df(30) + res = compile_formula("跌幅: -(C - REF(C, 1)) / REF(C, 1) * 100; 平: C == C;").compute(df) + assert (res.columns["平"] == 1.0).all() + assert "跌幅" in res.values + + +def test_warmup_nan_not_signal(): + """预热期 NaN 不产生信号(比较含 NaN → 0)。""" + df = _df(30) + res = compile_formula("信号: CROSS(MA(C, 20), MA(C, 25));").compute(df) + sig = res.columns["信号"] + assert np.nanmax(np.nan_to_num(sig[:25])) <= 1.0 + assert np.isnan(sig).sum() == 0 # 布尔输出不含 NaN + + +def test_output_classification_boolean_vs_numeric(): + """比较/逻辑输出 → 信号;数值输出 → 数值列;0/1 值域数值也归信号。""" + df = _df(40) + res = compile_formula( + """ + 布尔输出: C > REF(C, 1); + 数值输出: MA(C, 5) - MA(C, 20); + """ + ).compute(df) + assert res.signals == ["布尔输出"] + assert res.values == ["数值输出"] + + +def test_last_row_for_screening(): + df = _df(30) + res = compile_formula("买入: CROSS(MA(C, 3), MA(C, 10)); 值: MA(C, 5);").compute(df) + last = res.last_row() + assert set(last) == {"买入", "值"} + assert last["买入"] in (0.0, 1.0) + + +def test_chinese_identifier_and_comment(): + df = _df(30) + formula = compile_formula( + """ + {这是注释:N 周期} + 周期 := 5; + 均线: MA(C, 周期); + """ + ) + res = formula.compute(df) + assert res.columns["均线"][0] != res.columns["均线"][-1] + + +def test_compiled_formula_reusable_across_frames(): + f = compile_formula("值: MA(C, 5);") + r1 = f.compute(_df(30, seed=1)) + r2 = f.compute(_df(40, seed=2)) + assert len(r1.columns["值"]) == 30 + assert len(r2.columns["值"]) == 40 + + +def test_compiled_formula_is_dataclass_safe(): + """CompiledFormula 可 pickle(进程池/后台任务传输)。""" + import pickle + + f = compile_formula("值: MA(C, 5);") + f2 = pickle.loads(pickle.dumps(f)) + assert np.allclose( + f.compute(_df(20)).columns["值"], f2.compute(_df(20)).columns["值"], equal_nan=True + ) diff --git a/tests/unit/test_formula_integration.py b/tests/unit/test_formula_integration.py new file mode 100644 index 0000000..c03d4aa --- /dev/null +++ b/tests/unit/test_formula_integration.py @@ -0,0 +1,193 @@ +"""公式回测适配器 + REST 端点测试(三通道一致性)。""" + +from __future__ import annotations + +import time + +import numpy as np +import pandas as pd +import pytest + +pytest.importorskip("fastapi") + +from fastapi.testclient import TestClient # noqa: E402 + +from easy_tdx.backtest.formula_strategy import ( # noqa: E402 + FormulaStrategyError, + attach_formula_columns, + pick_signal_columns, + run_formula_backtest, +) +from easy_tdx.formula import compile_formula # noqa: E402 + + +def _df(n: int = 300, seed: int = 3, drift: float = 0.002) -> pd.DataFrame: + rng = np.random.default_rng(seed) + close = 10.0 * np.cumprod(1.0 + drift + rng.normal(0, 0.012, n)) + return pd.DataFrame( + { + "datetime": pd.date_range("2024-01-01", periods=n, freq="B"), + "open": close * 0.999, + "high": close * 1.02, + "low": close * 0.98, + "close": close, + "vol": 1e6, + "amount": close * 1e6, + } + ) + + +_MA_CROSS = "快: MA(C, 5);\n慢: MA(C, 20);\n买入: CROSS(快, 慢);\n卖出: CROSS(慢, 快);" + + +# ── attach / pick ───────────────────────────────────────────────────────────── + + +def test_attach_formula_columns(): + df = _df(100) + enriched, result = attach_formula_columns(df, compile_formula(_MA_CROSS)) + assert {"快", "慢", "买入", "卖出"} <= set(enriched.columns) + assert len(enriched) == len(df) + assert df is not enriched # 副本,不污染原 df + + +def test_pick_signal_columns_by_hint_and_order(): + _, result = attach_formula_columns(_df(60), compile_formula(_MA_CROSS)) + buy, sell = pick_signal_columns(result) + assert (buy, sell) == ("买入", "卖出") # 名称提示(买/卖)优先 + + _, r2 = attach_formula_columns(_df(60), compile_formula("A: C > MA(C, 5); B: C < MA(C, 5);")) + buy2, sell2 = pick_signal_columns(r2) + assert (buy2, sell2) == ("A", "B") # 无提示时按声明顺序 + + buy3, _ = pick_signal_columns(r2, buy_col="B") + assert buy3 == "B" # 显式指定优先 + + +def test_pick_requires_signal(): + from easy_tdx.formula import FormulaResult + + result = FormulaResult(columns={"x": np.ones(5)}, values=["x"]) + with pytest.raises(FormulaStrategyError, match="布尔信号"): + pick_signal_columns(result) + + +# ── run_formula_backtest ────────────────────────────────────────────────────── + + +def test_run_formula_backtest_full_report(): + out = run_formula_backtest(_df(300), _MA_CROSS) + assert out["performance"]["total_trades"] >= 1 + assert out["formula"]["buy_col"] == "买入" + assert out["formula"]["sell_col"] == "卖出" + assert out["grade"]["grade"] in ("S", "A", "B", "C", "D") + assert 0 <= out["score"]["total"] <= 100 + assert "trades" in out and "equity_curve" in out + + +def test_run_formula_backtest_no_sell_col_holds(): + """只有买入列 → 买入后持有到末尾(1 笔完成交易=0 卖出,持仓中)。""" + out = run_formula_backtest(_df(200, drift=0.004), "买入: CROSS(MA(C,3), MA(C,30));") + assert out["formula"]["sell_col"] is None + assert out["performance"]["total_return"] > 0 + + +def test_run_formula_backtest_accepts_compiled(): + compiled = compile_formula(_MA_CROSS) + out = run_formula_backtest(_df(200), compiled) + assert out["formula"]["buy_col"] == "买入" + + +def test_run_formula_backtest_rejects_no_signal(): + with pytest.raises(ValueError, match="布尔信号"): + run_formula_backtest(_df(60), "数值: MA(C, 5);") + + +def test_run_formula_backtest_json_serializable(): + import json + + out = run_formula_backtest(_df(150), _MA_CROSS) + json.dumps(out, default=str) + + +# ── REST 端点 ───────────────────────────────────────────────────────────────── + + +def _client() -> TestClient: + from easy_tdx.web import create_app + + return TestClient(create_app()) + + +def _ohlcv(n: int = 200) -> list[dict[str, object]]: + df = _df(n) + df["datetime"] = df["datetime"].dt.strftime("%Y-%m-%d") + return json_records(df) + + +def json_records(df: pd.DataFrame) -> list[dict[str, object]]: + import json + + return json.loads(df.to_json(orient="records", force_ascii=False)) + + +def test_rest_formula_validate_ok_and_error(): + client = _client() + r = client.post( + "/api/v1/formula/validate", json={"text": "金叉: CROSS(MA(C,5), MA(C,20)); 强度: MA(C,5);"} + ) + assert r.status_code == 200 + body = r.json() + assert body["ok"] is True + assert body["signals"] == ["金叉"] + assert body["values"] == ["强度"] + + r2 = client.post("/api/v1/formula/validate", json={"text": "A := ;"}) + assert r2.status_code == 200 + assert r2.json()["ok"] is False + assert r2.json()["error"] + + +def test_rest_formula_compute_inline_ohlcv(): + client = _client() + r = client.post( + "/api/v1/formula/compute", + json={"text": "买入: C > REF(C, 1); 值: MA(C, 5);", "ohlcv": _ohlcv(100), "tail": 5}, + ) + assert r.status_code == 200, r.text + body = r.json() + assert body["signals"] == ["买入"] + assert "last_row" in body and "值" in body["last_row"] + assert len(body["recent"]) == 5 + + +def test_rest_formula_backtest_async_task(): + client = _client() + r = client.post( + "/api/v1/formula/backtest/run/async", + json={"text": _MA_CROSS, "ohlcv": _ohlcv(300), "cash": 100000.0}, + ) + assert r.status_code == 202, r.text + task_id = r.json()["task_id"] + for _ in range(200): + st = client.get(f"/api/v1/backtest/tasks/{task_id}").json() + if st["status"] in ("done", "failed"): + break + time.sleep(0.05) + assert st["status"] == "done", st.get("error") + result = st["result"] + assert result["formula"]["buy_col"] == "买入" + assert result["performance"]["total_trades"] >= 1 + + +def test_rest_formula_screen_async_task(): + client = _client() + # 两份不同行情:A 上涨(末根 C>REF(C,1) 大概率真)、B 构造末根下跌 + up = _ohlcv(120) + r = client.post( + "/api/v1/formula/screen/run/async", + json={"text": "买入: C > REF(C, 1);", "symbols": ["SH:600519"], "ohlcv": up[:0]}, + ) + # symbols 路径需要行情连接——离线环境预期 400/500(无 mock client) + # 这里只验证请求校验(symbols 非空)不炸 + assert r.status_code in (400, 500, 202) diff --git a/web-ui/src/api.ts b/web-ui/src/api.ts index 9aaff88..1dd4c5b 100644 --- a/web-ui/src/api.ts +++ b/web-ui/src/api.ts @@ -151,6 +151,33 @@ export async function submitBacktestTask(req: BacktestRequest): Promise { + const resp = await fetch(`${BASE}/backtest/wf/run/async?n_windows=${nWindows}`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(req), + }) + if (!resp.ok) await throwError(resp) + return (await resp.json()) as TaskSubmitResponse +} + +/** 提交一条龙评估后台任务(回测+WF+适配性+评分+基准对比,v1.27)。 */ +export async function submitEvaluateTask( + req: BacktestRequest, +): Promise { + const resp = await fetch(`${BASE}/backtest/evaluate/run/async`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(req), + }) + if (!resp.ok) await throwError(resp) + return (await resp.json()) as TaskSubmitResponse +} + /** 提交组合回测后台任务,返回 task_id。 */ export async function submitPortfolioTask( req: PortfolioBacktestRequest, diff --git a/web-ui/src/components/EvaluatePanel.vue b/web-ui/src/components/EvaluatePanel.vue new file mode 100644 index 0000000..8a78e20 --- /dev/null +++ b/web-ui/src/components/EvaluatePanel.vue @@ -0,0 +1,255 @@ + + + + + diff --git a/web-ui/src/components/WalkForwardPanel.vue b/web-ui/src/components/WalkForwardPanel.vue new file mode 100644 index 0000000..9854118 --- /dev/null +++ b/web-ui/src/components/WalkForwardPanel.vue @@ -0,0 +1,168 @@ + + + + + diff --git a/web-ui/src/stores/backtest.ts b/web-ui/src/stores/backtest.ts index 2fec6fc..cb3f055 100644 --- a/web-ui/src/stores/backtest.ts +++ b/web-ui/src/stores/backtest.ts @@ -12,6 +12,8 @@ import { submitOptimizeAllTask, submitOptimizeTask, submitMultiStrategyTask, + submitWalkforwardTask, + submitEvaluateTask, fetchTask, } from '../api' import type { @@ -19,6 +21,7 @@ import type { BacktestResult, Bar, Category, + EvaluateReport, MultiStrategyBacktestRequest, PortfolioBacktestRequest, PortfolioResult, @@ -27,8 +30,22 @@ import type { OptimizeBacktestRequest, OptimizeResult, StrategySchema, + WalkForwardResult, } from '../types' +/** 轮询后台任务直到终态(done 返回 result,failed/超时抛错)。 */ +async function pollTask(taskId: string, timeoutMs: number, what: string): Promise { + const start = Date.now() + // eslint-disable-next-line no-constant-condition + while (true) { + const state = await fetchTask(taskId) + if (state.status === 'done' && state.result) return state.result as T + if (state.status === 'failed') throw new Error(state.error || `${what}失败`) + if (Date.now() - start > timeoutMs) throw new Error(`${what}超时(${timeoutMs / 1000}s)`) + await new Promise((r) => setTimeout(r, 400)) + } +} + export const useBacktestStore = defineStore('backtest', () => { // ── 策略 ───────────────────────────────────────────────────────────────── const strategies = ref([]) @@ -81,6 +98,56 @@ export const useBacktestStore = defineStore('backtest', () => { error.value = '' } + // ── 附加分析:Walk-Forward / 一条龙评估(v1.27) ───────────────────────── + const wfResult = ref(null) + const wfRunning = ref(false) + const wfError = ref('') + const evaluateResult = ref(null) + const evaluateRunning = ref(false) + const evaluateError = ref('') + + /** 提交 WF 样本外验证后台任务并轮询(与主回测共用同一份内联 OHLCV)。 */ + async function runWalkforward(req: Omit, nWindows = 7) { + if (!hasBars.value) return + wfRunning.value = true + wfError.value = '' + wfResult.value = null + try { + const { task_id } = await submitWalkforwardTask({ ...req, ohlcv: ohlcv.value }, nWindows) + const body = await pollTask<{ walkforward: WalkForwardResult }>(task_id, 180_000, 'WF 验证') + wfResult.value = body.walkforward + } catch (e) { + wfError.value = formatError(e) + wfResult.value = null + } finally { + wfRunning.value = false + } + } + + /** 提交一条龙评估后台任务并轮询(回测+WF+适配性+评分+基准对比)。 */ + async function runEvaluate(req: Omit) { + if (!hasBars.value) return + evaluateRunning.value = true + evaluateError.value = '' + evaluateResult.value = null + try { + const { task_id } = await submitEvaluateTask({ ...req, ohlcv: ohlcv.value }) + evaluateResult.value = await pollTask(task_id, 300_000, '一条龙评估') + } catch (e) { + evaluateError.value = formatError(e) + evaluateResult.value = null + } finally { + evaluateRunning.value = false + } + } + + function clearExtraAnalysis() { + wfResult.value = null + wfError.value = '' + evaluateResult.value = null + evaluateError.value = '' + } + // ── 组合回测(Phase 3) ─────────────────────────────────────────────────── const portfolioResult = ref(null) const portfolioRunning = ref(false) @@ -259,6 +326,12 @@ export const useBacktestStore = defineStore('backtest', () => { optimizeContext, optimizeAllResult, optimizeAllRunning, + wfResult, + wfRunning, + wfError, + evaluateResult, + evaluateRunning, + evaluateError, // getters hasBars, // actions @@ -266,6 +339,9 @@ export const useBacktestStore = defineStore('backtest', () => { setOhlcv, run, clearResult, + runWalkforward, + runEvaluate, + clearExtraAnalysis, runPortfolio, clearPortfolio, runMultiStrategy, diff --git a/web-ui/src/types.ts b/web-ui/src/types.ts index 7ba3974..b3241b1 100644 --- a/web-ui/src/types.ts +++ b/web-ui/src/types.ts @@ -137,6 +137,8 @@ export interface TaskState { | OptimizeResult | OptimizeAllResult | SignalScanResult + | WalkForwardResult + | EvaluateReport | null error: string | null description: string @@ -520,3 +522,94 @@ export interface RankRow { change_pct?: number [key: string]: unknown } + +// ── Walk-Forward 样本外验证(v1.27 POST /backtest/wf/run/async)────────────── + +export interface WalkForwardWindow { + index: number + start: string + end: string + bars: number + total_return: number + sharpe: number + max_drawdown: number + total_trades: number + win_rate: number +} + +export interface WalkForwardResult { + n_windows: number + warmup_ratio: number + windows: WalkForwardWindow[] + /** 盈利窗占比(0~1,时间稳定性核心指标) */ + consistency: number + /** 各窗收益连乘 - 1 */ + chained_return: number + mean_window_return: number + median_window_return: number + worst_window: number + best_window: number + mean_sharpe: number + worst_drawdown: number + total_trades: number +} + +// ── 一条龙评估(v1.27 POST /backtest/evaluate/run/async)───────────────────── + +export interface FitnessCheckRow { + name: string + passed: boolean + detail: string +} + +export interface FitnessSegmentRow { + name: string + start: string + end: string + bars: number + total_return: number + sharpe: number + max_drawdown: number + total_trades: number + win_rate: number +} + +export interface FitnessReport { + segments: FitnessSegmentRow[] + checks: FitnessCheckRow[] + pass_ratio: number + passed_count: number + total_checks: number + high_fitness: boolean + split: number[] +} + +/** 综合评分(0-100 加权:收益50/夏普15/回撤10/Sortino5/WF一致性20) */ +export interface StrategyScoreReport { + total: number + components: Record + weights_used: Record + wf_provided: boolean +} + +export interface EvaluateBenchmarkReport { + buy_hold: { + total_return: number + annual_return: number + max_drawdown: number + sharpe: number + calmar: number + volatility: number + } + /** 策略总收益 - 买入持有总收益 */ + excess_return: number +} + +export interface EvaluateReport { + performance: Performance + score: StrategyScoreReport + walkforward: WalkForwardResult + fitness: FitnessReport + benchmark: EvaluateBenchmarkReport + config: Record +} diff --git a/web-ui/src/views/BacktestView.vue b/web-ui/src/views/BacktestView.vue index e9e56fa..6e041ec 100644 --- a/web-ui/src/views/BacktestView.vue +++ b/web-ui/src/views/BacktestView.vue @@ -7,12 +7,14 @@ import { computed, nextTick, onMounted, ref } from 'vue' import { useRoute } from 'vue-router' import EquityChart from '../components/EquityChart.vue' +import EvaluatePanel from '../components/EvaluatePanel.vue' import GradeDetails from '../components/GradeDetails.vue' import KlineChart from '../components/KlineChart.vue' import MetricTable from '../components/MetricTable.vue' import StrategyPicker from '../components/StrategyPicker.vue' import SymbolPicker from '../components/SymbolPicker.vue' import TradeTable from '../components/TradeTable.vue' +import WalkForwardPanel from '../components/WalkForwardPanel.vue' import { formatError, saveStrategy } from '../api' import { detectMarket } from '../market' import { gradePerformance } from '../grading' @@ -85,21 +87,35 @@ onMounted(async () => { if (qCategory) category.value = qCategory }) +// 附加分析开关(v1.27):WF 样本外验证 / 一条龙评估, +// 勾选后随「开始回测」一起提交(与主回测共用同一份内联 OHLCV)。 +const wfEnabled = ref(false) +const wfWindows = ref(7) +const evaluateEnabled = ref(false) + // 取行情 + 回测 串联(点击「开始回测」触发) async function onRun() { store.error = '' + store.clearExtraAnalysis() // 1. 先取行情(SymbolPicker.loadBars 会校验并填充 store.ohlcv) const ok = await symbolPicker.value?.loadBars() if (!ok) return // 校验/取数失败,错误已在 store.error // 2. 再回测 - await store.run({ + const req = { strategy: strategy.value, params: params.value, cash: cash.value, commission: commission.value, slippage: slippage.value, execution: execution.value, - }) + } + await store.run(req) + // 3. 附加分析:勾选的 WF / 一条龙评估并行跑(互不阻塞,各自有独立错误提示) + if (!store.result) return + const jobs: Promise[] = [] + if (wfEnabled.value) jobs.push(store.runWalkforward(req, wfWindows.value)) + if (evaluateEnabled.value) jobs.push(store.runEvaluate(req)) + await Promise.allSettled(jobs) } // ── 保存策略(把当前结果 + 配置 + 上下文存进策略库)────────────────────────── @@ -235,12 +251,32 @@ async function onSave() { +
+

附加分析

+ +
+ + +
+ +

勾选后随「开始回测」自动附加运行

+
+ @@ -273,6 +309,30 @@ async function onSave() { + +
+

Walk-Forward 样本外验证

+

验证中…(逐窗独立回测,约需数秒)

+
⚠ {{ store.wfError }}
+ +
+ + +
+

一条龙评估

+

+ 评估中…(回测 + WF + 适配性 + 基准对比,约需数秒) +

+
⚠ {{ store.evaluateError }}
+ +
+

绩效指标

@@ -354,6 +414,37 @@ async function onSave() { color: var(--text-dim); font-size: 12px; } +/* 附加分析开关 */ +.check-row { + display: flex; + align-items: center; + gap: 6px; + font-size: 12px; + color: var(--text); + cursor: pointer; + margin-bottom: 8px; +} +.check-row input { + accent-color: var(--accent, #4a9eff); +} +.wf-windows { + margin: -2px 0 8px 20px; + max-width: 110px; +} +.wf-windows input { + width: 100%; + background: var(--bg); + border: 1px solid var(--border); + border-radius: var(--radius); + padding: 5px 8px; + font-size: 12px; + color: var(--text); +} +.extra-hint { + font-size: 11px; + color: var(--text-dim); + margin: 2px 0 0; +} .run-btn { margin-top: auto; width: 100%;