mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 13:24:18 +08:00
release: v1.27.0 — 通达信公式解析器三通道 + 轮动组合引擎 + 回测页WF/评估开关 + Docker 部署
升级计划 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 向量化
This commit is contained in:
@@ -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 原生提供。
|
||||
|
||||
+16
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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 字段),不静默。
|
||||
+1
-1
@@ -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"
|
||||
|
||||
@@ -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 全部通过"
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -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())
|
||||
@@ -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)
|
||||
|
||||
@@ -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))
|
||||
@@ -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<ws>\s+)
|
||||
| (?P<comment>\{[^}]*\})
|
||||
| (?P<num>\d+\.\d+|\.\d+|\d+)
|
||||
| (?P<name>[A-Za-z_\u4e00-\u9fff][A-Za-z0-9_\u4e00-\u9fff]*)
|
||||
| (?P<op>:=|>=|<=|==|&&|\|\||[-+*/(),:;><!=])
|
||||
""",
|
||||
re.VERBOSE,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Token:
|
||||
kind: str # num / name / op / eof
|
||||
value: str
|
||||
pos: int
|
||||
|
||||
|
||||
def _tokenize(text: str) -> 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)
|
||||
@@ -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")
|
||||
|
||||
@@ -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}
|
||||
@@ -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
|
||||
@@ -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
|
||||
)
|
||||
@@ -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)
|
||||
@@ -151,6 +151,33 @@ export async function submitBacktestTask(req: BacktestRequest): Promise<TaskSubm
|
||||
return (await resp.json()) as TaskSubmitResponse
|
||||
}
|
||||
|
||||
/** 提交 Walk-Forward 样本外验证后台任务(v1.27,n_windows 默认 7)。 */
|
||||
export async function submitWalkforwardTask(
|
||||
req: BacktestRequest,
|
||||
nWindows = 7,
|
||||
): Promise<TaskSubmitResponse> {
|
||||
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<TaskSubmitResponse> {
|
||||
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,
|
||||
|
||||
@@ -0,0 +1,255 @@
|
||||
<script setup lang="ts">
|
||||
// 一条龙评估面板:综合评分 + 适配性体检 + 买入持有基准对比。
|
||||
// 评级复用前端 grading(与后端口径对拍一致),评分/适配性/基准来自后端报告。
|
||||
|
||||
import { computed } from 'vue'
|
||||
|
||||
import GradeDetails from './GradeDetails.vue'
|
||||
import { gradePerformance } from '../grading'
|
||||
import type { EvaluateReport } from '../types'
|
||||
|
||||
const props = defineProps<{
|
||||
report: EvaluateReport
|
||||
}>()
|
||||
|
||||
/** 综合评分分项(含权重,展示顺序固定) */
|
||||
const scoreComponents = computed(() => {
|
||||
const labels: Record<string, string> = {
|
||||
total_return: '收益',
|
||||
sharpe: '夏普',
|
||||
max_drawdown: '回撤',
|
||||
sortino: '索提诺',
|
||||
wf_consistency: 'WF一致性',
|
||||
}
|
||||
return Object.entries(props.report.score.components).map(([key, v]) => ({
|
||||
key,
|
||||
label: labels[key] ?? key,
|
||||
value: v,
|
||||
weight: props.report.score.weights_used[key] ?? 0,
|
||||
}))
|
||||
})
|
||||
|
||||
const grade = computed(() => gradePerformance(props.report.performance))
|
||||
|
||||
const excess = computed(() => props.report.benchmark.excess_return)
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<div class="eval-panel">
|
||||
<!-- 顶部:综合评分 + 高适配 + 超额收益 -->
|
||||
<div class="eval-header">
|
||||
<div class="score-block">
|
||||
<span class="score-label">综合评分</span>
|
||||
<span class="score-value" :class="report.score.total >= 60 ? 'pos' : 'neg'">
|
||||
{{ report.score.total.toFixed(1) }}
|
||||
</span>
|
||||
<span class="score-unit">/100</span>
|
||||
</div>
|
||||
<div class="badge" :class="report.fitness.high_fitness ? 'badge-ok' : 'badge-warn'">
|
||||
{{ report.fitness.high_fitness ? '✓ 高适配' : '△ 适配性未达标' }}
|
||||
</div>
|
||||
<div class="excess-block">
|
||||
<span class="score-label">对比买入持有</span>
|
||||
<span class="score-value" :class="excess >= 0 ? 'pos' : 'neg'">
|
||||
{{ excess >= 0 ? '+' : '' }}{{ (excess * 100).toFixed(2) }}%
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- 分项评分条 -->
|
||||
<div class="score-components">
|
||||
<div v-for="c in scoreComponents" :key="c.key" class="comp">
|
||||
<span class="comp-label">{{ c.label }}<small> ×{{ (c.weight * 100).toFixed(0) }}%</small></span>
|
||||
<div class="comp-bar">
|
||||
<div class="comp-fill" :style="{ width: `${Math.min(c.value, 100)}%` }"></div>
|
||||
</div>
|
||||
<span class="comp-value">{{ c.value.toFixed(0) }}</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- 基准对比 -->
|
||||
<div class="bench-row">
|
||||
<div class="bench-cell">
|
||||
<span class="stat-label">策略总收益</span>
|
||||
<span :class="report.performance.total_return >= 0 ? 'pos' : 'neg'" class="mono">
|
||||
{{ (report.performance.total_return * 100).toFixed(2) }}%
|
||||
</span>
|
||||
</div>
|
||||
<div class="bench-cell">
|
||||
<span class="stat-label">买入持有</span>
|
||||
<span :class="report.benchmark.buy_hold.total_return >= 0 ? 'pos' : 'neg'" class="mono">
|
||||
{{ (report.benchmark.buy_hold.total_return * 100).toFixed(2) }}%
|
||||
</span>
|
||||
</div>
|
||||
<div class="bench-cell">
|
||||
<span class="stat-label">超额收益</span>
|
||||
<span :class="excess >= 0 ? 'pos' : 'neg'" class="mono">
|
||||
{{ excess >= 0 ? '+' : '' }}{{ (excess * 100).toFixed(2) }}%
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
<p v-if="excess < 0" class="bench-warn">⚠ 策略跑输同区间买入持有——研发阶段的一票否决信号。</p>
|
||||
|
||||
<!-- 适配性检查(8 项可解释) -->
|
||||
<h4 class="sub-title">
|
||||
适配性体检 {{ report.fitness.passed_count }}/{{ report.fitness.total_checks }}
|
||||
(train/valid/test = {{ report.fitness.split.map((s) => (s * 100).toFixed(0)).join('/') }})
|
||||
</h4>
|
||||
<ul class="check-list">
|
||||
<li v-for="c in report.fitness.checks" :key="c.name" :class="c.passed ? 'ok' : 'bad'">
|
||||
<span class="check-mark">{{ c.passed ? '✓' : '✗' }}</span>
|
||||
<span class="check-detail">{{ c.detail }}</span>
|
||||
</li>
|
||||
</ul>
|
||||
|
||||
<!-- 评级(复用本地评级,与后端字段口径一致) -->
|
||||
<h4 class="sub-title">评级(不看收益率)</h4>
|
||||
<GradeDetails :result="grade" />
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.eval-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 16px;
|
||||
flex-wrap: wrap;
|
||||
margin-bottom: 12px;
|
||||
}
|
||||
.score-block,
|
||||
.excess-block {
|
||||
display: flex;
|
||||
align-items: baseline;
|
||||
gap: 6px;
|
||||
}
|
||||
.score-label {
|
||||
font-size: 12px;
|
||||
color: var(--text-dim);
|
||||
}
|
||||
.score-value {
|
||||
font-size: 26px;
|
||||
font-weight: 700;
|
||||
font-family: var(--font-mono);
|
||||
}
|
||||
.score-unit {
|
||||
font-size: 12px;
|
||||
color: var(--text-dim);
|
||||
}
|
||||
.pos {
|
||||
color: var(--up);
|
||||
}
|
||||
.neg {
|
||||
color: #2ebd85;
|
||||
}
|
||||
.badge {
|
||||
font-size: 12px;
|
||||
padding: 4px 10px;
|
||||
border-radius: 10px;
|
||||
font-weight: 600;
|
||||
}
|
||||
.badge-ok {
|
||||
background: rgba(46, 189, 133, 0.15);
|
||||
color: #2ebd85;
|
||||
border: 1px solid #2ebd85;
|
||||
}
|
||||
.badge-warn {
|
||||
background: rgba(239, 65, 70, 0.12);
|
||||
color: var(--up);
|
||||
border: 1px solid var(--up);
|
||||
}
|
||||
.score-components {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fit, minmax(180px, 1fr));
|
||||
gap: 6px 14px;
|
||||
margin-bottom: 14px;
|
||||
}
|
||||
.comp {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
font-size: 12px;
|
||||
}
|
||||
.comp-label {
|
||||
width: 86px;
|
||||
color: var(--text-muted);
|
||||
flex-shrink: 0;
|
||||
}
|
||||
.comp-label small {
|
||||
color: var(--text-dim);
|
||||
}
|
||||
.comp-bar {
|
||||
flex: 1;
|
||||
height: 6px;
|
||||
background: var(--bg);
|
||||
border-radius: 3px;
|
||||
overflow: hidden;
|
||||
}
|
||||
.comp-fill {
|
||||
height: 100%;
|
||||
background: var(--accent, #4a9eff);
|
||||
border-radius: 3px;
|
||||
}
|
||||
.comp-value {
|
||||
width: 28px;
|
||||
text-align: right;
|
||||
font-family: var(--font-mono);
|
||||
color: var(--text-muted);
|
||||
}
|
||||
.bench-row {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(3, 1fr);
|
||||
gap: 8px;
|
||||
margin-bottom: 6px;
|
||||
}
|
||||
.bench-cell {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 2px;
|
||||
background: var(--bg);
|
||||
border: 1px solid var(--border);
|
||||
border-radius: var(--radius);
|
||||
padding: 8px 10px;
|
||||
}
|
||||
.mono {
|
||||
font-family: var(--font-mono);
|
||||
font-weight: 600;
|
||||
}
|
||||
.bench-warn {
|
||||
font-size: 12px;
|
||||
color: var(--up);
|
||||
margin: 4px 0 12px;
|
||||
}
|
||||
.sub-title {
|
||||
font-size: 12px;
|
||||
color: var(--text-muted);
|
||||
margin: 14px 0 8px;
|
||||
}
|
||||
.check-list {
|
||||
list-style: none;
|
||||
padding: 0;
|
||||
margin: 0 0 8px;
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fit, minmax(280px, 1fr));
|
||||
gap: 4px 14px;
|
||||
}
|
||||
.check-list li {
|
||||
display: flex;
|
||||
gap: 6px;
|
||||
font-size: 12px;
|
||||
align-items: baseline;
|
||||
}
|
||||
.check-mark {
|
||||
font-weight: 700;
|
||||
width: 14px;
|
||||
flex-shrink: 0;
|
||||
}
|
||||
.check-list li.ok .check-mark {
|
||||
color: #2ebd85;
|
||||
}
|
||||
.check-list li.bad .check-mark {
|
||||
color: var(--up);
|
||||
}
|
||||
.check-detail {
|
||||
color: var(--text-muted);
|
||||
}
|
||||
</style>
|
||||
@@ -0,0 +1,168 @@
|
||||
<script setup lang="ts">
|
||||
// Walk-Forward 样本外验证面板:逐窗收益柱状图 + 稳定性汇总。
|
||||
// 柱色按 A 股惯例:红涨绿跌。
|
||||
|
||||
import { onBeforeUnmount, onMounted, ref, watch } from 'vue'
|
||||
|
||||
import echarts from '../echarts-setup'
|
||||
import type { WalkForwardResult } from '../types'
|
||||
|
||||
const props = defineProps<{
|
||||
wf: WalkForwardResult
|
||||
}>()
|
||||
|
||||
const container = ref<HTMLDivElement>()
|
||||
let chart: echarts.ECharts | null = null
|
||||
|
||||
const nWin = ref(props.wf.windows.filter((w) => w.total_return > 0).length)
|
||||
|
||||
function render() {
|
||||
if (!container.value || props.wf.windows.length === 0) return
|
||||
chart ??= echarts.init(container.value, 'dark')
|
||||
chart.setOption(buildOption(), true)
|
||||
}
|
||||
|
||||
function buildOption(): echarts.EChartsCoreOption {
|
||||
const labels = props.wf.windows.map((w) => `窗${w.index + 1}\n${w.start.slice(5)}~${w.end.slice(5)}`)
|
||||
const values = props.wf.windows.map((w) => +(w.total_return * 100).toFixed(2))
|
||||
|
||||
return {
|
||||
backgroundColor: 'transparent',
|
||||
tooltip: {
|
||||
trigger: 'axis',
|
||||
valueFormatter: (v: number | string) => `${Number(v).toFixed(2)}%`,
|
||||
},
|
||||
grid: { left: '6%', right: '3%', top: 20, bottom: 40 },
|
||||
xAxis: {
|
||||
type: 'category',
|
||||
data: labels,
|
||||
axisLabel: { fontSize: 10, interval: 0 },
|
||||
},
|
||||
yAxis: {
|
||||
type: 'value',
|
||||
axisLabel: { formatter: (v: number) => `${v}%` },
|
||||
splitLine: { lineStyle: { color: '#2a2e3a' } },
|
||||
},
|
||||
series: [
|
||||
{
|
||||
name: '窗口收益',
|
||||
type: 'bar',
|
||||
data: values.map((v) => ({
|
||||
value: v,
|
||||
itemStyle: { color: v >= 0 ? '#ef4146' : '#2ebd85' }, // 红涨绿跌
|
||||
})),
|
||||
barMaxWidth: 36,
|
||||
label: {
|
||||
show: true,
|
||||
position: 'top',
|
||||
fontSize: 10,
|
||||
formatter: (p: { value: number }) => `${p.value.toFixed(1)}%`,
|
||||
},
|
||||
},
|
||||
],
|
||||
}
|
||||
}
|
||||
|
||||
function resize() {
|
||||
chart?.resize()
|
||||
}
|
||||
|
||||
onMounted(() => {
|
||||
render()
|
||||
window.addEventListener('resize', resize)
|
||||
})
|
||||
onBeforeUnmount(() => {
|
||||
window.removeEventListener('resize', resize)
|
||||
chart?.dispose()
|
||||
chart = null
|
||||
})
|
||||
watch(() => props.wf, render)
|
||||
watch(
|
||||
() => props.wf.windows,
|
||||
(ws) => {
|
||||
nWin.value = ws.filter((w) => w.total_return > 0).length
|
||||
},
|
||||
)
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<div class="wf-panel">
|
||||
<div class="wf-summary">
|
||||
<div class="stat">
|
||||
<span class="stat-label">盈利窗占比</span>
|
||||
<span class="stat-value" :class="wf.consistency >= 0.5 ? 'pos' : 'neg'">
|
||||
{{ nWin }}/{{ wf.windows.length }}({{ (wf.consistency * 100).toFixed(0) }}%)
|
||||
</span>
|
||||
</div>
|
||||
<div class="stat">
|
||||
<span class="stat-label">连乘收益</span>
|
||||
<span class="stat-value" :class="wf.chained_return >= 0 ? 'pos' : 'neg'">
|
||||
{{ (wf.chained_return * 100).toFixed(2) }}%
|
||||
</span>
|
||||
</div>
|
||||
<div class="stat">
|
||||
<span class="stat-label">最差窗</span>
|
||||
<span class="stat-value neg">{{ (wf.worst_window * 100).toFixed(2) }}%</span>
|
||||
</div>
|
||||
<div class="stat">
|
||||
<span class="stat-label">最好窗</span>
|
||||
<span class="stat-value pos">{{ (wf.best_window * 100).toFixed(2) }}%</span>
|
||||
</div>
|
||||
<div class="stat">
|
||||
<span class="stat-label">平均夏普</span>
|
||||
<span class="stat-value">{{ wf.mean_sharpe.toFixed(2) }}</span>
|
||||
</div>
|
||||
<div class="stat">
|
||||
<span class="stat-label">总交易</span>
|
||||
<span class="stat-value">{{ wf.total_trades }} 笔</span>
|
||||
</div>
|
||||
</div>
|
||||
<p class="wf-hint">
|
||||
前 {{ (wf.warmup_ratio * 100).toFixed(0) }}% 为预热区不参与评估;每窗独立开仓(窗口起点空仓),
|
||||
稳健策略的盈利窗占比应 ≥ 50%。
|
||||
</p>
|
||||
<div ref="container" class="wf-chart"></div>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.wf-summary {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fit, minmax(120px, 1fr));
|
||||
gap: 8px;
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.stat {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 2px;
|
||||
background: var(--bg);
|
||||
border: 1px solid var(--border);
|
||||
border-radius: var(--radius);
|
||||
padding: 8px 10px;
|
||||
}
|
||||
.stat-label {
|
||||
font-size: 11px;
|
||||
color: var(--text-dim);
|
||||
}
|
||||
.stat-value {
|
||||
font-size: 14px;
|
||||
font-weight: 600;
|
||||
font-family: var(--font-mono);
|
||||
}
|
||||
.stat-value.pos {
|
||||
color: var(--up);
|
||||
}
|
||||
.stat-value.neg {
|
||||
color: #2ebd85;
|
||||
}
|
||||
.wf-hint {
|
||||
font-size: 11px;
|
||||
color: var(--text-dim);
|
||||
margin: 0 0 8px;
|
||||
}
|
||||
.wf-chart {
|
||||
width: 100%;
|
||||
height: 280px;
|
||||
}
|
||||
</style>
|
||||
@@ -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<T>(taskId: string, timeoutMs: number, what: string): Promise<T> {
|
||||
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<StrategySchema[]>([])
|
||||
@@ -81,6 +98,56 @@ export const useBacktestStore = defineStore('backtest', () => {
|
||||
error.value = ''
|
||||
}
|
||||
|
||||
// ── 附加分析:Walk-Forward / 一条龙评估(v1.27) ─────────────────────────
|
||||
const wfResult = ref<WalkForwardResult | null>(null)
|
||||
const wfRunning = ref(false)
|
||||
const wfError = ref<string>('')
|
||||
const evaluateResult = ref<EvaluateReport | null>(null)
|
||||
const evaluateRunning = ref(false)
|
||||
const evaluateError = ref<string>('')
|
||||
|
||||
/** 提交 WF 样本外验证后台任务并轮询(与主回测共用同一份内联 OHLCV)。 */
|
||||
async function runWalkforward(req: Omit<BacktestRequest, 'ohlcv'>, 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<BacktestRequest, 'ohlcv'>) {
|
||||
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<EvaluateReport>(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<PortfolioResult | null>(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,
|
||||
|
||||
@@ -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<string, number>
|
||||
weights_used: Record<string, number>
|
||||
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<string, unknown>
|
||||
}
|
||||
|
||||
@@ -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<void>[] = []
|
||||
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() {
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section class="panel-section">
|
||||
<h3>附加分析</h3>
|
||||
<label class="check-row" title="把时间轴切 7 窗独立回测,检验跨时段稳定性(每窗独立开仓)">
|
||||
<input v-model="wfEnabled" type="checkbox" />
|
||||
Walk-Forward 样本外验证
|
||||
</label>
|
||||
<div v-if="wfEnabled" class="field wf-windows">
|
||||
<label>窗口数</label>
|
||||
<input v-model.number="wfWindows" type="number" min="2" max="12" step="1" />
|
||||
</div>
|
||||
<label
|
||||
class="check-row"
|
||||
title="回测+WF+适配性体检+综合评分+买入持有基准对比,一份报告"
|
||||
>
|
||||
<input v-model="evaluateEnabled" type="checkbox" />
|
||||
一条龙评估
|
||||
</label>
|
||||
<p class="extra-hint">勾选后随「开始回测」自动附加运行</p>
|
||||
</section>
|
||||
|
||||
<button
|
||||
class="primary run-btn"
|
||||
:disabled="store.running"
|
||||
:disabled="store.running || store.wfRunning || store.evaluateRunning"
|
||||
@click="onRun"
|
||||
>
|
||||
{{ store.running ? '取行情+回测中…' : '开始回测' }}
|
||||
{{ store.running || store.wfRunning || store.evaluateRunning ? '取行情+回测中…' : '开始回测' }}
|
||||
</button>
|
||||
</aside>
|
||||
|
||||
@@ -273,6 +309,30 @@ async function onSave() {
|
||||
<GradeDetails :result="grade" expanded />
|
||||
</section>
|
||||
|
||||
<!-- 附加分析:WF 样本外验证(v1.27) -->
|
||||
<section
|
||||
v-if="store.wfRunning || store.wfResult || store.wfError"
|
||||
class="report-section"
|
||||
>
|
||||
<h3>Walk-Forward 样本外验证</h3>
|
||||
<p v-if="store.wfRunning" class="loading-text">验证中…(逐窗独立回测,约需数秒)</p>
|
||||
<div v-else-if="store.wfError" class="error-banner">⚠ {{ store.wfError }}</div>
|
||||
<WalkForwardPanel v-else-if="store.wfResult" :wf="store.wfResult" />
|
||||
</section>
|
||||
|
||||
<!-- 附加分析:一条龙评估(v1.27) -->
|
||||
<section
|
||||
v-if="store.evaluateRunning || store.evaluateResult || store.evaluateError"
|
||||
class="report-section"
|
||||
>
|
||||
<h3>一条龙评估</h3>
|
||||
<p v-if="store.evaluateRunning" class="loading-text">
|
||||
评估中…(回测 + WF + 适配性 + 基准对比,约需数秒)
|
||||
</p>
|
||||
<div v-else-if="store.evaluateError" class="error-banner">⚠ {{ store.evaluateError }}</div>
|
||||
<EvaluatePanel v-else-if="store.evaluateResult" :report="store.evaluateResult" />
|
||||
</section>
|
||||
|
||||
<section class="report-section">
|
||||
<h3>绩效指标</h3>
|
||||
<MetricTable :perf="store.result.performance" />
|
||||
@@ -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%;
|
||||
|
||||
Reference in New Issue
Block a user