release: v1.32.6 — 两周改动深度审查全面修复(回测口径三件套/LLM 安全加固/涨停价舍入/时区统一/缓存与竞态等 58 处)

对 v1.21→v1.32.5 的 249 文件 4.2 万行改动做六路专项审查,本轮落地全部发现:

回测正确性:组合收益 fillna(0) 虚增、轮动停牌日过期价成交、单标的 WF 逐窗指标
被预热区稀释(三件套均带先红后绿回归);worst_drawdown 方向、grading 容错、
组合体检品种费率、寻优端点费率透传。

安全:LLM api_url 仅 http/https 且禁 userinfo(封死 file:// 读取与 Key 外送链)、
错误响应不回显原始 body、响应体 2MB 上限、配置原子写、坏配置字段级防御。

数据:涨跌停价整数分币舍入(67/318/90 个价位错 1 分漏判清零)、交易时段/采样/
provisional 统一沪时区、warehouse 增量缺口自动全量重拉、provisional 定点转正、
baostock 真故障抛错 + W/M 去 tradestatus(实测服务端报错,周月兜底此前从未工作)
+ 指数 vol 股→手(实测锚定)、ccpm 结构变更抛错。

Web API:缓存键补 count/vipdoc、NaN 清洗先于缓存、count>800 分页取全量、
submit 透传真实状态、pending 不再被淘汰成幽灵、watchlist/server 入参约束。

公式:FILTER 去副作用、0-1 值域误判收严、递归深度上限、REF 负移位显式禁止。

前端:4 处请求竞态序号守卫、Sparkline viewBox、北交所 market=2 映射、
空数据缓存死角、AI 弹窗卸载中止轮询、量能/资金日历口径修正。

CLI/CI:warehouse sync 失败 exit 1、参数校验干净报错、release 真实发布 SHA256、
CI 超时与缓存、spec 补 baostock 前提。

约 60 条回归测试先红后绿;pytest 1820 全过,ruff/mypy/vue-tsc/node --test 全绿。
This commit is contained in:
Justin Gu
2026-09-06 22:16:48 +08:00
parent d5cf9c14a9
commit e374a0da28
86 changed files with 4245 additions and 442 deletions
+17
View File
@@ -9,6 +9,7 @@ on:
jobs: jobs:
test: test:
runs-on: ${{ matrix.os }} runs-on: ${{ matrix.os }}
timeout-minutes: 60 # Windows runner 偶发卡顿,6 格矩阵整体留足余量
strategy: strategy:
fail-fast: false fail-fast: false
matrix: matrix:
@@ -22,6 +23,10 @@ jobs:
- uses: actions/setup-python@v5 - uses: actions/setup-python@v5
with: with:
python-version: ${{ matrix.python-version }} python-version: ${{ matrix.python-version }}
cache: "pip"
cache-dependency-path: |
pyproject.toml
requirements-dev.txt
# 前端 dist 被 pyproject.toml 的 force-include 引用,pip install -e . # 前端 dist 被 pyproject.toml 的 force-include 引用,pip install -e .
# 要求它存在。先构建前端(type-check + vite build)。 # 要求它存在。先构建前端(type-check + vite build)。
- uses: actions/setup-node@v4 - uses: actions/setup-node@v4
@@ -43,11 +48,14 @@ jobs:
mypy: mypy:
runs-on: ubuntu-latest runs-on: ubuntu-latest
timeout-minutes: 30
steps: steps:
- uses: actions/checkout@v4 - uses: actions/checkout@v4
- uses: actions/setup-python@v5 - uses: actions/setup-python@v5
with: with:
python-version: "3.13" python-version: "3.13"
cache: "pip"
cache-dependency-path: pyproject.toml
# 同 test jobforce-include 要求 web-ui/dist 存在。 # 同 test jobforce-include 要求 web-ui/dist 存在。
- uses: actions/setup-node@v4 - uses: actions/setup-node@v4
with: with:
@@ -67,6 +75,7 @@ jobs:
# 与 Python 版本无关,故独立 job,避免在 test matrix 里跑 6 遍。 # 与 Python 版本无关,故独立 job,避免在 test matrix 里跑 6 遍。
frontend: frontend:
runs-on: ubuntu-latest runs-on: ubuntu-latest
timeout-minutes: 60 # Playwright 安装 + E2E 首跑较慢
steps: steps:
- uses: actions/checkout@v4 - uses: actions/checkout@v4
- uses: actions/setup-node@v4 - uses: actions/setup-node@v4
@@ -84,7 +93,15 @@ jobs:
- uses: actions/setup-python@v5 - uses: actions/setup-python@v5
with: with:
python-version: "3.13" python-version: "3.13"
cache: "pip"
cache-dependency-path: pyproject.toml
- run: pip install -e ".[web]" - run: pip install -e ".[web]"
# 浏览器二进制按 lockfile 哈希缓存(版本变 → 新 key → 重下)
- name: Cache Playwright browsers
uses: actions/cache@v4
with:
path: ~/.cache/ms-playwright
key: playwright-${{ runner.os }}-${{ hashFiles('web-ui/package-lock.json') }}
- run: npx playwright install --with-deps chromium - run: npx playwright install --with-deps chromium
working-directory: web-ui working-directory: web-ui
- run: npm run test:e2e - run: npm run test:e2e
+70 -31
View File
@@ -57,10 +57,22 @@ jobs:
mv dist/easy-tdx.exe "dist/easy-tdx-${VERSION}-windows.exe" mv dist/easy-tdx.exe "dist/easy-tdx-${VERSION}-windows.exe"
ls -lh dist/ ls -lh dist/
# 生成 SHA256 清单并随产物上传——Release 正文引导用户核对该哈希
# (防杀软误报时自行验证完整性),必须真实存在而非只写在文案里。
- name: Compute SHA256
shell: bash
run: |
cd dist
FILE=$(ls easy-tdx-*-windows.exe)
sha256sum "$FILE" > SHA256SUMS.txt
cat SHA256SUMS.txt
- uses: actions/upload-artifact@v4 - uses: actions/upload-artifact@v4
with: with:
name: easy-tdx-windows-exe name: easy-tdx-windows-exe
path: dist/easy-tdx-*-windows.exe path: |
dist/easy-tdx-*-windows.exe
dist/SHA256SUMS.txt
if-no-files-found: error if-no-files-found: error
release: release:
@@ -73,37 +85,64 @@ jobs:
name: easy-tdx-windows-exe name: easy-tdx-windows-exe
path: dist path: dist
# 重算哈希(与构建 job 的 SHA256SUMS.txt 同源同值),嵌入 Release 正文,
# 使「核对 Release 说明中公布的 SHA256」的指引真实可执行。
- name: Compute SHA256 for release notes
run: |
EXE=$(ls dist/easy-tdx-*-windows.exe)
SHA=$(sha256sum "$EXE" | cut -d' ' -f1)
echo "EXE_NAME=$(basename "$EXE")" >> "$GITHUB_ENV"
echo "EXE_SHA256=$SHA" >> "$GITHUB_ENV"
# 引号 heredoc 防止正文里的反引号被 bash 当命令替换,先写占位符再 sed 替换
- name: Generate release body
run: |
cat > release_body.md <<'BODY'
## 下载使用
1. 下载下方 `__EXE_NAME__`(约 40-50MB
2. 双击运行
3. 浏览器会自动打开回测界面(地址 `http://localhost:8000`
## ✅ 完整性核对(SHA256
```
__EXE_SHA256__ __EXE_NAME__
```
PowerShell 核对:`Get-FileHash __EXE_NAME__`,与上方一致即为官方构建。
同份哈希也在附件 `SHA256SUMS.txt` 中。
## ⚠️ SmartScreen 提示
本版本未做代码签名,首次运行 Windows 会弹出"已保护你的电脑":
1. 点击 **更多信息**
2. 点击 **仍要运行**
Phase 2 将引入代码签名消除此提示。
## 🛡️ 杀软误报说明(Windows Defender 报毒?请先读这里)
部分 Windows Defender 用户会看到 `Trojan:Win32/Wacatac.C!ml` 报警。
`.exe` 由 GitHub Actions 在公开的 tag 提交上从源码构建(本文件即构建脚本),
为**未签名的 PyInstaller 单文件**打包——"自解压 + 无签名"特征是杀软
机器学习引擎(`!ml` 后缀即 ML 判定)的经典误报源头,v1.32.1 起各版本
构建方式相同。自行核验:
1. 核对上方公布的 SHA256PowerShell`Get-FileHash __EXE_NAME__`
2. 可上传 [VirusTotal](https://www.virustotal.com/) 交叉验证:典型误报特征是
少数 ML 启发式引擎报警、主流特征码引擎(Kaspersky/ESET/BitDefender 等)不报
3. 哈希一致仍想加速白名单,可向微软提交误报申诉:
[microsoft.com/en-us/wdsi/filesubmission](https://www.microsoft.com/en-us/wdsi/filesubmission)(选"软件开发者"
BODY
sed -i "s/__EXE_NAME__/${EXE_NAME}/g; s/__EXE_SHA256__/${EXE_SHA256}/g" release_body.md
- name: Create release - name: Create release
uses: softprops/action-gh-release@v2 uses: softprops/action-gh-release@v2
with: with:
files: dist/easy-tdx-*-windows.exe files: |
dist/easy-tdx-*-windows.exe
dist/SHA256SUMS.txt
generate_release_notes: true generate_release_notes: true
body: | body_path: release_body.md
## 下载使用
1. 下载下方 `easy-tdx-*-windows.exe`(约 40-50MB
2. 双击运行
3. 浏览器会自动打开回测界面(地址 `http://localhost:8000`
## ⚠️ SmartScreen 提示
本版本未做代码签名,首次运行 Windows 会弹出"已保护你的电脑":
1. 点击 **更多信息**
2. 点击 **仍要运行**
Phase 2 将引入代码签名消除此提示。
## 🛡️ 杀软误报说明(Windows Defender 报毒?请先读这里)
部分 Windows Defender 用户会看到 `Trojan:Win32/Wacatac.C!ml` 报警。
`.exe` 由 GitHub Actions 在公开的 tag 提交上从源码构建(本文件即构建脚本),
为**未签名的 PyInstaller 单文件**打包——"自解压 + 无签名"特征是杀软
机器学习引擎(`!ml` 后缀即 ML 判定)的经典误报源头,v1.32.1 起各版本
构建方式相同。自行核验:
1. 核对 Release 说明中公布的 SHA256PowerShell`Get-FileHash *.exe`
2. 可上传 [VirusTotal](https://www.virustotal.com/) 交叉验证:典型误报特征是
少数 ML 启发式引擎报警、主流特征码引擎(Kaspersky/ESET/BitDefender 等)不报
3. 哈希一致仍想加速白名单,可向微软提交误报申诉:
[microsoft.com/en-us/wdsi/filesubmission](https://www.microsoft.com/en-us/wdsi/filesubmission)(选"软件开发者"
+42
View File
@@ -2,6 +2,48 @@
本文件记录 easy-tdx 的版本变更。格式遵循 [Keep a Changelog](https://keepachangelog.com/zh-CN/)。 本文件记录 easy-tdx 的版本变更。格式遵循 [Keep a Changelog](https://keepachangelog.com/zh-CN/)。
## [1.32.6] — 2026-09-06
**两周改动深度审查后的全面修复**——对 v1.21→v1.32.5 的 249 个文件、4.2 万行改动做六路专项审查,本轮落地全部修复:回测数字可信性(组合收益虚增 / 轮动停牌成交 / WF 指标稀释三件套)、LLM 安全加固(封死 file:// 读取与 API Key 外送链)、涨跌停价舍入漏判、数据源与缓存正确性、前端请求竞态等,共 58 处行为修复、约 60 条"先红后绿"回归测试。
### 回测正确性(影响数字可信性,建议重点升级)
- **组合收益不再虚增**[portfolio_engine.py](src/easy_tdx/backtest/portfolio_engine.py)、[multi_strategy_engine.py](src/easy_tdx/backtest/multi_strategy_engine.py)):组合内标的起始日期不齐时(次新股/取数截断),合并净值曲线前导缺口旧代码填 0,导致曲线首值小于总投入、`total_return` 被系统性虚增(实测两标的各投 10 万可显示 +126%,真实约 +13%);改为前导缺口按初始资金回填,与组合 Walk-Forward 口径一致,`total_return` 恒等于资金加权收益率。
- **轮动回测停牌日不再按过期价成交**([rotation.py](src/easy_tdx/backtest/rotation.py)):停牌日挂单此前会以停牌前旧开盘价成交(实测卖出价虚增 43%);改为当日真实有交易才可成交,挂单顺延至复牌开盘,符合真实挂单语义。另修:首个调仓日照常产生信号、历史不足 5 根的次新股不再以 0 分混入买入候选。
- **单标的 Walk-Forward 逐窗指标不再被预热区稀释**([walkforward.py](src/easy_tdx/backtest/walkforward.py)):上下文预热 bar 此前计入窗口绩效,默认参数下逐窗 sharpe/年化被零收益段大幅稀释(实测 8.84 vs 正确 14.96);改为只用窗内净值与成交计算,与组合级同口径。另修:`worst_drawdown` 方向取反、单窗不足 20 根强制跳过、窗口失败记 warning 不再静默、int 日期列窗口标签显示 1970-01-01。
- **评分容错**[grading.py](src/easy_tdx/backtest/grading.py)):组合回撤序列缺行/NaN 不再截断或放大水下期计数;组合级"适配性体检"对 ETF/可转债组合按品种费率解析(此前错收股票印花税)。
### 安全
- **LLM 接口加固**[ai/llm.py](src/easy_tdx/ai/llm.py)):`api_url` 仅允许 http/https 且禁止携带 URL 凭据(封死 `file://` 读本地配置文件与内网 SSRF 链);HTTP 错误不再回显原始响应体(改为只提取 provider 错误 message,杜绝错误通道变任意内容回读);响应体 2MB 上限;配置文件原子写(崩溃不再半写损坏);手工编辑的 llm.json 坏字段自动回退默认并告警,不再打挂全部 AI 端点;anthropic 路径思考型模型空正文给出可操作报错而非静默空回复。
### 数据与统计
- **涨跌停价舍入修复**[limitup.py](src/easy_tdx/screen/limitup.py)):旧浮点实现 `floor(x*100+0.5)` 在半分边界受浮点误差影响,±10% 档 67/318 个、±5% 档 90 个价位会算低 1 分(如 33.05×1.1 误算 36.35,交易所 36.36),真实涨跌停被静默漏判;改为纯整数分币运算,全价位对账与交易所零差异。涨停家数/连板高度/炸板率/离线回补一并修正,ST 5% 与 3 元低价门槛改为逐 bar 判定。
- **时区统一沪市时间**[session.py](src/easy_tdx/realtime/session.py) 等):`is_trading_time`、情绪采样、warehouse provisional 标记/转正、热点"今日"列此前用主机本地时区,非中国时区主机上情绪页永久无数据、当日 K 线被长期标 provisional;统一为固定 UTC+8(中国无夏令时),与主机时区无关。
- **warehouse 增量同步**[sync.py](src/easy_tdx/warehouse/sync.py)、[store.py](src/easy_tdx/warehouse/store.py)):停用超过约三周后再同步,尾部 15 根覆盖不到的缺口此前永久丢失且 summary 照常 ok;现在检测到缺口自动全量重拉并告警。provisional 盘中临时值改为"拉取成功后仅转正到本次拉到的最新 bar",数据源失败不再把未定值洗成 completed。
- **baostock 兜底**[sources/baostock.py](src/easy_tdx/sources/baostock.py)):真故障(登录失败/查询报错)改为记日志并抛错,仅真无数据返回空,`--source baostock` 不再静默"无数据";周/月线去掉服务端明确报错的 `tradestatus` 字段(实测 error_code=10004012——周/月兜底此前从未工作过);指数兜底 vol 按实测(sh.000001)从股换算为手对齐 /bars/index 契约;`/bars` 数字周期串(如 `category=4`)不再绕过兜底。
- **ccpm**:中金所页面结构变更时抛 `CcpmError` 并附结构线索,不再静默返回空表被当成"无数据"。
- **Web API 正确性**[board_mac.py](src/easy_tdx/web/routers/board_mac.py)、[market.py](src/easy_tdx/web/routers/market.py)、[schemas.py](src/easy_tdx/web/schemas.py)、[bars.py](src/easy_tdx/web/routers/bars.py) 等):板块总览/涨停缓存键补齐 `count`/`vipdoc`(15s TTL 内不同参数不再互相串台);响应 NaN 递归清洗且先清洗后入缓存(一行 NaN 导致稳定 500 的问题消除);回测/寻优/公式端点 `count>800` 改分页取全量(TDX 单次协议上限,此前静默截短回测窗口);任务提交响应透传真实状态(极快完成的任务不再谎报 running);任务淘汰不再把未起跑的 pending 淘汰成"永久 pending 幽灵"。
### 公式解析
- [formula.py](src/easy_tdx/formula.py)、[MyTT.py](src/easy_tdx/MyTT.py):FILTER 不再原地改写输入序列(`FILTER(C,2); MA(C,2)` 的 MA 此前被污染);值域兜底收严(RSI/100、价格比率等 0~1 数值列不再被误判为买卖信号列);解析深度上限 100 层(5000 层嵌套从 RecursionError 裸崩改为 FormulaError);REF 负移位显式禁止(未来函数入口不再依赖类型巧合拦截)。
### 前端
- **请求竞态守卫**[HotspotView.vue](web-ui/src/views/HotspotView.vue)、HotspotCorrelation、BoardOverviewView、BoardDialog 四处同款):快速切换类型/参数时旧响应后到不再覆盖新数据、不再杀死构建轮询、不再产生幽灵翻红翻绿事件,休市/后台标签页不再停留错误状态。
- 其余:Dashboard 指数分时不再被裁掉约 39%Sparkline viewBox 随尺寸);北交所个股从榜单点开能正确拉到行情(market=2 映射与 920xxx 前缀修正);大盘日历空数据不再被永久缓存(重试可用);AI 解读弹窗关闭即中止轮询(此前最长空转 20 分钟);量能图与资金日历口径/正负号展示修正。
### CLI / CI / 打包
- `warehouse sync` 全部失败退出码改 1(对齐 ccpm);`--period` 参数枚举校验;`--source baostock` + 非日线周期前置拦截;公式/仓库命令缺冒号参数改为干净报错(此前裸 traceback)。
- Release 工作流现在真实计算并发布 SHA256SUMS.txt(正文"核对哈希"指引从此成立);CI 各 job 加超时与 pip/Playwright 缓存;PyInstaller spec 构建前提补 baostock extra 并在缺包时明确报错。
### 测试
- 约 60 条回归测试全部"先在旧码跑出失败、再修复转绿";两项存疑以实测定案:baostock 周/月线字段与指数 vol 单位(真实登录对账)、深市 ETF vol 单位(本机 vipdoc 与服务器双源一致,现有 ÷100 换算正确)。全量 1820 通过,ruff / ruff format / mypy 严格模式 / `vue-tsc` / 前端 node --test 23 条全绿。
## [1.32.5] — 2026-09-06 ## [1.32.5] — 2026-09-06
**回撤持续统计修复:不再恒为 1**——单标的回测「绩效指标 → 风险 → 回撤持续」此前无论什么股票都显示 1,本版修复计算错误,并把三处实现的口径统一为指标文档承诺的「最长水下期」:从净值峰值跌落到重新创新高的最长天数(末日仍未修复则计到最后一天),即"最长一次套牢了多久"。 **回撤持续统计修复:不再恒为 1**——单标的回测「绩效指标 → 风险 → 回撤持续」此前无论什么股票都显示 1,本版修复计算错误,并把三处实现的口径统一为指标文档承诺的「最长水下期」:从净值峰值跌落到重新创新高的最长天数(末日仍未修复则计到最后一天),即"最长一次套牢了多久"。
+15 -2
View File
@@ -3,7 +3,9 @@
构建前提(CI 会自动完成,本地手动构建需自行执行):: 构建前提(CI 会自动完成,本地手动构建需自行执行)::
1. pip install -e ".[web]" pyinstaller 1. pip install -e ".[web,packaging,baostock]" pyinstaller
# baostock 为 EXE 内置兜底数据源(见下方 hiddenimports),缺它会
# 被 spec 顶部的显式检查拦下——这是有意设计,防止静默产出无兜底的 EXE
2. cd web-ui && npm ci && npm run build # 产出 web-ui/dist/ 2. cd web-ui && npm ci && npm run build # 产出 web-ui/dist/
3. pyinstaller easy_tdx.spec # 产出 dist/easy-tdx.exe 3. pyinstaller easy_tdx.spec # 产出 dist/easy-tdx.exe
@@ -24,8 +26,19 @@ from PyInstaller.utils.hooks import collect_data_files, collect_submodules
hiddenimports: list[str] = [] hiddenimports: list[str] = []
hiddenimports += collect_submodules("uvicorn") hiddenimports += collect_submodules("uvicorn")
hiddenimports += collect_submodules("easy_tdx") hiddenimports += collect_submodules("easy_tdx")
# baostock 兜底数据源(v1.33/bars 与 warehouse 的最后一级回退)在 # baostock 兜底数据源(/bars 与 warehouse 的最后一级回退)在
# sources/baostock.py 里经 importlib 懒加载,静态分析扫不到,需显式声明。 # sources/baostock.py 里经 importlib 懒加载,静态分析扫不到,需显式声明。
# 显式检查而非依赖 collect_submodules 的失败形态(不同 PyInstaller 版本下
# 可能返回空列表静默跳过)——缺 baostock 时立即报错,防止静默产出
# 「兜底源缺失」的 EXE(运行期 /bars 自动回退静默失效,极难排查)。
import importlib.util
if importlib.util.find_spec("baostock") is None:
raise SystemExit(
": baostockEXE "
'请先执行 pip install baostock(或 pip install -e ".[web,packaging,baostock]"'
" pyinstaller easy_tdx.spec"
)
hiddenimports += collect_submodules("baostock") hiddenimports += collect_submodules("baostock")
# pandas / numpy / scipy 由 PyInstaller 自带 hook 处理(见 # pandas / numpy / scipy 由 PyInstaller 自带 hook 处理(见
# PyInstaller/hooks/hook-pandas.* 等),无需手动 collect_submodules—— # PyInstaller/hooks/hook-pandas.* 等),无需手动 collect_submodules——
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project] [project]
name = "easy-tdx" name = "easy-tdx"
version = "1.32.5" version = "1.32.6"
description = "通达信 TCP 协议行情数据客户端,支持在线行情、离线数据读取与写入同步" description = "通达信 TCP 协议行情数据客户端,支持在线行情、离线数据读取与写入同步"
readme = "README.md" readme = "README.md"
requires-python = ">=3.10" requires-python = ">=3.10"
+7 -3
View File
@@ -183,9 +183,13 @@ def EXIST(S, N): # EXIST(CLOSE>3010, N=5) n日内是否存在一天大于3000
def FILTER(S, N): # FILTER函数,S满足条件后,将其后N周期内的数据置为0, FILTER(C==H,5) def FILTER(S, N): # FILTER函数,S满足条件后,将其后N周期内的数据置为0, FILTER(C==H,5)
for i in range(len(S)): # 无副作用实现:在副本上置零。曾直接改写输入序列——公式通道里
S[i + 1 : i + 1 + N] = 0 if S[i] else S[i + 1 : i + 1 + N] # FILTER(C, N) 会把同一公式后续语句引用的 C 一并污染(或对只读
return S # 例:FILTER(C==H,5) 涨停后,后5天不再发出信号 # 数组直接报错)。
out = np.array(S, copy=True)
for i in range(len(out)):
out[i + 1 : i + 1 + N] = 0 if out[i] else out[i + 1 : i + 1 + N]
return out # 例:FILTER(C==H,5) 涨停后,后5天不再发出信号
def BARSLAST(S): # 上一次条件成立到当前的周期, BARSLAST(C/REF(C,1)>=1.1) 上一次涨停到今天的天数 def BARSLAST(S): # 上一次条件成立到当前的周期, BARSLAST(C/REF(C,1)>=1.1) 上一次涨停到今天的天数
+260 -30
View File
@@ -29,9 +29,12 @@ from __future__ import annotations
import asyncio import asyncio
import json import json
import logging import logging
import math
import os import os
import tempfile
import time import time
import urllib.error import urllib.error
import urllib.parse
import urllib.request import urllib.request
from dataclasses import asdict, dataclass, field, replace from dataclasses import asdict, dataclass, field, replace
from pathlib import Path from pathlib import Path
@@ -45,6 +48,7 @@ __all__ = [
"mask_key", "mask_key",
"resolve_config", "resolve_config",
"save_config", "save_config",
"validate_api_url",
] ]
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -157,29 +161,108 @@ def _read_config_file() -> dict[str, Any]:
def load_config() -> LlmConfig: def load_config() -> LlmConfig:
"""加载配置:llm.json 显式字段 > 环境变量兜底(未填字段仍为空,调用时再取预设)。""" """加载配置:llm.json 显式字段 > 环境变量兜底(未填字段仍为空,调用时再取预设)。
字段级防御:llm.json 常被手工编辑,单个字段类型不对(``"temperature":
null`` / ``"abc"`` 等)只记 warning 并回退默认值,绝不让 load_config
抛异常打挂全部 /llm/* 端点。
"""
data = _read_config_file() data = _read_config_file()
env_url = os.environ.get("LLM_BASE_URL", "") env_url = os.environ.get("LLM_BASE_URL", "")
cfg = LlmConfig( return LlmConfig(
provider=str(data.get("provider") or os.environ.get("LLM_PROVIDER", "") or "deepseek"), provider=_clean_str(data.get("provider") or os.environ.get("LLM_PROVIDER", ""), "provider")
api_url=str(data.get("api_url") or env_url or ""), or "deepseek",
api_key=str(data.get("api_key") or os.environ.get("LLM_API_KEY", "") or ""), api_url=_clean_str(data.get("api_url") or env_url, "api_url"),
model=str(data.get("model") or os.environ.get("LLM_MODEL", "") or ""), api_key=_clean_str(data.get("api_key") or os.environ.get("LLM_API_KEY", ""), "api_key"),
temperature=float(data.get("temperature", 0.3)), model=_clean_str(data.get("model") or os.environ.get("LLM_MODEL", ""), "model"),
max_tokens=int(data.get("max_tokens", 16000)), temperature=_clean_float(data.get("temperature"), "temperature", 0.3, minimum=0.0),
timeout=float(data.get("timeout", 180.0)), max_tokens=_clean_int(data.get("max_tokens"), "max_tokens", 16000, minimum=1),
system_prompt=str(data.get("system_prompt", "") or LlmConfig.system_prompt), timeout=_clean_float(data.get("timeout"), "timeout", 180.0, minimum=0.1),
system_prompt=_clean_str(
data.get("system_prompt") or LlmConfig.system_prompt, "system_prompt"
)
or LlmConfig.system_prompt,
) )
return cfg
def _clean_str(value: Any, field_name: str) -> str:
"""字符串字段清洗:非 None 的非字符串(数字/列表等)记 warning 后回退空串。"""
if value is None:
return ""
if not isinstance(value, str):
logger.warning(
"LLM 配置字段 %s 应为字符串,收到 %s(%r),已忽略并回退默认值",
field_name,
type(value).__name__,
value,
)
return ""
return value
def _clean_float(
value: Any, field_name: str, default: float, *, minimum: float | None = None
) -> float:
"""数值字段清洗:非法/非有限/低于下限均 warning 后回退默认。"""
if value is None:
return default
try:
f = float(value)
except (TypeError, ValueError):
logger.warning(
"LLM 配置字段 %s 应为数字,收到 %r,回退默认值 %s", field_name, value, default
)
return default
if not math.isfinite(f) or (minimum is not None and f < minimum):
logger.warning(
"LLM 配置字段 %s 超出合理范围(%r),回退默认值 %s", field_name, value, default
)
return default
return f
def _clean_int(value: Any, field_name: str, default: int, *, minimum: int | None = None) -> int:
"""整数字段清洗:接受 "8192.9" 这类字符串的宽松矫正(截断到 int)。"""
if value is None:
return default
try:
i = int(value)
except (TypeError, ValueError):
try:
i = int(float(value))
except (TypeError, ValueError):
logger.warning(
"LLM 配置字段 %s 应为整数,收到 %r,回退默认值 %s", field_name, value, default
)
return default
if not math.isfinite(i) or (minimum is not None and i < minimum):
logger.warning(
"LLM 配置字段 %s 超出合理范围(%r),回退默认值 %s", field_name, value, default
)
return default
return i
def save_config(cfg: LlmConfig) -> Path: def save_config(cfg: LlmConfig) -> Path:
"""写入 llm.json(WebUI 保存入口;目录惰性创建)。""" """写入 llm.json(WebUI 保存入口;目录惰性创建;原子写)。
先写同目录临时文件再 ``os.replace``——写一半崩溃/断电不会留下损坏的
llm.json(损坏的后果是下次 load 静默回空配置,用户要重填 key)。
"""
p = config_path() p = config_path()
p.parent.mkdir(parents=True, exist_ok=True) p.parent.mkdir(parents=True, exist_ok=True)
p.write_text( text = json.dumps(cfg.to_dict(), ensure_ascii=False, indent=2)
json.dumps(cfg.to_dict(), ensure_ascii=False, indent=2), encoding="utf-8", newline="\n" fd, tmp_name = tempfile.mkstemp(dir=str(p.parent), prefix=".llm-", suffix=".tmp")
) try:
with os.fdopen(fd, "w", encoding="utf-8", newline="\n") as f:
f.write(text)
os.replace(tmp_name, p)
except BaseException:
try:
os.unlink(tmp_name)
except OSError:
pass
raise
return p return p
@@ -188,10 +271,12 @@ def resolve_config(cfg: LlmConfig | None = None) -> LlmConfig:
- ``api_url`` 空 → 预设 ``base_url`` - ``api_url`` 空 → 预设 ``base_url``
- ``model`` 空 → 预设 ``default_model`` - ``model`` 空 → 预设 ``default_model``
- provider 无预设(拼错)→ 按 custom 处理,url/model 必须已填 - provider 无预设(拼错)→ 按 custom 处理,url/model 必须已填
- ``api_url`` 强制 http/https 且禁 userinfoSSRF/本地文件读取防线)。
Raises: Raises:
ValueError: 补齐后仍缺 api_url 或 modelcustom 未填全) ValueError: 补齐后仍缺 api_url 或 modelcustom 未填全)
或 api_url 非法(非 http/https、携带 user:pass@)。
""" """
c = replace(cfg or load_config()) c = replace(cfg or load_config())
preset = PROVIDER_PRESETS.get(c.provider, PROVIDER_PRESETS["custom"]) preset = PROVIDER_PRESETS.get(c.provider, PROVIDER_PRESETS["custom"])
@@ -203,6 +288,7 @@ def resolve_config(cfg: LlmConfig | None = None) -> LlmConfig:
raise ValueError( raise ValueError(
f"LLM 配置不完整:provider={c.provider} 缺少 api_url 或 model,请在 AI 设置中补全" f"LLM 配置不完整:provider={c.provider} 缺少 api_url 或 model,请在 AI 设置中补全"
) )
validate_api_url(c.api_url)
return c return c
@@ -226,6 +312,80 @@ class LlmError(RuntimeError):
self.status = status self.status = status
#: 响应体大小上限:正常 chat 响应远小于此(max_tokens 128k 的纯文本约几百 KB),
#: 超限说明对端异常(如把 api_url 配成了下载地址),及时中止防内存被撑爆。
_MAX_RESPONSE_BYTES = 2 * 1024 * 1024
_READ_CHUNK = 64 * 1024
#: 允许的 URL schemeSSRF / 本地文件读取防线:file:// 会读本地文件、
#: ftp:// 与内网 http 可被当跳板——resolve_config 里强制校验)。
_ALLOWED_URL_SCHEMES = ("http", "https")
def validate_api_url(url: str) -> None:
"""api_url 安全校验:仅 http/https、禁止携带 userinfouser:pass@)。
Raises:
ValueError: scheme 非法/缺失,或 URL 携带用户凭据(web 层已有
ValueError→错误响应通道,CLI 场景同样可直接展示)。
"""
if not url:
return
parts = urllib.parse.urlsplit(url)
scheme = parts.scheme.lower()
if scheme not in _ALLOWED_URL_SCHEMES:
raise ValueError(
f"LLM api_url 非法:仅允许 http/https 地址(收到 scheme={scheme!r})——"
"file/ftp 等协议已禁用;缺前缀时请补 http:// 或 https://"
)
if parts.username or parts.password:
raise ValueError(
"LLM api_url 非法:不允许携带用户凭据(user:pass@host 形式)——"
"请把鉴权放到 API Key 字段(请求头),而不是 URL 里"
)
def _read_capped(fp: Any, cap: int = _MAX_RESPONSE_BYTES) -> bytes:
"""分块读响应体,超过 cap 字节即中止(防异常网关撑爆内存)。"""
buf = bytearray()
while True:
chunk = fp.read(_READ_CHUNK)
if not chunk:
break
buf.extend(chunk)
if len(buf) > cap:
raise LlmError(
f"LLM API 响应超过 {cap // (1024 * 1024)}MB 上限——对端不是正常的 chat 接口"
"(请检查 api_url 是否填错),已中止"
)
return bytes(buf)
def _extract_error_message(body: str) -> str | None:
"""从 provider 错误响应中提取可读 message(不回显原始 body)。
OpenAI/Anthropic 兼容网关的惯例是 ``{"error": {"message": ...}}``
部分网关 error 直接是字符串,或把 message 放顶层。提取结果截到 300 字符。
"""
try:
obj = json.loads(body)
except (json.JSONDecodeError, ValueError):
return None
if not isinstance(obj, dict):
return None
err = obj.get("error")
if isinstance(err, dict):
msg = err.get("message") or err.get("msg") or err.get("code")
if msg:
return str(msg)[:300]
elif isinstance(err, str) and err:
return err[:300]
msg = obj.get("message")
if msg:
return str(msg)[:300]
return None
def _post_json( def _post_json(
url: str, headers: dict[str, str], payload: dict[str, Any], timeout: float url: str, headers: dict[str, str], payload: dict[str, Any], timeout: float
) -> dict[str, Any]: ) -> dict[str, Any]:
@@ -235,6 +395,7 @@ def _post_json(
超时单独成类报错:非流式 chat 接口要等模型**整段回复生成完**才回包, 超时单独成类报错:非流式 chat 接口要等模型**整段回复生成完**才回包,
大 Prompt(如整份回测报告解读)生成 1-3 分钟很正常,读超时≠网络故障, 大 Prompt(如整份回测报告解读)生成 1-3 分钟很正常,读超时≠网络故障,
报错必须把「调大超时」这个动作说清楚(v1.29.1 实测踩坑)。 报错必须把「调大超时」这个动作说清楚(v1.29.1 实测踩坑)。
HTTP 错误只回显 provider 的 error.message(≤300 字符),不透传原始 body。
""" """
req = urllib.request.Request( req = urllib.request.Request(
url, url,
@@ -244,11 +405,16 @@ def _post_json(
) )
try: try:
with urllib.request.urlopen(req, timeout=timeout) as resp: with urllib.request.urlopen(req, timeout=timeout) as resp:
data: dict[str, Any] = json.loads(resp.read().decode("utf-8")) raw = _read_capped(resp)
data: dict[str, Any] = json.loads(raw.decode("utf-8", errors="replace"))
return data return data
except urllib.error.HTTPError as exc: except urllib.error.HTTPError as exc:
body = exc.read().decode("utf-8", errors="replace")[:500] try:
raise LlmError(f"LLM API HTTP {exc.code}: {body}", status=exc.code) from exc body = exc.read(_READ_CHUNK).decode("utf-8", errors="replace")
except OSError:
body = ""
detail = _extract_error_message(body) or "接口返回错误响应"
raise LlmError(f"LLM API HTTP {exc.code}: {detail}", status=exc.code) from exc
except urllib.error.URLError as exc: except urllib.error.URLError as exc:
if isinstance(exc.reason, TimeoutError): if isinstance(exc.reason, TimeoutError):
raise LlmError(_timeout_message(timeout)) from exc raise LlmError(_timeout_message(timeout)) from exc
@@ -256,7 +422,9 @@ def _post_json(
except TimeoutError as exc: except TimeoutError as exc:
raise LlmError(_timeout_message(timeout)) from exc raise LlmError(_timeout_message(timeout)) from exc
except json.JSONDecodeError as exc: except json.JSONDecodeError as exc:
raise LlmError(f"LLM API 响应不是合法 JSON: {exc}") from exc raise LlmError(
"LLM API 响应不是合法 JSON——api_url 可能不是 chat 接口端点,请检查 AI 设置"
) from exc
def _timeout_message(timeout: float) -> str: def _timeout_message(timeout: float) -> str:
@@ -325,8 +493,14 @@ class LlmClient:
except LlmError: except LlmError:
raise raise
except (KeyError, IndexError, TypeError) as exc: except (KeyError, IndexError, TypeError) as exc:
raw = json.dumps(data, ensure_ascii=False)[:300] # 不回显原始响应体:内容可能包含网关内部信息,且 historical 上
raise LlmError(f"LLM 响应格式异常: {raw}") from exc # 曾被当作任意 URL 响应的回读通道。只描述缺什么 + 顶层键名。
hint = sorted(data.keys()) if isinstance(data, dict) else type(data).__name__
raise LlmError(
"LLM 响应格式异常:未找到 choices[0].message 字段"
f"(响应顶层字段: {hint})——请检查 api_url 是否为正确的"
" chat/completions 端点、模型名是否正确"
) from exc
def _extract_reply_openai(self, message: dict[str, Any], finish: str) -> str: def _extract_reply_openai(self, message: dict[str, Any], finish: str) -> str:
"""从 OpenAI 兼容响应的 message 里提取正文,处理思考型模型的空白正文。 """从 OpenAI 兼容响应的 message 里提取正文,处理思考型模型的空白正文。
@@ -352,8 +526,11 @@ class LlmClient:
raise LlmError( raise LlmError(
"模型输出被 max_tokens 截断且无正文,请在「AI 设置」调大 Max Tokens 后重试" "模型输出被 max_tokens 截断且无正文,请在「AI 设置」调大 Max Tokens 后重试"
) )
raw = json.dumps(message, ensure_ascii=False)[:300] keys = sorted(message.keys()) if isinstance(message, dict) else type(message).__name__
raise LlmError(f"LLM 响应 message.content 为空: {raw}") raise LlmError(
f"LLM 响应 message.content 为空(message 字段: {keys}"
f"finish_reason={finish or 'unknown'})——请检查模型名与 api_url 是否匹配"
)
def _chat_anthropic(self, prompt: str, system: str) -> str: def _chat_anthropic(self, prompt: str, system: str) -> str:
cfg = self._cfg cfg = self._cfg
@@ -371,11 +548,64 @@ class LlmClient:
url = f"{cfg.api_url.rstrip('/')}/messages" url = f"{cfg.api_url.rstrip('/')}/messages"
data = _post_json(url, headers, payload, cfg.timeout) data = _post_json(url, headers, payload, cfg.timeout)
try: try:
blocks = data["content"] return self._extract_reply_anthropic(data)
return "".join(str(b.get("text", "")) for b in blocks if b.get("type") == "text") except LlmError:
except (KeyError, TypeError) as exc: raise
raw = json.dumps(data, ensure_ascii=False)[:300] except (KeyError, TypeError, AttributeError) as exc:
raise LlmError(f"LLM 响应格式异常: {raw}") from exc # 与 openai 路径同口径:不回显原始响应体,只描述问题
hint = sorted(data.keys()) if isinstance(data, dict) else type(data).__name__
raise LlmError(
f"LLM 响应格式异常:content 字段不可解析(响应顶层字段: {hint})——"
"请确认 api_url 指向 Anthropic /messages 端点、模型名正确"
) from exc
def _extract_reply_anthropic(self, data: dict[str, Any]) -> str:
"""从 Anthropic 响应提取正文,与 openai 路径同口径:绝不返回空串。
- content 是块列表:拼接 text 块,统计 thinking 块字数(思考耗尽
max_tokens 时报可操作错误,而非静默成功空串);
- content 是字符串(部分网关):直接作为正文;
- 空白正文:按 stop_reason 给「调大 Max Tokens」指引。
"""
content = data["content"]
if isinstance(content, str):
text = content
thinking_len = 0
block_types: list[str] = []
elif isinstance(content, list):
parts: list[str] = []
thinking_len = 0
block_types = []
for b in content:
if not isinstance(b, dict):
continue
b_type = str(b.get("type") or "")
block_types.append(b_type)
if b_type == "text":
parts.append(str(b.get("text") or ""))
elif b_type in ("thinking", "redacted_thinking"):
thinking_len += len(str(b.get("thinking") or b.get("data") or ""))
text = "".join(parts)
else:
raise TypeError(f"content 应为块列表或字符串,收到 {type(content).__name__}")
if text.strip():
return text
finish = str(data.get("stop_reason") or "")
if thinking_len:
raise LlmError(
f"模型只返回了思考链(thinking {thinking_len} 字),未生成正文——"
f"max_tokens={self._cfg.max_tokens} 大概率被思考耗尽"
f"stop_reason={finish or 'unknown'})。"
"请在「AI 设置」把 Max Tokens 调大(思考型模型建议 ≥16000)后重试"
)
if finish == "max_tokens":
raise LlmError(
"模型输出被 max_tokens 截断且无正文,请在「AI 设置」调大 Max Tokens 后重试"
)
raise LlmError(
f"LLM 响应 content 为空(block 类型: {block_types or ''}"
f"stop_reason={finish or 'unknown'})——请检查模型名与 api_url 是否匹配"
)
async def test(self) -> dict[str, Any]: async def test(self) -> dict[str, Any]:
"""连通性测试:发一句极短 ping,返回 ok/延迟/样例回复。""" """连通性测试:发一句极短 ping,返回 ok/延迟/样例回复。"""
+9 -2
View File
@@ -357,13 +357,20 @@ def evaluate_portfolio(
**engine_kwargs, **engine_kwargs,
).run() ).run()
# 3. 适配性体检:逐标的跑三段体检,跨标的多数口径聚合 # 3. 适配性体检:逐标的跑三段体检,跨标的多数口径聚合
# 每标的传入各自 symbol,使 auto_fees 按品种解析费率——与组合回测主路径
# PortfolioBacktestEngine 逐标的 resolve_fee_model)同口径。此前漏传
# symbol:ETF/可转债组合的三段体检被按股票口径错收印花税。
fitness_kwargs: dict[str, Any] = { fitness_kwargs: dict[str, Any] = {
k: v for k, v in engine_kwargs.items() if k not in ("total_cash", "chanlun_level") k: v for k, v in engine_kwargs.items() if k not in ("total_cash", "chanlun_level")
} }
per_stock_fitness = [ per_stock_fitness = [
FitnessEngine( FitnessEngine(
strategy=strategy, split=split, context_bars=context_bars, **fitness_kwargs strategy=strategy,
split=split,
context_bars=context_bars,
symbol=f"{stock.market}{stock.code}",
**fitness_kwargs,
).evaluate(stock.df) ).evaluate(stock.df)
for stock in stocks for stock in stocks
] ]
+17
View File
@@ -115,6 +115,16 @@ def backtest(
# 1. 加载策略(单策略 or 多因子组合) # 1. 加载策略(单策略 or 多因子组合)
is_combo = combo_strategies is not None is_combo = combo_strategies is not None
# 组合模式暂不支持的分析旗标:显式告警而非静默吞掉(审查修复)
if is_combo and walk_forward:
click.echo(
"警告: --wfWalk-Forward 样本外验证)暂不支持 --combo-strategies 组合模式,已忽略",
err=True,
)
if is_combo and full_evaluate:
click.echo(
"警告: --evaluate(一条龙评估)暂不支持 --combo-strategies 组合模式,已忽略", err=True
)
if is_combo: if is_combo:
assert combo_strategies is not None # narrowed by is_combo assert combo_strategies is not None # narrowed by is_combo
@@ -728,6 +738,13 @@ def optimize(
raise SystemExit(1) raise SystemExit(1)
custom_grid = _parse_param_grid(param_pairs) if param_pairs else None custom_grid = _parse_param_grid(param_pairs) if param_pairs else None
if optimize_all and custom_grid is not None:
# --all 逐策略使用各自预设网格,--param 无处安放:显式告警而非静默忽略(审查修复)
click.echo(
"警告: --param 在 --all 模式下被忽略(一键寻优逐策略使用各自预设网格;"
"如需自定义网格请指定 --strategy 单策略寻优)",
err=True,
)
if not optimize_all: if not optimize_all:
assert strategy_name is not None assert strategy_name is not None
try: try:
+25 -2
View File
@@ -131,6 +131,11 @@ THRESHOLDS: dict[str, tuple[str, tuple[Anchor, ...]]] = {
Anchor(0.6, 0), Anchor(0.6, 0),
), ),
), ),
# 回撤持续:该维度输入的单位是「bar 数」(performance.max_dd_duration 输出
# 水下期的 bar 数),下方天数锚点按日线(1 bar ≈ 1 交易日)校准。分钟级
# 周期下 bar 数远大于天数,得分会系统性偏低(偏保守)。注意:评分维度
# grade_performance / grade_portfolio_equity)当前均未使用该表项,
# 锚点仅供展示与前端对照。
"max_dd_duration": ( "max_dd_duration": (
"回撤持续", "回撤持续",
( (
@@ -400,7 +405,13 @@ def _downweight_unreliable(dimensions: list[DimensionScore]) -> bool:
@dataclass @dataclass
class CombinedMetrics: class CombinedMetrics:
"""从净值序列重算的组合级指标(净值可推导的字段子集)。""" """从净值序列重算的组合级指标(净值可推导的字段子集)。
``max_dd_duration`` 的单位是 **bar 数**(一根 K 线计 1),与
``performance.max_dd_duration`` 同口径;THRESHOLDS 里对应锚点的天数
30/90/365…)按日线(1 bar ≈ 1 交易日)校准,分钟级周期下该值按 bar
直读会显著大于天数(评级维度未使用,仅展示)。
"""
total_return: float = 0.0 total_return: float = 0.0
annual_return: float = 0.0 annual_return: float = 0.0
@@ -463,12 +474,24 @@ def compute_combined_metrics(equity: list[dict[str, Any]]) -> CombinedMetrics:
# 最大回撤:优先用 drawdown_pct(与前端一致),缺则从 totals 反推。 # 最大回撤:优先用 drawdown_pct(与前端一致),缺则从 totals 反推。
# 持续 = 最长水下期(峰值 → 重新创新高;末日未修复则计到最后一点), # 持续 = 最长水下期(峰值 → 重新创新高;末日未修复则计到最后一点),
# 与 performance.py / combinedMetrics.ts / max_dd_duration 锚点量纲同口径。 # 与 performance.py / combinedMetrics.ts / max_dd_duration 锚点量纲同口径。
# 缺行 / None / NaN 的 drawdown_pct 按「状态延续」处理:沿用上一根的
# 水下/峰值状态——此前缺行被当 0(创新高)截断水下期、NaN 永久脱离
# 峰值判定,两者都会扭曲 max_dd_duration。
max_dd = 0.0 max_dd = 0.0
max_dd_dur = 0 max_dd_dur = 0
if equity[0].get("drawdown_pct") is not None: if equity[0].get("drawdown_pct") is not None:
last_peak = 0 last_peak = 0
prev_dd = 0.0 # 上一根的有效回撤(首根之前视作峰值状态)
for i, e in enumerate(equity): for i, e in enumerate(equity):
dd = float(e.get("drawdown_pct") or 0.0) raw = e.get("drawdown_pct")
dd: float | None
try:
dd = None if raw is None else float(raw)
except (TypeError, ValueError):
dd = None
if dd is None or not math.isfinite(dd):
dd = prev_dd
prev_dd = dd
if dd > max_dd: if dd > max_dd:
max_dd = dd max_dd = dd
if dd == 0: if dd == 0:
+12 -3
View File
@@ -61,7 +61,8 @@ class MultiStrategyResult:
total_performance: 组合整体绩效(资金加权收益率 + 策略数 + 总资金)。 total_performance: 组合整体绩效(资金加权收益率 + 策略数 + 总资金)。
individual_results: 每个策略槽位的独立回测结果,key 形如 "{label}@{symbol}" individual_results: 每个策略槽位的独立回测结果,key 形如 "{label}@{symbol}"
equity_allocation: 每个槽位的资金分配比例(均分时各 1/N)。 equity_allocation: 每个槽位的资金分配比例(均分时各 1/N)。
combined_equity: 组合整体净值曲线(各槽位按日期并集 ffill 对齐后求和 combined_equity: 组合整体净值曲线(各槽位按日期并集 ffill 对齐后求和,
晚起步槽位的前导缺口按首个净值=初始资金回填),
列: datetime / total / drawdown / drawdown_pct。 列: datetime / total / drawdown / drawdown_pct。
""" """
@@ -200,7 +201,8 @@ class MultiStrategyEngine:
算法与 ``PortfolioBacktestEngine._build_combined_equity`` 一致: 算法与 ``PortfolioBacktestEngine._build_combined_equity`` 一致:
各策略回测日期范围可能不同(取数差异、停牌),取 datetime 并集, 各策略回测日期范围可能不同(取数差异、停牌),取 datetime 并集,
每个策略的 total 列 forward-fill 对齐到并集后求和得组合总净值, 每个策略的 total 列 forward-fill 对齐到并集后求和得组合总净值,
再算回撤。 再算回撤。前导缺口(晚起步槽位)按其首个净值(=初始资金)回填
(bfill),保证合并曲线首值等于总投入资金。
""" """
del allocations # 资金分配不参与曲线形状(各策略独立 full cash 回测, del allocations # 资金分配不参与曲线形状(各策略独立 full cash 回测,
# 合并的是 normalized 的净值贡献;保持签名与 Portfolio 版一致便于对照) # 合并的是 normalized 的净值贡献;保持签名与 Portfolio 版一致便于对照)
@@ -224,8 +226,15 @@ class MultiStrategyEngine:
if not series_list: if not series_list:
return empty return empty
# 外连接对齐(并集日期):各槽位在缺失日期 forward-fill(持有不动);
# 前导缺口(晚起步槽位)用每列首个有效值回填(bfill)——资金在组合
# 起点即已分配,建仓前按初始资金趴账,与 PortfolioBacktestEngine 及
# 组合 Walk-Forward 的口径一致。此前前导缺口填 0:合并曲线首值会小于
# 总投入资金,total_return 被系统性虚增。
# 退化兜底:整列全 NaN(理论不可达——无数据的槽位不会进入 series_list
# 显式落 0,避免 sum 传播 NaN。
aligned = pd.concat(series_list, axis=1).sort_index() aligned = pd.concat(series_list, axis=1).sort_index()
aligned = aligned.ffill().fillna(0) aligned = aligned.ffill().bfill().fillna(0.0)
total = aligned.sum(axis=1) total = aligned.sum(axis=1)
# 回撤:drawdown 为绝对回撤额(峰值-当前,正值),drawdown_pct 为相对当时 # 回撤:drawdown 为绝对回撤额(峰值-当前,正值),drawdown_pct 为相对当时
+14 -5
View File
@@ -44,7 +44,9 @@ class PortfolioResult:
equity_allocation: 每只标的的资金分配比例 equity_allocation: 每只标的的资金分配比例
combined_equity: 组合整体净值曲线(按日期对齐各标的求和), combined_equity: 组合整体净值曲线(按日期对齐各标的求和),
列: datetime/total/drawdown/drawdown_pct。各标的独立回测日期范围 列: datetime/total/drawdown/drawdown_pct。各标的独立回测日期范围
可能不同,此处按日期并集 forward-fill 对齐后求和 可能不同,此处按日期并集 forward-fill 对齐后求和;晚上市标的的
前导缺口按其首个净值(=初始资金)回填,保证合并曲线首值等于
总投入资金。
trades: 组合层汇总成交(各标的 concat + ``symbol`` 列标注来源标的), trades: 组合层汇总成交(各标的 concat + ``symbol`` 列标注来源标的),
供组合级绩效统计(逐标的 FIFO 配对持仓天数)与前端明细表使用。 供组合级绩效统计(逐标的 FIFO 配对持仓天数)与前端明细表使用。
""" """
@@ -265,7 +267,9 @@ class PortfolioBacktestEngine:
"""把各标的独立净值曲线按日期对齐求和,生成组合整体净值曲线。 """把各标的独立净值曲线按日期对齐求和,生成组合整体净值曲线。
各标的独立回测的日期范围可能不同(取数差异、停牌等),这里取所有标的 各标的独立回测的日期范围可能不同(取数差异、停牌等),这里取所有标的
datetime 的并集,每个标的的 total 列 forward-fill 对齐到并集后求和 datetime 的并集,每个标的的 total 列 forward-fill 对齐到并集后求和
前导缺口(晚上市标的)按其首个净值(=初始资金)回填(bfill),保证
合并曲线首值等于总投入资金。
Returns: Returns:
DataFrame: datetime / total / drawdown / drawdown_pct。 DataFrame: datetime / total / drawdown / drawdown_pct。
@@ -293,10 +297,15 @@ class PortfolioBacktestEngine:
if not series_list: if not series_list:
return empty return empty
# 外连接对齐(并集日期)forward-fill 各标的在缺失日期的净值(持有不动) # 外连接对齐(并集日期):各标的在缺失日期 forward-fill(持有不动)
# 再求和得组合总净值。缺失值填 0 是为应对某标的完全无该日期数据的情况。 # 前导缺口(晚上市 / 取数晚于组合起点)用每列首个有效值回填(bfill)——
# 资金在组合起点即已分配,建仓前按初始资金趴账,与组合 Walk-Forward 的
# ffill().bfill() 口径一致。此前前导缺口填 0:晚上市标的上市前贡献 0,
# 合并曲线首值 < 总投入资金,total_return 被系统性虚增。
# 退化兜底:整列全 NaN(理论不可达——无数据的标的不会进入 series_list
# 显式落 0,避免 sum 传播 NaN。
aligned = pd.concat(series_list, axis=1).sort_index() aligned = pd.concat(series_list, axis=1).sort_index()
aligned = aligned.ffill().fillna(0) aligned = aligned.ffill().bfill().fillna(0.0)
total = aligned.sum(axis=1) total = aligned.sum(axis=1)
# 回撤:drawdown 为绝对回撤额(峰值-当前,正值),drawdown_pct 为相对 # 回撤:drawdown 为绝对回撤额(峰值-当前,正值),drawdown_pct 为相对
+20 -7
View File
@@ -225,16 +225,24 @@ class RotationEngine:
# 1. 推进各标的指针到 ≤ d 的最新一根 # 1. 推进各标的指针到 ≤ d 的最新一根
bar_today: dict[str, pd.Series] = {} bar_today: dict[str, pd.Series] = {}
traded_today: set[str] = set()
for sym, df in self._dfs.items(): for sym, df in self._dfs.items():
dts = self._dt_index(sym) dts = self._dt_index(sym)
while pointers[sym] + 1 < len(dts) and dts[pointers[sym] + 1] <= d: while pointers[sym] + 1 < len(dts) and dts[pointers[sym] + 1] <= d:
pointers[sym] += 1 pointers[sym] += 1
if pointers[sym] >= 0: if pointers[sym] >= 0:
bar_today[sym] = df.iloc[pointers[sym]] bar_today[sym] = df.iloc[pointers[sym]]
# 当日真实有 bar 才可成交;停牌标的只有旧 bar(估值用)
if dts[pointers[sym]] == d:
traded_today.add(sym)
# 2. 次开执行昨日信号(用当日开盘价) # 2. 次开执行挂单(用当日开盘价)。停牌标的当日不可成交,挂单顺延、
# 复牌开盘成交(真实挂单语义);第 4 步排队按 symbol+方向去重,
# 不会与后续新信号重复排队。
still_pending: list[tuple[str, str, str]] = []
for sym, direction, reason in pending: for sym, direction, reason in pending:
if sym not in bar_today: if sym not in traded_today:
still_pending.append((sym, direction, reason))
continue continue
price = float(bar_today[sym]["open"]) price = float(bar_today[sym]["open"])
if not math.isfinite(price) or price <= 0: if not math.isfinite(price) or price <= 0:
@@ -268,7 +276,7 @@ class RotationEngine:
day_i, d_str, sym, "BUY", shares, price, fee, 0.0, reason=reason day_i, d_str, sym, "BUY", shares, price, fee, 0.0, reason=reason
) )
) )
pending = [] pending = still_pending
# 3. 止盈止损检查(收盘口径,次日执行) # 3. 止盈止损检查(收盘口径,次日执行)
for sym in list(positions): for sym in list(positions):
@@ -281,7 +289,8 @@ class RotationEngine:
elif self._take_profit is not None and close >= cost * (1 + self._take_profit): elif self._take_profit is not None and close >= cost * (1 + self._take_profit):
pending.append((sym, "SELL", "take_profit")) pending.append((sym, "SELL", "take_profit"))
# 4. 调仓判定 # 4. 调仓判定(day0 即可产生初始调仓信号,次日开盘执行;排名只用
# 截至 day0 收盘的数据,无未来泄漏)
key = ( key = (
(d.isocalendar()[0], d.isocalendar()[1]) (d.isocalendar()[0], d.isocalendar()[1])
if self._refresh == "weekly" if self._refresh == "weekly"
@@ -289,7 +298,7 @@ class RotationEngine:
) )
is_rebalance = key != prev_key is_rebalance = key != prev_key
prev_key = key prev_key = key
if is_rebalance and day_i >= 1: if is_rebalance:
rebalances.append(d_str) rebalances.append(d_str)
ranked = self._rank_all(pointers, d) ranked = self._rank_all(pointers, d)
top_keep = [s for s, _ in ranked[: self._keep_rank]] top_keep = [s for s, _ in ranked[: self._keep_rank]]
@@ -355,12 +364,16 @@ class RotationEngine:
return pd.DatetimeIndex(self._dfs[sym]["_ts"]) return pd.DatetimeIndex(self._dfs[sym]["_ts"])
def _rank_all(self, pointers: dict[str, int], d: Any) -> list[tuple[str, float]]: def _rank_all(self, pointers: dict[str, int], d: Any) -> list[tuple[str, float]]:
"""对全部标的按截至 d 的前缀数据打分并降序排名。""" """对全部标的按截至 d 的前缀数据打分并降序排名。
历史不足(< 5 根,如次新股)打不出有效分,直接从排名(买入候选)中
剔除而非按 0 分参与排序——0 分会排在负动量标的之前导致误买。持仓
标的买入时即已满足 ≥5 根且指针只进不退,不受影响。
"""
scored: list[tuple[str, float]] = [] scored: list[tuple[str, float]] = []
for sym, df in self._dfs.items(): for sym, df in self._dfs.items():
idx = pointers[sym] idx = pointers[sym]
if idx < 5: if idx < 5:
scored.append((sym, 0.0))
continue continue
start = max(0, idx - self._max_history) start = max(0, idx - self._max_history)
prefix = df.iloc[start : idx + 1].drop(columns=["_ts"], errors="ignore") prefix = df.iloc[start : idx + 1].drop(columns=["_ts"], errors="ignore")
+56 -12
View File
@@ -30,6 +30,7 @@
from __future__ import annotations from __future__ import annotations
import logging
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any from typing import Any
@@ -37,9 +38,15 @@ import numpy as np
import pandas as pd import pandas as pd
from easy_tdx.backtest.engine import BacktestEngine from easy_tdx.backtest.engine import BacktestEngine
from easy_tdx.backtest.performance import PerformanceAnalyzer
from easy_tdx.backtest.strategy import Strategy from easy_tdx.backtest.strategy import Strategy
from easy_tdx.backtest.types import to_json_native from easy_tdx.backtest.types import to_json_native
logger = logging.getLogger(__name__)
# 单个测试窗最少 bar 数(低于此的窗口跳过,不参与评估)
MIN_WINDOW_BARS = 20
__all__ = [ __all__ = [
"WalkForwardWindow", "WalkForwardWindow",
"WalkForwardResult", "WalkForwardResult",
@@ -198,7 +205,7 @@ class WalkForwardEngine:
for i in range(self._n_windows): for i in range(self._n_windows):
s = eval_start + i * window_len s = eval_start + i * window_len
e = s + window_len if i < self._n_windows - 1 else n # 末窗吃到尾部 e = s + window_len if i < self._n_windows - 1 else n # 末窗吃到尾部
if e - s < 5: if e - s < MIN_WINDOW_BARS:
continue continue
win = self._run_window(df, s, e, i) win = self._run_window(df, s, e, i)
if win is not None: if win is not None:
@@ -211,7 +218,9 @@ class WalkForwardEngine:
"""独立回测单个窗口 [s, e)。 """独立回测单个窗口 [s, e)。
带前置上下文(指标预热),用 warmup_bars 压制上下文区间的信号; 带前置上下文(指标预热),用 warmup_bars 压制上下文区间的信号;
窗口起点空仓(每窗独立开仓语义)。 窗口起点空仓(每窗独立开仓语义)。绩效只用窗内净值与成交计算,
上下文预热区不稀释 sharpe/年化/波动等时间口径指标(与组合级
``_ComboWalkForwardBase._run_window`` 的 ``ec.iloc[lead:]`` 同口径)。
""" """
ctx_s = max(0, s - self._context_bars) ctx_s = max(0, s - self._context_bars)
lead = s - ctx_s # 上下文 bar 数 = 需压制的信号数 lead = s - ctx_s # 上下文 bar 数 = 需压制的信号数
@@ -226,9 +235,10 @@ class WalkForwardEngine:
) )
try: try:
bt = engine.run(sub) bt = engine.run(sub)
except Exception: # noqa: BLE001 — 单窗失败不拖垮整组,跳过该窗 except Exception as exc: # noqa: BLE001 — 单窗失败不拖垮整组,跳过该窗
logger.warning("WF 第 %s 窗回测失败,跳过该窗:%s: %s", index, type(exc).__name__, exc)
return None return None
perf = bt.performance perf = self._window_performance(bt, lead)
dt = self._dates(sub, lead) dt = self._dates(sub, lead)
return WalkForwardWindow( return WalkForwardWindow(
@@ -241,21 +251,48 @@ class WalkForwardEngine:
max_drawdown=float(perf.get("max_drawdown", 0.0)), max_drawdown=float(perf.get("max_drawdown", 0.0)),
total_trades=int(perf.get("total_trades", 0)), total_trades=int(perf.get("total_trades", 0)),
win_rate=float(perf.get("win_rate", 0.0)), win_rate=float(perf.get("win_rate", 0.0)),
performance={k: v for k, v in perf.items()}, performance=perf,
) )
@staticmethod
def _window_performance(bt: Any, lead: int) -> dict[str, Any]:
"""只用窗内净值 + 成交重算绩效(上下文预热区不参与窗指标)。
上下文区净值恒为初始现金(warmup 压制信号),窗口内回撤/收益不变,
但 sharpe/年化/波动等按全序列(含上下文)计算会被零收益段稀释。
"""
equity = bt.equity_curve
if len(equity) <= lead:
return dict(bt.performance)
window_equity = equity.iloc[lead:].reset_index(drop=True)
return dict(PerformanceAnalyzer(equity_curve=window_equity, trades=bt.trades).compute())
@staticmethod @staticmethod
def _dates(sub: pd.DataFrame, lead: int) -> tuple[str, str]: def _dates(sub: pd.DataFrame, lead: int) -> tuple[str, str]:
"""取窗口起止日期(跳过 lead 根上下文)。""" """取窗口起止日期(跳过 lead 根上下文)。
int/np 整数(YYYYMMDDTDX 日线原样)先 str 再按 %Y%m%d 解析——
直接 ``pd.Timestamp(int)`` 会被当纳秒换算成 1970 年。
"""
col = "datetime" if "datetime" in sub.columns else "date" col = "datetime" if "datetime" in sub.columns else "date"
vals = sub[col].iloc[lead:] vals = sub[col].iloc[lead:]
if len(vals) == 0: if len(vals) == 0:
return "", "" return "", ""
return ( return (
pd.Timestamp(vals.iloc[0]).strftime("%Y-%m-%d"), WalkForwardEngine._fmt_date(vals.iloc[0]),
pd.Timestamp(vals.iloc[-1]).strftime("%Y-%m-%d"), WalkForwardEngine._fmt_date(vals.iloc[-1]),
) )
@staticmethod
def _fmt_date(v: Any) -> str:
"""单个日期值 → YYYY-MM-DDint/float YYYYMMDD 与 Timestamp/datetime64 兼容)。"""
ts: str
if isinstance(v, int | float | np.integer | np.floating) and not isinstance(v, bool):
ts = str(pd.to_datetime(str(int(v)), format="%Y%m%d").strftime("%Y-%m-%d"))
else:
ts = str(pd.Timestamp(v).strftime("%Y-%m-%d"))
return ts
@staticmethod @staticmethod
def _aggregate(result: WalkForwardResult) -> None: def _aggregate(result: WalkForwardResult) -> None:
"""聚合各窗指标(空列表安全)。""" """聚合各窗指标(空列表安全)。"""
@@ -270,7 +307,8 @@ class WalkForwardEngine:
result.worst_window = float(np.min(rets)) result.worst_window = float(np.min(rets))
result.best_window = float(np.max(rets)) result.best_window = float(np.max(rets))
result.mean_sharpe = float(np.mean([w.sharpe for w in ws])) result.mean_sharpe = float(np.mean([w.sharpe for w in ws]))
result.worst_drawdown = float(min(w.max_drawdown for w in ws)) # max_drawdown 为正数幅度((peak-total)/peak),“最差窗回撤”应取最大值
result.worst_drawdown = float(max(w.max_drawdown for w in ws))
result.total_trades = int(sum(w.total_trades for w in ws)) result.total_trades = int(sum(w.total_trades for w in ws))
@@ -352,7 +390,7 @@ class _ComboWalkForwardBase:
for i in range(self._n_windows): for i in range(self._n_windows):
s = eval_start + i * window_len s = eval_start + i * window_len
e = s + window_len if i < self._n_windows - 1 else n # 末窗吃到尾部 e = s + window_len if i < self._n_windows - 1 else n # 末窗吃到尾部
if e - s < 5: if e - s < MIN_WINDOW_BARS:
continue continue
win = self._run_window(timeline, s, e, i) win = self._run_window(timeline, s, e, i)
if win is not None: if win is not None:
@@ -413,7 +451,14 @@ class _ComboWalkForwardBase:
) )
try: try:
bt = engine.run(sub) bt = engine.run(sub)
except Exception: # noqa: BLE001 — 单槽位失败不拖垮整窗 except Exception as exc: # noqa: BLE001 — 单槽位失败不拖垮整窗
logger.warning(
"WF 组合第 %s 窗槽位 %s 回测失败,跳过:%s: %s",
index,
slot.key,
type(exc).__name__,
exc,
)
continue continue
# 只取窗内净值点(上下文区恒为现金,不参与窗指标,避免稀释波动率) # 只取窗内净值点(上下文区恒为现金,不参与窗指标,避免稀释波动率)
@@ -453,7 +498,6 @@ class _ComboWalkForwardBase:
if trade_frames if trade_frames
else pd.DataFrame(columns=["symbol", "direction", "pnl", "rejected"]) else pd.DataFrame(columns=["symbol", "direction", "pnl", "rejected"])
) )
from easy_tdx.backtest.performance import PerformanceAnalyzer
perf = PerformanceAnalyzer(equity_curve=window_equity, trades=all_trades).compute() perf = PerformanceAnalyzer(equity_curve=window_equity, trades=all_trades).compute()
+14
View File
@@ -174,6 +174,20 @@ def parse_xml(text: str) -> list[dict[str, Any]]:
"product": _g(node, "productid"), "product": _g(node, "productid"),
} }
if not cells:
if len(root) > 0:
# XML 合法但 0 个 <data> 数据节点,且根元素下存在其他子结构——
# 官网模板/字段改版的典型形态(正常发布日必有 <data>;真实
# 「无数据日」走 302 → CcpmNoDataError)。静默返回空表会让
# 改版长期伪装成"无数据",这里显式报错。
tags = ",".join(sorted({child.tag for child in root}))[:200]
raise CcpmError(
f"XML 解析出 0 条 <data> 数据节点(根元素 <{root.tag}>"
f"子节点: {tags})——中金所页面结构可能已变更"
)
# 根元素无任何子节点(如空 <positionRank/>):按当日无数据处理
return []
rows: list[dict[str, Any]] = [] rows: list[dict[str, Any]] = []
ranks = sorted({r for (_, _, r) in cells}) ranks = sorted({r for (_, _, r) in cells})
for instrument in sorted(meta): for instrument in sorted(meta):
+14 -2
View File
@@ -23,6 +23,17 @@ def formula() -> None:
"""通达信公式:计算 / 选股 / 回测(命名布尔输出即信号)。""" """通达信公式:计算 / 选股 / 回测(命名布尔输出即信号)。"""
def _parse_symbol(sym: str) -> tuple[str, str]:
"""解析 ``市场:代码``;格式不对抛 click.BadParameter(而非裸 ValueError)。"""
market, sep, code = sym.strip().partition(":")
if not sep or not market.strip() or not code.strip() or ":" in code:
raise click.BadParameter(
f"标的格式应为 市场:代码(如 SH:600519),收到: {sym!r}",
param_hint="--symbols",
)
return market.strip().upper(), code.strip()
def _load_formula(text: str | None, file: str | None) -> str: def _load_formula(text: str | None, file: str | None) -> str:
from easy_tdx.formula import compile_formula from easy_tdx.formula import compile_formula
@@ -130,12 +141,13 @@ def formula_screen(
] ]
else: else:
symbol_list = [s.strip() for s in symbols.split(",") if s.strip()] symbol_list = [s.strip() for s in symbols.split(",") if s.strip()]
# 前置校验全部标的格式(缺冒号等在发请求前就报错,而非循环中途裸崩)
symbol_pairs = [_parse_symbol(s) for s in symbol_list]
compiled = compile_formula(source) compiled = compile_formula(source)
hits: list[dict[str, Any]] = [] hits: list[dict[str, Any]] = []
errors: list[dict[str, str]] = [] errors: list[dict[str, str]] = []
for sym in symbol_list: for sym, (market, code) in zip(symbol_list, symbol_pairs, strict=True):
market, code = sym.split(":", 1)
try: try:
df = _fetch(market, code, count, adjust) df = _fetch(market, code, count, adjust)
if df is None or len(df) == 0: if df is None or len(df) == 0:
+4
View File
@@ -561,6 +561,10 @@ def _sync_one_daily(client: TdxClient, filepath: Path) -> tuple[int, str]:
# 协议返回的成交量单位是股;encode_daily_bar 在 vol_coeff=0.01 时按 ×100 # 协议返回的成交量单位是股;encode_daily_bar 在 vol_coeff=0.01 时按 ×100
# 写入(.day 原始字段为股,读取端 ×0.01 还原为手),故写入前须换算为手。 # 写入(.day 原始字段为股,读取端 ×0.01 还原为手),故写入前须换算为手。
# 否则日成交 > 4295 万股(如招商银行等大盘股)的 bar 会溢出 uint32。 # 否则日成交 > 4295 万股(如招商银行等大盘股)的 bar 会溢出 uint32。
# 深市基金/ETF 同口径已实测锚定(2026-09-06):sz159915.day 2026-06-12
# 原始 vol=1,249,608,655 与 get_security_bars 返回值完全一致,且
# amount/vol≈close——.day 原始字段与协议 vol 同为「股」,SZ_FUND
# vol_coeff=0.01)走本分支换算正确,无需例外。
if vol_coeff == 0.01: if vol_coeff == 0.01:
for b in bars: for b in bars:
b.vol /= 100 b.vol /= 100
+55 -12
View File
@@ -20,6 +20,32 @@ from typing import Any
import click import click
#: sync 支持的周期(与 Period 枚举名对齐,仓库按同名键存储)。
_PERIOD_CHOICES = ["DAILY", "WEEKLY", "MONTHLY", "MIN_1", "MIN_5", "MIN_15", "MIN_30", "MIN_60"]
#: baostock 是 EOD 源,只覆盖日线及以上(与 sources/baostock.py 的能力边界一致)。
_BAOSTOCK_PERIODS = ("DAILY", "WEEKLY", "MONTHLY")
def _parse_symbol(sym: str) -> tuple[str, str]:
"""解析 ``市场:代码``;格式不对抛 click.BadParameter(而非裸 ValueError)。"""
market, sep, code = sym.strip().partition(":")
if not sep or not market.strip() or not code.strip() or ":" in code:
raise click.BadParameter(
f"标的格式应为 市场:代码(如 SH:600519),收到: {sym!r}",
param_hint="--symbols",
)
return market.strip().upper(), code.strip()
def _validate_source_period(source: str, period: str) -> None:
"""--source baostock 不支持的周期在参数层直接报错(而非静默空转 exit 0)。"""
if source == "baostock" and period.upper() not in _BAOSTOCK_PERIODS:
raise click.BadParameter(
f"--source baostock 仅支持日线及以上周期 {'/'.join(_BAOSTOCK_PERIODS)}"
f"收到: {period}baostock 为 EOD 源,无分钟线)",
param_hint="--period",
)
def _require_warehouse(db_path: str | None) -> Any: def _require_warehouse(db_path: str | None) -> Any:
"""惰性导入 warehouse(duckdb 可选依赖),失败给友好错误。""" """惰性导入 warehouse(duckdb 可选依赖),失败给友好错误。"""
@@ -40,7 +66,13 @@ def warehouse() -> None:
@click.option( @click.option(
"--symbols", required=True, help="标的列表:逗号分隔(SH:600519,SZ:000001)或 @文件(每行一个)" "--symbols", required=True, help="标的列表:逗号分隔(SH:600519,SZ:000001)或 @文件(每行一个)"
) )
@click.option("--period", default="DAILY", help="K 线周期(默认 DAILY") @click.option(
"--period",
"period",
default="DAILY",
type=click.Choice(_PERIOD_CHOICES, case_sensitive=False),
help="K 线周期(默认 DAILY;baostock 源仅支持日线及以上)",
)
@click.option("--max-bars", default=8000, type=int, help="首同步最大拉取根数(默认 8000") @click.option("--max-bars", default=8000, type=int, help="首同步最大拉取根数(默认 8000")
@click.option("--tail-bars", default=15, type=int, help="增量同步尾部根数(默认 15") @click.option("--tail-bars", default=15, type=int, help="增量同步尾部根数(默认 15")
@click.option( @click.option(
@@ -72,6 +104,7 @@ def warehouse_sync(
db_path: str | None, db_path: str | None,
) -> None: ) -> None:
"""增量同步行情进仓库(首同步全量、此后只补尾部)。""" """增量同步行情进仓库(首同步全量、此后只补尾部)。"""
_validate_source_period(source, period)
if symbols.startswith("@"): if symbols.startswith("@"):
from pathlib import Path from pathlib import Path
@@ -86,6 +119,9 @@ def warehouse_sync(
] ]
else: else:
symbol_list = [s.strip() for s in symbols.split(",") if s.strip()] symbol_list = [s.strip() for s in symbols.split(",") if s.strip()]
# 前置校验全部标的格式(缺冒号等在发请求前就报错,而非循环中途裸崩)
for sym in symbol_list:
_parse_symbol(sym)
with _require_warehouse(db_path) as wh: with _require_warehouse(db_path) as wh:
from easy_tdx.warehouse import WarehouseSyncer from easy_tdx.warehouse import WarehouseSyncer
@@ -100,7 +136,7 @@ def warehouse_sync(
syncer = WarehouseSyncer( syncer = WarehouseSyncer(
kline_client, wh, max_bars=max_bars, tail_bars=tail_bars, adjust=adjust kline_client, wh, max_bars=max_bars, tail_bars=tail_bars, adjust=adjust
) )
summary = syncer.sync(symbol_list, period=period, progress=_progress) summary = syncer.sync(symbol_list, period=period.upper(), progress=_progress)
else: else:
from ..cli.conn import get_mac_client from ..cli.conn import get_mac_client
@@ -115,13 +151,18 @@ def warehouse_sync(
syncer = WarehouseSyncer( syncer = WarehouseSyncer(
kline_client, wh, max_bars=max_bars, tail_bars=tail_bars, adjust=adjust kline_client, wh, max_bars=max_bars, tail_bars=tail_bars, adjust=adjust
) )
summary = syncer.sync(symbol_list, period=period, progress=_progress) summary = syncer.sync(symbol_list, period=period.upper(), progress=_progress)
click.echo( # source 口径标注进 summary(与 /bars 响应的 source 字段呼应)
json.dumps( payload = {k: v for k, v in summary.items() if k != "details"}
{k: v for k, v in summary.items() if k != "details"}, payload["source"] = source
ensure_ascii=False, # 失败明细打到 stderr(JSON 主体保持机器可读)
) for d in summary.get("details", []):
) if isinstance(d, dict) and d.get("error"):
click.echo(f" {d['symbol']}: {d['error']}", err=True)
click.echo(json.dumps(payload, ensure_ascii=False))
if summary.get("failed"):
click.echo(f"错误: {summary['failed']}/{summary['total']} 个标的同步失败", err=True)
raise SystemExit(1)
@warehouse.command("query") @warehouse.command("query")
@@ -189,11 +230,13 @@ def warehouse_check(symbols: str | None, db_path: str | None) -> None:
with _require_warehouse(db_path) as wh: with _require_warehouse(db_path) as wh:
market = code = None market = code = None
if symbols: if symbols:
first = [s.strip() for s in symbols.split(",") if s.strip()][0] # 先校验「只支持一个标的」,再解析(旧码顺序颠倒:含冒号的多标的
market, code = first.split(":", 1) # 绕过校验后在 split 处裸崩;无冒号单标的直接 ValueError)
if ":" not in symbols and len(symbols.split(",")) > 1: sym_list = [s.strip() for s in symbols.split(",") if s.strip()]
if len(sym_list) > 1:
click.echo("错误: --symbols 自检模式一次只支持一个标的", err=True) click.echo("错误: --symbols 自检模式一次只支持一个标的", err=True)
raise SystemExit(1) raise SystemExit(1)
market, code = _parse_symbol(sym_list[0])
report = wh.health_check(market=market, code=code) report = wh.health_check(market=market, code=code)
click.echo(json.dumps(report, ensure_ascii=False, default=str)) click.echo(json.dumps(report, ensure_ascii=False, default=str))
if report["issues"]: if report["issues"]:
+43 -10
View File
@@ -43,6 +43,10 @@ import pandas as pd
__all__ = ["FormulaError", "FormulaResult", "CompiledFormula", "compile_formula"] __all__ = ["FormulaError", "FormulaResult", "CompiledFormula", "compile_formula"]
# 表达式嵌套深度上限(递归下降防爆栈:每层约 8 个 Python 栈帧,100 层
# 远低于 CPython 默认递归上限,超出按 FormulaError 报错而非 RecursionError
_MAX_EXPRESSION_DEPTH = 100
# ── Token ───────────────────────────────────────────────────────────────────── # ── Token ─────────────────────────────────────────────────────────────────────
_TOKEN_RE = re.compile( _TOKEN_RE = re.compile(
@@ -108,6 +112,7 @@ class _Parser:
def __init__(self, tokens: list[_Token]) -> None: def __init__(self, tokens: list[_Token]) -> None:
self._tokens = tokens self._tokens = tokens
self._i = 0 self._i = 0
self._depth = 0
def _peek(self) -> _Token: def _peek(self) -> _Token:
return self._tokens[self._i] return self._tokens[self._i]
@@ -172,7 +177,13 @@ class _Parser:
# 表达式优先级:OR < AND < 比较 < 加减 < 乘除 < 一元 < 原子 # 表达式优先级:OR < AND < 比较 < 加减 < 乘除 < 一元 < 原子
def parse_expression(self) -> _Node: def parse_expression(self) -> _Node:
return self._parse_or() self._depth += 1
if self._depth > _MAX_EXPRESSION_DEPTH:
raise FormulaError(f"公式嵌套过深(超过 {_MAX_EXPRESSION_DEPTH} 层)")
try:
return self._parse_or()
finally:
self._depth -= 1
def _parse_or(self) -> _Node: def _parse_or(self) -> _Node:
left = self._parse_and() left = self._parse_and()
@@ -213,15 +224,22 @@ class _Parser:
return left return left
def _parse_unary(self) -> _Node: def _parse_unary(self) -> _Node:
if tok := self._match_op("-", "+"): tok = self._match_op("-", "+", "!", "NOT")
child = self._parse_unary() if tok is None:
if tok.value == "-": return self._parse_primary()
return _Node(kind="un", value="neg", children=[child]) # 一元运算符链也计入深度(防 "!!!!…" 型超长链爆栈)
return child self._depth += 1
if tok := self._match_op("!", "NOT"): if self._depth > _MAX_EXPRESSION_DEPTH:
raise FormulaError(f"公式嵌套过深(超过 {_MAX_EXPRESSION_DEPTH} 层)")
try:
child = self._parse_unary() child = self._parse_unary()
finally:
self._depth -= 1
if tok.value == "-":
return _Node(kind="un", value="neg", children=[child])
if tok.value in ("!", "NOT"):
return _Node(kind="un", value="not", children=[child]) return _Node(kind="un", value="not", children=[child])
return self._parse_primary() return child # 一元正号
def _parse_primary(self) -> _Node: def _parse_primary(self) -> _Node:
tok = self._peek() tok = self._peek()
@@ -319,6 +337,18 @@ def _build_functions() -> dict[str, Callable[..., Any]]:
): ):
if hasattr(mytt, name): if hasattr(mytt, name):
fns[name] = getattr(mytt, name) fns[name] = getattr(mytt, name)
# REF 负移位 = 引用未来数据(未来函数),显式禁止。此前仅靠负数字面量
# 经一元负号转成 float 在 pandas 层报错这一巧合拦截。MyTT 库内直调
# (如 ICHIMOKU 迟行带画图 REF(C, -SHIFT))不走公式白名单,不受影响。
def _ref_no_lookahead(S: Any, N: Any = 1) -> Any:
n_arr = np.asarray(N)
if n_arr.size and float(np.min(n_arr)) < 0:
raise FormulaError(f"REF 不允许负移位(引用未来数据):N={N}")
return mytt.REF(S, N)
fns["REF"] = _ref_no_lookahead
# numpy 补齐(TDX 语义) # numpy 补齐(TDX 语义)
fns["POW"] = np.power fns["POW"] = np.power
fns["SQRT"] = np.sqrt fns["SQRT"] = np.sqrt
@@ -373,7 +403,10 @@ class _Evaluator:
@staticmethod @staticmethod
def _is_boolean(expr: _Node, val: Any) -> bool: def _is_boolean(expr: _Node, val: Any) -> bool:
"""输出归类:比较/逻辑/CROSS 节点或 0/1 值域 → 信号列。""" """输出归类:比较/逻辑/CROSS 节点 → 信号列;否则仅当全部有限值
{0.0, 1.0} 才兜底判为信号(真 0/1 布尔指标)——含其他小数的
连续值(价格比率、归一化振荡器等)一律归数值列。
"""
if expr.kind in ("cmp", "logic"): if expr.kind in ("cmp", "logic"):
return True return True
if expr.kind == "call" and expr.value in _BOOL_FUNCS: if expr.kind == "call" and expr.value in _BOOL_FUNCS:
@@ -382,7 +415,7 @@ class _Evaluator:
finite = arr[np.isfinite(arr)] finite = arr[np.isfinite(arr)]
if finite.size == 0: if finite.size == 0:
return False return False
return bool(finite.min() >= 0.0 and finite.max() <= 1.0) return bool(np.all((finite == 0.0) | (finite == 1.0)))
def eval(self, node: _Node) -> Any: def eval(self, node: _Node) -> Any:
if node.kind == "num": if node.kind == "num":
+6 -2
View File
@@ -34,8 +34,12 @@ UNUSUAL_TYPE_NAMES: dict[int, str] = {
} }
def _describe_unusual(unusual_type: int, data: bytes, hour: int = 9) -> tuple[str, str]: def _describe_unusual(unusual_type: int, data: bytes, hour: int) -> tuple[str, str]:
"""根据异动类型解析描述和数值。hour 用于区分竞价/尾盘双时刻信号(0x15)。""" """根据异动类型解析描述和数值。
hour 必传:来自报文时间槽(offset 29),用于区分竞价/尾盘双时刻信号
(0x15)——缺省会按 9 点把 15:00 的收盘信号误标成「竞价」。
"""
if len(data) < 13: if len(data) < 13:
return "", "" return "", ""
v1, v2, v3, v4 = struct.unpack_from("<B2fI", data) v1, v2, v3, v4 = struct.unpack_from("<B2fI", data)
+3 -1
View File
@@ -145,7 +145,9 @@ def detect_ex_dividend_gaps(
with np.errstate(divide="ignore", invalid="ignore"): with np.errstate(divide="ignore", invalid="ignore"):
ratio = open_arr[1:] / prev_close[:-1] - 1.0 ratio = open_arr[1:] / prev_close[:-1] - 1.0
out: list[str] = [] out: list[str] = []
for i in np.where(~np.isfinite(ratio) | (ratio < threshold))[0]: # 仅「有限且低于阈值」计跳空:前收缺失(NaN/0)产生的非有限比率是数据
# 缺口而非除权(除权跳空必然是有限值),计入只会制造假 unexplained_gap。
for i in np.where(np.isfinite(ratio) & (ratio < threshold))[0]:
out.append(_fmt(df["_dt"].iloc[i + 1])) out.append(_fmt(df["_dt"].iloc[i + 1]))
return out return out
+10 -3
View File
@@ -212,7 +212,11 @@ class RealtimeDataFeed:
:meth:`stop`。 :meth:`stop`。
""" """
if self._stop_requested: if self._stop_requested:
return # 启动前已请求停止(见 __init__ 的竞态说明) # 启动前已请求停止(见 __init__ 的竞态说明):消费该一次性标志并
# 立即退出;复位后同一实例可再次 start(修复 start→stop→start
# 二次启动被永久静默吞掉的问题)。
self._stop_requested = False
return
self._running = True self._running = True
try: try:
count = 0 count = 0
@@ -262,7 +266,9 @@ class RealtimeDataFeed:
async def _run_sync_loop(self, client: Any, max_iterations: int | None) -> None: async def _run_sync_loop(self, client: Any, max_iterations: int | None) -> None:
"""同步客户端的轮询循环:阻塞调用丢到 executor。""" """同步客户端的轮询循环:阻塞调用丢到 executor。"""
if self._stop_requested: if self._stop_requested:
return # 启动前已请求停止(见 __init__ 的竞态说明 # 同 run_async:消费一次性停止标志并复位(实例可复用
self._stop_requested = False
return
self._running = True self._running = True
try: try:
count = 0 count = 0
@@ -279,7 +285,8 @@ class RealtimeDataFeed:
"""请求停止轮询(下一轮 sleep 结束后生效)。 """请求停止轮询(下一轮 sleep 结束后生效)。
在 ``run_async`` / ``run_sync`` 首次获得调度之前调用同样有效 在 ``run_async`` / ``run_sync`` 首次获得调度之前调用同样有效
(启动即退出),见 ``__init__`` 的竞态说明。 (启动即退出);该停止请求是一次性的——被某次 run_* 消费后,
同一实例可以再次 start。
""" """
self._stop_requested = True self._stop_requested = True
self._running = False self._running = False
+18 -7
View File
@@ -17,12 +17,22 @@
from __future__ import annotations from __future__ import annotations
from datetime import datetime, time, tzinfo from datetime import datetime, time, timedelta, timezone, tzinfo
from typing import Any from typing import Any
__all__ = ["SESSION_WINDOWS", "SESSION_DESC", "is_trading_time", "session_info"] __all__ = [
"SHANGHAI_TZ",
"SESSION_WINDOWS",
"SESSION_DESC",
"is_trading_time",
"session_info",
]
#: 有效行情时段(本地时间)。窗口 = (start, end),含两端。 #: A 股行情统一按沪市时区判断。中国无夏令时,固定 UTC+8 即可精确表达,
#: 不依赖系统时区/zoneinfo 数据库(Windows 无 IANA tzdata)。
SHANGHAI_TZ = timezone(timedelta(hours=8), "Asia/Shanghai")
#: 有效行情时段(沪市时间)。窗口 = (start, end),含两端。
#: - 早盘 09:15:00-11:30:3009:15 起集合竞价可看,11:30:30 容纳尾单撮合散点; #: - 早盘 09:15:00-11:30:3009:15 起集合竞价可看,11:30:30 容纳尾单撮合散点;
#: - 午盘 13:00:00-15:05:0015:00-15:03 为收盘集合竞价,留 2 分钟余量。 #: - 午盘 13:00:00-15:05:0015:00-15:03 为收盘集合竞价,留 2 分钟余量。
SESSION_WINDOWS: tuple[tuple[time, time], ...] = ( SESSION_WINDOWS: tuple[tuple[time, time], ...] = (
@@ -42,13 +52,14 @@ def is_trading_time(now: datetime | None = None, *, tz: tzinfo | None = None) ->
而节假日行情本就不动,手动刷新始终可用)。 而节假日行情本就不动,手动刷新始终可用)。
Args: Args:
now: 待判断时间,None = 取本地当前时间。 now: 待判断时间,None = 取当前时间。
tz: 未传 ``now`` 时使用的时区,None = 系统本地时区。 tz: 未传 ``now`` 时使用的时区,None = :data:`SHANGHAI_TZ`(与主机
时区无关;非中国时区的服务器/海外机器不会错位)。
Returns: Returns:
True = 盘中(含集合竞价缓冲窗)。 True = 盘中(含集合竞价缓冲窗)。
""" """
t = now or datetime.now(tz=tz) t = now or datetime.now(tz=tz or SHANGHAI_TZ)
if t.weekday() >= 5: # 周六/周日 if t.weekday() >= 5: # 周六/周日
return False return False
for start, end in SESSION_WINDOWS: for start, end in SESSION_WINDOWS:
@@ -62,7 +73,7 @@ def session_info(now: datetime | None = None, *, tz: tzinfo | None = None) -> di
前端以本地判断为主(每 15s 重估),本接口用于校准服务器侧视角。 前端以本地判断为主(每 15s 重估),本接口用于校准服务器侧视角。
""" """
t = now or datetime.now(tz=tz) t = now or datetime.now(tz=tz or SHANGHAI_TZ)
return { return {
"is_trading_time": is_trading_time(t), "is_trading_time": is_trading_time(t),
"sessions": [ "sessions": [
+52 -40
View File
@@ -17,7 +17,6 @@
from __future__ import annotations from __future__ import annotations
import math
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
@@ -34,13 +33,25 @@ __all__ = [
] ]
def _round_price(x: float) -> float: def _to_cents(price: float) -> int:
"""四舍五入到分(Python round 是银行家舍入,交易所是四舍五入,不能混用)""" """元 → 分。.day 价格本身按 ×100 存 uint,round 消除读回的浮点表示误差"""
return math.floor(x * 100 + 0.5) / 100 return int(round(price * 100))
def _eq_price(a: float, b: float) -> bool: def _limit_price_cents(prev_cents: int, pct: int) -> int:
return abs(a - b) < 1e-4 """交易所涨跌停价(分):前收 × (1 ± pct%),四舍五入到分(半进位)。
纯整数运算 ``(prev_cents * (100 + pct) + 50) // 100``与交易所逐价位
对账零差异不能用 float 乘完再 ``floor(x*100+0.5)``乘法在半分边界
受浮点表示误差影响±10% 67/318 ±5% 90/884 个价位会算低
1 33.05×1.1 误算 36.35交易所 36.36导致真实涨跌停被漏判
"""
return (prev_cents * (100 + pct) + 50) // 100
def _limit_price(prev: float, pct: int) -> float:
"""交易所涨跌停价(元):pct 为整数百分数(正=涨停档,负=跌停档)。"""
return _limit_price_cents(_to_cents(prev), pct) / 100.0
def _limit_ratio(code: str) -> float: def _limit_ratio(code: str) -> float:
@@ -113,27 +124,29 @@ def _entry_from_closes(
pct = (closes[-1] / prev - 1.0) * 100.0 pct = (closes[-1] / prev - 1.0) * 100.0
entry = LimitUpEntry(code=code, market=market, pct=round(pct, 2)) entry = LimitUpEntry(code=code, market=market, pct=round(pct, 2))
up_ratio = _limit_ratio(code) up_pct = 20 if _limit_ratio(code) == 0.20 else 10
limit_up_price = _round_price(prev * (1 + up_ratio)) # 全程分币整数比较,杜绝浮点舍入在半分边界错 1 分(漏判涨跌停)
cents = [_to_cents(c) for c in closes]
high_cents = _to_cents(last_high)
prev_c = cents[-2]
limit_up_c = _limit_price_cents(prev_c, up_pct)
limit_down_c = _limit_price_cents(prev_c, -up_pct)
# 主板 5%:疑似 ST 涨停。低价股(< 3 元)最小报价单位 0.01 占比过大, # 主板 5%:疑似 ST 涨停。低价股(< 3 元)最小报价单位 0.01 占比过大,
# +5% 整的巧合概率骤增,跳过 ST 判定(宁可漏报不误报)。 # +5% 整的巧合概率骤增,跳过 ST 判定(宁可漏报不误报)。
st_applicable = up_ratio == 0.10 and prev >= 3.0 st_applicable = up_pct == 10 and prev >= 3.0
st_price = _round_price(prev * 1.05) if st_applicable else None st_price_c = _limit_price_cents(prev_c, 5) if st_applicable else None
limit_down_price = _round_price(prev * (1 - up_ratio)) st_down_price_c = _limit_price_cents(prev_c, -5) if st_applicable else None
st_down_price = _round_price(prev * 0.95) if st_applicable else None
def _eq(a: float, b: float) -> bool:
return abs(a - b) < 1e-4
def _is_up(i: int) -> bool: def _is_up(i: int) -> bool:
"""第 i 根是否涨停(用第 i-1 根收盘作前收)。""" """第 i 根是否涨停(用第 i-1 根收盘作前收ST/3 元门槛逐 bar 判定,
避免按最新前收定性整段历史在价格穿越 3 元时漏计/多计"""
if i < 1: if i < 1:
return False return False
p = closes[i - 1] p_c = cents[i - 1]
c = closes[i] c_c = cents[i]
if _eq(c, _round_price(p * (1 + up_ratio))): if c_c == _limit_price_cents(p_c, up_pct):
return True return True
return st_applicable and _eq(c, _round_price(p * 1.05)) return up_pct == 10 and p_c >= 300 and c_c == _limit_price_cents(p_c, 5)
# 连板高度(截至最后一根) # 连板高度(截至最后一根)
streak = 0 streak = 0
@@ -142,28 +155,28 @@ def _entry_from_closes(
streak += 1 streak += 1
i -= 1 i -= 1
entry.streak = streak entry.streak = streak
entry.st = bool(streak > 0 and st_price is not None and _eq(closes[-1], st_price)) entry.st = bool(streak > 0 and st_price_c is not None and cents[-1] == st_price_c)
if streak > 0: if streak > 0:
entry.blown = False entry.blown = False
return entry return entry
# 未封住的场合:炸板(high 触及涨停价)或跌停 # 未封住的场合:炸板(high 触及涨停价)或跌停
if _eq(last_high, limit_up_price): # 口径说明:炸板仅按 10%/20% 档判定——.day 文件无法识别 ST,若对主板
# 额外按 5% 判炸板,非 ST 股恰好摸到 +5.00% 的会误报,故维持漏报方向。
if high_cents == limit_up_c:
entry.blown = True entry.blown = True
return entry return entry
if _eq(closes[-1], limit_down_price) or ( if cents[-1] == limit_down_c or (st_down_price_c is not None and cents[-1] == st_down_price_c):
st_down_price is not None and _eq(closes[-1], st_down_price)
):
down_streak = 0 down_streak = 0
j = len(closes) - 1 j = len(closes) - 1
while j >= 1: while j >= 1:
p = closes[j - 1] p_c = cents[j - 1]
c = closes[j] c_c = cents[j]
hit = _eq(c, _round_price(p * (1 - up_ratio))) hit = c_c == _limit_price_cents(p_c, -up_pct) or (
if not hit and st_applicable: up_pct == 10 and p_c >= 300 and c_c == _limit_price_cents(p_c, -5)
hit = _eq(c, _round_price(p * 0.95)) )
if not hit: if not hit:
break break
down_streak += 1 down_streak += 1
@@ -300,22 +313,21 @@ def compute_limitup_history(
n_files += 1 n_files += 1
if n_files >= max_files: if n_files >= max_files:
break break
up_ratio = _limit_ratio(code) up_pct = 20 if _limit_ratio(code) == 0.20 else 10
closes = [b.close for b in tail] cents = [_to_cents(b.close) for b in tail]
date_ints = [b.year * 10000 + b.month * 100 + b.day for b in tail] date_ints = [b.year * 10000 + b.month * 100 + b.day for b in tail]
for i in range(1, len(tail)): for i in range(1, len(tail)):
p, c = closes[i - 1], closes[i] p_c, c_c = cents[i - 1], cents[i]
if p <= 0: if p_c <= 0:
continue continue
st_applicable = up_ratio == 0.10 and p >= 3.0
d = date_ints[i] d = date_ints[i]
bucket = counts.setdefault(d, {"limit_up": 0, "limit_down": 0}) bucket = counts.setdefault(d, {"limit_up": 0, "limit_down": 0})
if _eq_price(c, _round_price(p * (1 + up_ratio))) or ( if c_c == _limit_price_cents(p_c, up_pct) or (
st_applicable and _eq_price(c, _round_price(p * 1.05)) up_pct == 10 and p_c >= 300 and c_c == _limit_price_cents(p_c, 5)
): ):
bucket["limit_up"] += 1 bucket["limit_up"] += 1
elif _eq_price(c, _round_price(p * (1 - up_ratio))) or ( elif c_c == _limit_price_cents(p_c, -up_pct) or (
st_applicable and _eq_price(c, _round_price(p * 0.95)) up_pct == 10 and p_c >= 300 and c_c == _limit_price_cents(p_c, -5)
): ):
bucket["limit_down"] += 1 bucket["limit_down"] += 1
if n_files >= max_files: if n_files >= max_files:
+49 -27
View File
@@ -8,20 +8,28 @@
即启用未安装时本模块整体静默关闭核心功能零影响 即启用未安装时本模块整体静默关闭核心功能零影响
- baostock 客户端是单条全局连接且非线程安全本模块内部全程持锁串行 - baostock 客户端是单条全局连接且非线程安全本模块内部全程持锁串行
async 调用方经 ``asyncio.to_thread`` 使用 async 调用方经 ``asyncio.to_thread`` 使用
- 数据口径volume 为股 /bars 输出契约一致无需换算停牌日 - 数据口径个股 volume 为股 /bars 输出契约一致无需换算指数
tradestatus=0 volume=0剔除与通达信 K 线不含停牌日的口径对齐 ``is_index=True`` 显式声明后 vol ÷100 转为手baostock 指数 volume
复权经 adjustflag 原生支持QFQ/HFQ/NONENorth ExchangeBJ不覆盖 单位为股 /bars/index 契约为手停牌日tradestatus=0
volume=0剔除与通达信 K 线不含停牌日的口径对齐复权经 adjustflag
原生支持QFQ/HFQ/NONENorth ExchangeBJ不覆盖
- 拉取失败登录失败 / 查询 error_code0 warning 日志并抛
``RuntimeError``auto 兜底路径以 except 包裹调用不受影响
``--source baostock`` 显式使用时不会被伪装成"无数据"
""" """
from __future__ import annotations from __future__ import annotations
import importlib import importlib
import logging
import os import os
import threading import threading
from datetime import datetime, timedelta from datetime import datetime, timedelta
import pandas as pd import pandas as pd
logger = logging.getLogger(__name__)
BAOSTOCK_DISABLE_ENV = "EASY_TDX_BAOSTOCK" BAOSTOCK_DISABLE_ENV = "EASY_TDX_BAOSTOCK"
# baostock 的全局连接锁(该库单连接、非线程安全) # baostock 的全局连接锁(该库单连接、非线程安全)
@@ -33,6 +41,10 @@ _LOCK_TIMEOUT_SECONDS = 30.0
# 支持兜底的周期(baostock frequency):日线及以上;分钟线/季年线不兜 # 支持兜底的周期(baostock frequency):日线及以上;分钟线/季年线不兜
_FREQ_BY_CATEGORY: dict[str, str] = {"DAY": "d", "WEEK": "w", "MONTH": "m"} _FREQ_BY_CATEGORY: dict[str, str] = {"DAY": "d", "WEEK": "w", "MONTH": "m"}
# 请求字段:baostock 周线/月线不支持 tradestatus(实测 error_code=10004012
# 「周线指标参数传入错误:tradestatus」,2026-09-06),仅日线可传。
_FIELDS_DAILY = "date,open,high,low,close,volume,amount,tradestatus"
_FIELDS_WEEKLY = "date,open,high,low,close,volume,amount"
# 复权映射:baostock adjustflag — 1=后复权 2=前复权 3=不复权 # 复权映射:baostock adjustflag — 1=后复权 2=前复权 3=不复权
_ADJUST_FLAG = {"NONE": "3", "QFQ": "2", "HFQ": "1"} _ADJUST_FLAG = {"NONE": "3", "QFQ": "2", "HFQ": "1"}
_MARKET_PREFIX = {"SZ": "sz", "SH": "sh"} # BJ baostock 不覆盖 _MARKET_PREFIX = {"SZ": "sz", "SH": "sh"} # BJ baostock 不覆盖
@@ -85,6 +97,7 @@ def fetch_bars(
start: int, start: int,
count: int, count: int,
adjust: str, adjust: str,
is_index: bool = False,
) -> pd.DataFrame | None: ) -> pd.DataFrame | None:
"""拉取日线及以上 K 线,输出对齐 /bars 契约的 DataFrame。 """拉取日线及以上 K 线,输出对齐 /bars 契约的 DataFrame。
@@ -95,11 +108,19 @@ def fetch_bars(
start: 跳过最新 start TDX offset 语义一致 start: 跳过最新 start TDX offset 语义一致
count: 最多返回 count count: 最多返回 count
adjust: "NONE" / "QFQ" / "HFQ" adjust: "NONE" / "QFQ" / "HFQ"
is_index: 标的是指数 sh.000001baostock 指数 volume 单位为
/bars/index 输出契约为手通达信指数日线原样/
×100 还原后均为手True vol ÷100 转手实测
sh.000001 2026-09-04volume=53,728,616,100
÷100 = 537,286,161
Returns: Returns:
[date, open, close, high, low, vol, amount] 列序时间升序的 [date, open, close, high, low, vol, amount] 列序时间升序的
DataFrame兜底不可用 / 不适用 / 无数据时返回 None调用方继续 DataFrame兜底不可用 / 不适用 / 无数据时返回 None
维持原错误不吞异常
Raises:
RuntimeError: baostock 登录或查询失败已记 warning 日志auto
兜底调用方以 except 包裹即可维持原错误路径
""" """
global _logged_in global _logged_in
if not is_enabled(): if not is_enabled():
@@ -116,46 +137,47 @@ def fetch_bars(
coef, buffer_days = _WINDOW_DAYS[frequency] coef, buffer_days = _WINDOW_DAYS[frequency]
end_date = datetime.now() end_date = datetime.now()
start_date = end_date - timedelta(days=total * coef + buffer_days) start_date = end_date - timedelta(days=total * coef + buffer_days)
fields = _FIELDS_DAILY if frequency == "d" else _FIELDS_WEEKLY
# baostock 全局单连接:持锁串行;等待超时则放弃本次兜底 # baostock 全局单连接:持锁串行;等待超时则放弃本次兜底
if not _bs_lock.acquire(timeout=_LOCK_TIMEOUT_SECONDS): if not _bs_lock.acquire(timeout=_LOCK_TIMEOUT_SECONDS):
return None return None
try: try:
bs = importlib.import_module("baostock") bs = importlib.import_module("baostock")
try: _login_if_needed(bs)
_login_if_needed(bs) rows = _query_rows(
rows = _query_rows( bs,
bs, code=f"{prefix}.{code}",
code=f"{prefix}.{code}", fields=fields,
fields="date,open,high,low,close,volume,amount,tradestatus", start_date=start_date.strftime("%Y-%m-%d"),
start_date=start_date.strftime("%Y-%m-%d"), end_date=end_date.strftime("%Y-%m-%d"),
end_date=end_date.strftime("%Y-%m-%d"), frequency=frequency,
frequency=frequency, adjustflag=adjustflag,
adjustflag=adjustflag, )
) except Exception as exc:
except Exception: # 连接可能中途断开:重置登录态,下次兜底重新登录。
# 连接可能中途断开:重置登录态,下次兜底重新登录 # 真故障记日志并上抛——auto 兜底调用方(/bars 的 except 分支)接住
_logged_in = False # 后维持原错误;显式 --source baostock 不会被伪装成"无数据"。
raise _logged_in = False
except Exception: logger.warning("baostock 拉取失败(%s.%s %s):%s", prefix, code, frequency, exc)
# 兜底源自身的任何失败都不向上抛:调用方按"无兜底数据"处理 raise RuntimeError(f"baostock 拉取失败: {exc}") from exc
return None
finally: finally:
_bs_lock.release() _bs_lock.release()
if not rows: if not rows:
return None return None
df = pd.DataFrame( df = pd.DataFrame(rows, columns=fields.split(",")).rename(columns={"volume": "vol"})
rows, columns=["date", "open", "high", "low", "close", "vol", "amount", "tradestatus"]
)
for col in ("open", "high", "low", "close", "vol", "amount"): for col in ("open", "high", "low", "close", "vol", "amount"):
df[col] = pd.to_numeric(df[col], errors="coerce") df[col] = pd.to_numeric(df[col], errors="coerce")
# 停牌日剔除(tradestatus=0 或无成交),对齐通达信 K 线不含停牌日的口径 # 停牌日剔除(tradestatus=0 或无成交),对齐通达信 K 线不含停牌日的口径
# W/M 无 tradestatus 列(baostock 不支持),停牌周/月靠 vol>0 兜底剔除。
if "tradestatus" in df.columns: if "tradestatus" in df.columns:
df = df[df["tradestatus"] != "0"] df = df[df["tradestatus"] != "0"]
df = df.dropna(subset=["close"]) df = df.dropna(subset=["close"])
df = df[df["close"] > 0] df = df[df["close"] > 0]
df = df[df["vol"] > 0] df = df[df["vol"] > 0]
if is_index:
df["vol"] = df["vol"] / 100.0 # 股 → 手,见 docstring is_index 说明
if df.empty: if df.empty:
return None return None
df["date"] = pd.to_datetime(df["date"]).dt.normalize() df["date"] = pd.to_datetime(df["date"]).dt.normalize()
+60 -12
View File
@@ -29,6 +29,8 @@ from typing import Any
import pandas as pd import pandas as pd
from easy_tdx.realtime.session import SHANGHAI_TZ
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
__all__ = ["KlineWarehouse", "default_warehouse_path", "MARKET_TO_TDX"] __all__ = ["KlineWarehouse", "default_warehouse_path", "MARKET_TO_TDX"]
@@ -36,6 +38,17 @@ __all__ = ["KlineWarehouse", "default_warehouse_path", "MARKET_TO_TDX"]
# 未收盘 cutoff15:05A股 15:00 收盘 + 5 分钟数据落定余量) # 未收盘 cutoff15:05A股 15:00 收盘 + 5 分钟数据落定余量)
_MARKET_CLOSE_CUTOFF = dt_time(15, 5) _MARKET_CLOSE_CUTOFF = dt_time(15, 5)
def _shanghai_now() -> datetime:
"""沪市墙钟(naive):provisional 判定与主机时区无关。
UTC+8 主机海外服务器的系统本地时间会把当日 / 15:05 判错
UTC 主机上沪市 18:00 收盘后本地才 10:00当日 bar 被误标 provisional
而被默认查询隐藏故统一按 A 股时区取墙钟中国无夏令时固定 UTC+8
"""
return datetime.now(SHANGHAI_TZ).replace(tzinfo=None)
MARKET_TO_TDX: dict[str, int] = {"SZ": 0, "SH": 1, "BJ": 2} MARKET_TO_TDX: dict[str, int] = {"SZ": 0, "SH": 1, "BJ": 2}
_TDX_TO_MARKET: dict[int, str] = {v: k for k, v in MARKET_TO_TDX.items()} _TDX_TO_MARKET: dict[int, str] = {v: k for k, v in MARKET_TO_TDX.items()}
@@ -90,7 +103,13 @@ class KlineWarehouse:
self._duckdb = _require_duckdb() self._duckdb = _require_duckdb()
self._path = Path(db_path) if db_path is not None else default_warehouse_path() self._path = Path(db_path) if db_path is not None else default_warehouse_path()
self._path.parent.mkdir(parents=True, exist_ok=True) self._path.parent.mkdir(parents=True, exist_ok=True)
self._conn = self._duckdb.connect(str(self._path)) try:
self._conn = self._duckdb.connect(str(self._path))
except Exception as exc:
raise RuntimeError(
f"无法打开 K 线仓库 {self._path}{exc}"
"(DuckDB 为单写者——请检查是否另有 easy-tdx 进程/CLI 正占用该文件)"
) from exc
self._conn.execute(_SCHEMA) self._conn.execute(_SCHEMA)
# ── 基本属性 ───────────────────────────────────────────────────────────── # ── 基本属性 ─────────────────────────────────────────────────────────────
@@ -143,7 +162,7 @@ class KlineWarehouse:
if c not in src.columns: if c not in src.columns:
src[c] = float("nan") src[c] = float("nan")
now = datetime.now() now = _shanghai_now()
today = now.date() today = now.date()
before_close = now.time() < _MARKET_CLOSE_CUTOFF before_close = now.time() < _MARKET_CLOSE_CUTOFF
@@ -201,15 +220,44 @@ class KlineWarehouse:
updated = len(rows) - inserted updated = len(rows) - inserted
return (inserted, updated) return (inserted, updated)
def promote_provisional(self) -> int: def promote_provisional(
"""把「日期已过」的 provisional 行转正(收盘后的临时值已被次日增量覆盖)。""" self,
today = datetime.now().date() market: str | None = None,
code: str | None = None,
before: datetime | pd.Timestamp | None = None,
) -> int:
"""把 provisional 行转正为 completed,返回转正行数。
两种用法
- **同步流程推荐**``promote_provisional(market, code,
before=max_dt)``在拉取成功并 upsert 之后调用只转正本次成功
拉到的数据已覆盖datetime <= before的行拉取失败/为空时不
调用盘中临时值不会被洗成 completed
- **无参维护**仅转正日期早于沪市今日的行历史遗留清理
Args:
market: 限定市场None = 全仓库
code: 限定标的None = 全市场
before: 只转正 datetime <= 该时刻的行None = 日期早于沪市今日
"""
conds = ["status = 'provisional'"]
params: list[Any] = []
if before is not None:
conds.append("datetime <= ?")
params.append(pd.Timestamp(before).to_pydatetime())
else:
conds.append("CAST(datetime AS DATE) < ?")
params.append(_shanghai_now().date())
if market is not None:
conds.append("market = ?")
params.append(market.upper())
if code is not None:
conds.append("code = ?")
params.append(code)
cur = self._conn.execute( cur = self._conn.execute(
""" f"UPDATE klines SET status = 'completed' WHERE {' AND '.join(conds)}",
UPDATE klines SET status = 'completed' params,
WHERE status = 'provisional' AND CAST(datetime AS DATE) < ?
""",
[today],
) )
return int(cur.fetchone()[0]) if cur.description else 0 return int(cur.fetchone()[0]) if cur.description else 0
@@ -347,7 +395,7 @@ class KlineWarehouse:
).df() ).df()
issues: list[dict[str, Any]] = [] issues: list[dict[str, Any]] = []
today = date.today() today = _shanghai_now().date()
stale: list[dict[str, Any]] = [] stale: list[dict[str, Any]] = []
total_provisional = 0 total_provisional = 0
@@ -409,7 +457,7 @@ class KlineWarehouse:
"symbols_with_issues": len({i["symbol"] for i in issues}), "symbols_with_issues": len({i["symbol"] for i in issues}),
"stale_symbols": stale[:20], "stale_symbols": stale[:20],
"provisional_rows": total_provisional, "provisional_rows": total_provisional,
"checked_at": datetime.now().isoformat(timespec="seconds"), "checked_at": _shanghai_now().isoformat(timespec="seconds"),
}, },
} }
+77 -11
View File
@@ -4,8 +4,10 @@
- **首同步全量**仓库无该标的数据时按 ``max_bars``默认 8000 拉取 - **首同步全量**仓库无该标的数据时按 ``max_bars``默认 8000 拉取
- **增量补缺**已有数据时只拉最近 ``tail_bars``默认 15 覆盖 - **增量补缺**已有数据时只拉最近 ``tail_bars``默认 15 覆盖
覆盖同日 bar收盘价修正 / provisional 转正不动更早历史 覆盖同日 bar收盘价修正 / provisional 转正不动更早历史尾部窗口
- 同步前自动 :meth:`promote_provisional`过期临时行转正 覆盖不到上次同步点超过 tail_bars 个交易日未同步时自动改全量重拉
- provisional 转正在每标的**拉取成功后**进行且只转正本次数据已覆盖的行
拉取失败/为空的标的盘中临时值保持 provisional不会被洗成 completed
客户端只需具备 ``get_stock_kline(market:int, code, period, start, count, 客户端只需具备 ``get_stock_kline(market:int, code, period, start, count,
adjust)`` 签名``MacClient`` / ``AsyncMacClient`` 均可本同步器只用同步 adjust)`` 签名``MacClient`` / ``AsyncMacClient`` 均可本同步器只用同步
@@ -18,6 +20,8 @@ import logging
from collections.abc import Callable from collections.abc import Callable
from typing import Any from typing import Any
import pandas as pd
from easy_tdx.warehouse.store import MARKET_TO_TDX, KlineWarehouse from easy_tdx.warehouse.store import MARKET_TO_TDX, KlineWarehouse
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -34,6 +38,22 @@ def _period_name(period: str) -> str:
return period.upper() return period.upper()
def _df_time_column(df: pd.DataFrame) -> str | None:
"""取 df 的时间列名(兼容 ``datetime`` / ``date`` 两种客户端输出)。"""
for col in ("datetime", "date"):
if col in df.columns:
return col
return None
def _df_max_datetime(df: pd.DataFrame) -> pd.Timestamp | None:
"""df 内最大 bar 时间(无时间列返回 None)。"""
col = _df_time_column(df)
if col is None:
return None
return pd.to_datetime(df[col]).max()
class WarehouseSyncer: class WarehouseSyncer:
"""把客户端行情增量同步进仓库。 """把客户端行情增量同步进仓库。
@@ -69,6 +89,53 @@ class WarehouseSyncer:
self._tail_bars = max(int(tail_bars), 5) self._tail_bars = max(int(tail_bars), 5)
self._adjust = adjust self._adjust = adjust
def _fetch(
self,
market: str,
code: str,
period: str,
count: int,
) -> pd.DataFrame:
return self._client.get_stock_kline(
MARKET_TO_TDX[market.upper()],
code,
period=period,
start=0,
count=count,
adjust=self._adjust,
)
def _refetch_full_on_gap(
self,
df: pd.DataFrame,
existing_last: pd.Timestamp | None,
market: str,
code: str,
period: str,
symbol: str,
) -> pd.DataFrame:
"""增量尾部覆盖不到上次同步点时全量重拉,消除静默缺口。
尾部只拉 ``tail_bars`` 若超过 tail_bars 个交易日未同步窗口
bar 会晚于仓库最新 bar中间日期不补就永久缺失检测到该形态
bar 时间 > existing_last即改全量重拉一次
"""
if existing_last is None or df is None or len(df) == 0:
return df
col = _df_time_column(df)
if col is None:
return df
first_dt = pd.to_datetime(df[col]).min()
if first_dt <= existing_last:
return df
logger.warning(
"仓库同步 %s:增量尾部(首根 %s)覆盖不到上次同步点 %s,存在缺口,改全量重拉",
symbol,
first_dt,
existing_last,
)
return self._fetch(market, code, period, self._max_bars)
def sync_symbol( def sync_symbol(
self, self,
market: str, market: str,
@@ -80,17 +147,17 @@ class WarehouseSyncer:
try: try:
existing_last = self._wh.last_datetime(market, code, period) existing_last = self._wh.last_datetime(market, code, period)
count = self._tail_bars if existing_last is not None else self._max_bars count = self._tail_bars if existing_last is not None else self._max_bars
df = self._client.get_stock_kline( df = self._fetch(market, code, period, count)
MARKET_TO_TDX[market.upper()], df = self._refetch_full_on_gap(df, existing_last, market, code, period, symbol)
code,
period=period,
start=0,
count=count,
adjust=self._adjust,
)
if df is None or len(df) == 0: if df is None or len(df) == 0:
return {"symbol": symbol, "added": 0, "updated": 0, "skipped": 1, "error": None} return {"symbol": symbol, "added": 0, "updated": 0, "skipped": 1, "error": None}
added, updated = self._wh.upsert_bars(market, code, df, period=period) added, updated = self._wh.upsert_bars(market, code, df, period=period)
# provisional 转正:只在拉取成功后进行,且只转正本次数据已覆盖
# datetime <= 拉取最大时间)的行——拉取失败/为空时盘中临时值
# 不会被洗成 completed。
max_dt = _df_max_datetime(df)
if max_dt is not None:
self._wh.promote_provisional(market=market, code=code, before=max_dt)
return { return {
"symbol": symbol, "symbol": symbol,
"added": added, "added": added,
@@ -118,7 +185,6 @@ class WarehouseSyncer:
Returns: Returns:
``{"total", "ok", "added", "updated", "skipped", "failed", "details"}`` ``{"total", "ok", "added", "updated", "skipped", "failed", "details"}``
""" """
self._wh.promote_provisional()
p = _period_name(period) p = _period_name(period)
parsed: list[tuple[str, str]] = [] parsed: list[tuple[str, str]] = []
+12 -2
View File
@@ -134,8 +134,14 @@ class OptimizeBacktestRequest(BaseModel):
strategy: str = Field(..., description="策略名") strategy: str = Field(..., description="策略名")
cash: float = Field(default=1_000_000.0, gt=0) cash: float = Field(default=1_000_000.0, gt=0)
commission: float = Field(default=0.0003, ge=0, le=0.01) commission: float = Field(default=0.0003, ge=0, le=0.01)
min_commission: float = Field(default=5.0, ge=0)
stamp_tax: float = Field(default=0.001, ge=0, le=0.01)
slippage: float = Field(default=0.0, ge=0, le=0.05) slippage: float = Field(default=0.0, ge=0, le=0.05)
execution: Literal["next_open", "next_close"] = Field(default="next_open") execution: Literal["next_open", "next_close"] = Field(default="next_open")
auto_fees: bool = Field(
default=False,
description="按标的品种自动解析费率(ETF/可转债免印花税等);显式非默认费率优先",
)
workers: int = Field( workers: int = Field(
default=1, default=1,
ge=0, ge=0,
@@ -304,10 +310,14 @@ class BacktestResultResponse(BaseModel):
class TaskSubmitResponse(BaseModel): class TaskSubmitResponse(BaseModel):
"""后台任务提交响应。""" """后台任务提交响应。
``status`` 透传提交瞬间的真实状态通常是 pending/running极快任务在
拿到响应前可能已 done/failed如实上报前端轮询一次即见分晓
"""
task_id: str task_id: str
status: Literal["pending", "running"] status: Literal["pending", "running", "done", "failed"]
class TaskStateResponse(BaseModel): class TaskStateResponse(BaseModel):
+96 -58
View File
@@ -40,6 +40,10 @@ from easy_tdx.web.task_runner import get_runner
router = APIRouter(tags=["backtest"]) router = APIRouter(tags=["backtest"])
# 标准 TdxClient 单次 get_security_bars 取数上限(协议约束,服务器对更大
# 请求静默截断)。所有按标的取数路径都必须经 _fetch_bars_paged 翻页。
_BARS_PAGE_SIZE = 800
# ── 策略枚举 ─────────────────────────────────────────────────────────────────── # ── 策略枚举 ───────────────────────────────────────────────────────────────────
@@ -104,10 +108,9 @@ async def run_backtest_async(
# 3. 提交后台任务 # 3. 提交后台任务
runner = get_runner() runner = get_runner()
task_id = runner.submit(lambda: _run_backtest(df, snapshot), description=description) task_id = runner.submit(lambda: _run_backtest(df, snapshot), description=description)
# 提交瞬间通常是 pending/running;极快任务可能已 done/failed,如实上报
state = runner.get(task_id) state = runner.get(task_id)
# 提交瞬间任务应是 pending/running;极端情况下线程已跑完则报实际状态 return TaskSubmitResponse(task_id=task_id, status=state.status)
status: Any = state.status if state.status in ("pending", "running") else "running"
return TaskSubmitResponse(task_id=task_id, status=status)
@router.get("/backtest/tasks", response_model=TaskListResponse) @router.get("/backtest/tasks", response_model=TaskListResponse)
@@ -252,8 +255,7 @@ async def run_portfolio_backtest_async(
description=description, description=description,
) )
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
# ── 多策略组合回测(资金分仓) ─────────────────────────────────────────────── # ── 多策略组合回测(资金分仓) ───────────────────────────────────────────────
@@ -285,8 +287,7 @@ async def run_multi_strategy_backtest_async(
description=description, description=description,
) )
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
@router.post( @router.post(
@@ -316,8 +317,7 @@ async def run_multi_strategy_walkforward_async(
description=description, description=description,
) )
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
@router.post( @router.post(
@@ -348,8 +348,7 @@ async def run_multi_strategy_evaluate_async(
description=description, description=description,
) )
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
@router.post("/backtest/optimize/run/async", response_model=TaskSubmitResponse, status_code=202) @router.post("/backtest/optimize/run/async", response_model=TaskSubmitResponse, status_code=202)
@@ -388,8 +387,7 @@ async def run_optimize_async(
description=description, description=description,
) )
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
# ── 一键寻优所有策略 ─────────────────────────────────────────────────────────── # ── 一键寻优所有策略 ───────────────────────────────────────────────────────────
@@ -429,8 +427,7 @@ async def run_optimize_all_async(
description=description, description=description,
) )
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
# ── 信号雷达(一键扫描已保存策略)──────────────────────────────────────────── # ── 信号雷达(一键扫描已保存策略)────────────────────────────────────────────
@@ -467,8 +464,7 @@ async def run_signal_scan_async(
description=description, description=description,
) )
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
# ── Walk-Forward / 一条龙评估(v1.25 防过拟合链)────────────────────────────── # ── Walk-Forward / 一条龙评估(v1.25 防过拟合链)──────────────────────────────
@@ -494,8 +490,7 @@ async def run_walkforward_async(
lambda: _run_walkforward(df, snapshot, n_windows), description=description lambda: _run_walkforward(df, snapshot, n_windows), description=description
) )
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
@router.post("/backtest/evaluate/run/async", response_model=TaskSubmitResponse, status_code=202) @router.post("/backtest/evaluate/run/async", response_model=TaskSubmitResponse, status_code=202)
@@ -516,8 +511,7 @@ async def run_evaluate_async(
runner = get_runner() runner = get_runner()
task_id = runner.submit(lambda: _run_evaluate(df, snapshot), description=description) task_id = runner.submit(lambda: _run_evaluate(df, snapshot), description=description)
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
# ── 组合级 Walk-Forward / 一条龙评估(对齐单标的防过拟合链)────────────────── # ── 组合级 Walk-Forward / 一条龙评估(对齐单标的防过拟合链)──────────────────
@@ -550,8 +544,7 @@ async def run_portfolio_walkforward_async(
description=description, description=description,
) )
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
@router.post( @router.post(
@@ -581,8 +574,7 @@ async def run_portfolio_evaluate_async(
description=description, description=description,
) )
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
async def _resolve_df(client: Any, req: BacktestRequest) -> pd.DataFrame: async def _resolve_df(client: Any, req: BacktestRequest) -> pd.DataFrame:
@@ -606,19 +598,10 @@ async def run_multiseed_async(
平均收益结果含 per_seed_positive_ratio 稳定性列通过 平均收益结果含 per_seed_positive_ratio 稳定性列通过
GET /backtest/tasks/{task_id} 轮询 GET /backtest/tasks/{task_id} 轮询
""" """
from easy_tdx.web.convert import category_from_str, market_from_str
stock_dfs: dict[str, pd.DataFrame] = {} stock_dfs: dict[str, pd.DataFrame] = {}
for symbol in req.stocks: for symbol in req.stocks:
market_str, code = symbol.split(":", 1)
try: try:
page = await client.get_security_bars( page = await _fetch_bars_paged(client, symbol, req.category, req.count)
market_from_str(market_str),
code,
category_from_str(req.category),
0,
req.count,
)
except Exception: # noqa: BLE001 — 单标的失败跳过 except Exception: # noqa: BLE001 — 单标的失败跳过
continue continue
if len(page) >= 30: if len(page) >= 30:
@@ -631,8 +614,7 @@ async def run_multiseed_async(
runner = get_runner() runner = get_runner()
task_id = runner.submit(lambda: _run_multiseed(stock_dfs, snapshot), description=description) task_id = runner.submit(lambda: _run_multiseed(stock_dfs, snapshot), description=description)
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
def _run_multiseed(stock_dfs: dict[str, pd.DataFrame], req: MultiSeedRequest) -> dict[str, Any]: def _run_multiseed(stock_dfs: dict[str, pd.DataFrame], req: MultiSeedRequest) -> dict[str, Any]:
@@ -781,15 +763,10 @@ async def run_rotation_async(
可选槽内止盈止损打分支持内置动量``score="momentum"`` + ``period`` 可选槽内止盈止损打分支持内置动量``score="momentum"`` + ``period``
或通达信公式数值输出``score="formula"`` + ``formula_text`` + ``score_col`` 或通达信公式数值输出``score="formula"`` + ``formula_text`` + ``score_col``
""" """
from easy_tdx.web.convert import category_from_str, market_from_str
stock_dfs: dict[str, pd.DataFrame] = {} stock_dfs: dict[str, pd.DataFrame] = {}
for symbol in req.stocks: for symbol in req.stocks:
market_str, code = symbol.split(":", 1)
try: try:
page = await client.get_security_bars( page = await _fetch_bars_paged(client, symbol, req.category, req.count)
market_from_str(market_str), code, category_from_str(req.category), 0, req.count
)
except Exception: # noqa: BLE001 — 单标的失败跳过 except Exception: # noqa: BLE001 — 单标的失败跳过
continue continue
if page is not None and len(page) >= 30: if page is not None and len(page) >= 30:
@@ -802,8 +779,7 @@ async def run_rotation_async(
runner = get_runner() runner = get_runner()
task_id = runner.submit(lambda: _run_rotation(stock_dfs, snapshot), description=description) task_id = runner.submit(lambda: _run_rotation(stock_dfs, snapshot), description=description)
state = runner.get(task_id) state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running" return TaskSubmitResponse(task_id=task_id, status=state.status)
return TaskSubmitResponse(task_id=task_id, status=status)
def _run_rotation( def _run_rotation(
@@ -916,18 +892,44 @@ def _normalize_bars_dt(df: pd.DataFrame) -> pd.DataFrame:
return df return df
async def _fetch_bars(client: Any, symbol: str, category: str, count: int) -> pd.DataFrame: async def _fetch_bars_paged(client: Any, symbol: str, category: str, count: int) -> pd.DataFrame:
"""标的取 K 线(async,必须在 event loop 内调用)。""" """ 800/页翻页取最多 ``count`` 根 K 线,返回时间升序 DataFrame。
TDX 协议单次 get_security_bars 最多返回 800 count>800 的单次调用会被
服务器静默截断multiseed / rotation / formula 曾各自单页取数悄悄少
数据本辅助按 start=0,800,1600 翻页拼接页间按时间升序排序
末页不足 800 根视为数据起点提前停止列结构与 get_security_bars
原始输出一致日线 ``date`` / 分钟 ``datetime``不做改名/类型规整
"""
from easy_tdx.web.convert import category_from_str, market_from_str from easy_tdx.web.convert import category_from_str, market_from_str
market_str, code = symbol.split(":", 1) market_str, code = symbol.split(":", 1)
df = await client.get_security_bars( market = market_from_str(market_str)
market_from_str(market_str), cat = category_from_str(category)
code, frames: list[pd.DataFrame] = []
category_from_str(category), fetched = 0
0, while fetched < count:
count, page_size = min(_BARS_PAGE_SIZE, count - fetched)
) page_df = await client.get_security_bars(market, code, cat, fetched, page_size)
if page_df is None or len(page_df) == 0:
break
frames.append(page_df)
fetched += len(page_df)
if len(page_df) < page_size:
break # 数据起点
if not frames:
return pd.DataFrame()
df = pd.concat(frames, ignore_index=True)
dt_col = "datetime" if "datetime" in df.columns else "date"
if dt_col in df.columns:
# 页间天然逆序(page0=最新一页),拼接后按时间升序
df = df.sort_values(dt_col).reset_index(drop=True)
return df
async def _fetch_bars(client: Any, symbol: str, category: str, count: int) -> pd.DataFrame:
"""按标的取 K 线(async,必须在 event loop 内调用)。"""
df = await _fetch_bars_paged(client, symbol, category, count)
if len(df) == 0: if len(df) == 0:
raise ValueError(f"标的 {symbol} 未取到任何 K 线数据") raise ValueError(f"标的 {symbol} 未取到任何 K 线数据")
return _normalize_bars_dt(df) return _normalize_bars_dt(df)
@@ -1182,31 +1184,67 @@ def _run_multi_strategy_evaluate(
) )
def _resolve_effective_fees(
auto_fees: bool,
symbol: str | None,
commission: float,
min_commission: float,
stamp_tax: float,
) -> tuple[float, float, float]:
"""auto_fees 品种费率解析(与 BacktestEngine 同款口径)。
显式非默认值优先调用方有意覆盖默认值按品种费率表替换
ETF/可转债免印花税ParamGridOptimizer auto_fees 参数寻优端点
web 层预解析成具体费率再传入保证与单标的回测同口径
"""
if not auto_fees or not symbol:
return commission, min_commission, stamp_tax
from easy_tdx.backtest.fees import resolve_fee_model
fee = resolve_fee_model(symbol)
if commission == 0.0003:
commission = fee.commission
if min_commission == 5.0:
min_commission = fee.min_commission
if stamp_tax == 0.001:
stamp_tax = fee.stamp_tax
return commission, min_commission, stamp_tax
def _run_optimize(df: pd.DataFrame, req: OptimizeBacktestRequest) -> dict[str, Any]: def _run_optimize(df: pd.DataFrame, req: OptimizeBacktestRequest) -> dict[str, Any]:
"""执行参数网格寻优并返回清洗后的结果字典(后台线程内调用)。""" """执行参数网格寻优并返回清洗后的结果字典(后台线程内调用)。"""
from easy_tdx.backtest.benchmark import run_buy_hold_benchmark from easy_tdx.backtest.benchmark import run_buy_hold_benchmark
from easy_tdx.backtest.optimizer import ParamGridOptimizer from easy_tdx.backtest.optimizer import ParamGridOptimizer
commission, min_commission, stamp_tax = _resolve_effective_fees(
req.auto_fees, req.symbol, req.commission, req.min_commission, req.stamp_tax
)
optimizer = ParamGridOptimizer( optimizer = ParamGridOptimizer(
strategy_name=req.strategy, strategy_name=req.strategy,
param_grid=req.param_grid, param_grid=req.param_grid,
df=df, df=df,
cash=req.cash, cash=req.cash,
commission=req.commission, commission=commission,
min_commission=min_commission,
stamp_tax=stamp_tax,
slippage=req.slippage, slippage=req.slippage,
execution=req.execution, execution=req.execution,
workers=req.workers, workers=req.workers,
) )
result = optimizer.run() result = optimizer.run()
out = result.to_dict() out = result.to_dict()
# 买入持有基准(同区间/同费率/同资金,与一条龙评估同口径), # 买入持有基准(同区间/同费率/同资金,与一条龙评估同口径),供前端在
# 供前端在最优结果旁直观对比「策略 vs 买入不动」。 # 最优结果旁直观对比「策略 vs 买入不动」。
out["buy_hold"] = run_buy_hold_benchmark( out["buy_hold"] = run_buy_hold_benchmark(
df, df,
cash=req.cash, cash=req.cash,
commission=req.commission, commission=commission,
min_commission=min_commission,
stamp_tax=stamp_tax,
slippage=req.slippage, slippage=req.slippage,
execution=req.execution, execution=req.execution,
symbol=req.symbol,
auto_fees=req.auto_fees,
) )
return out return out
+36 -7
View File
@@ -218,13 +218,24 @@ async def _fetch_120m(
async def _baostock_last_resort( async def _baostock_last_resort(
market: str, code: str, category: str, start: int, count: int, adjust: str market: str,
code: str,
category: str,
start: int,
count: int,
adjust: str,
is_index: bool = False,
) -> tuple[pd.DataFrame | None, str | None]: ) -> tuple[pd.DataFrame | None, str | None]:
"""TDX 全部路径失败/为空后的最后一级兜底:baostock(仅日线及以上)。 """TDX 全部路径失败/为空后的最后一级兜底:baostock(仅日线及以上)。
未安装 baostock / 设置了 EASY_TDX_BAOSTOCK=0 / 周期不适用 / 查询失败 未安装 baostock / 设置了 EASY_TDX_BAOSTOCK=0 / 周期不适用 / 查询失败
一律返回 ``(None, None)``兜底源自身的任何失败都不影响原错误语义 一律返回 ``(None, None)``兜底源自身的任何失败含新版 fetch_bars
baostock 客户端阻塞且非线程安全丢线程池执行模块内部持锁串行 真故障抛出的 RuntimeError都按"兜底不可用"处理调用方继续维持原
TDX 错误语义baostock 客户端阻塞且非线程安全丢线程池执行模块内部
持锁串行
Args:
is_index: 标的是指数/bars/index 兜底传 Truevol 手对齐契约
""" """
from easy_tdx.sources import baostock as baostock_source from easy_tdx.sources import baostock as baostock_source
@@ -232,14 +243,28 @@ async def _baostock_last_resort(
return None, None return None, None
try: try:
df = await asyncio.to_thread( df = await asyncio.to_thread(
baostock_source.fetch_bars, market, code, category, start, count, adjust baostock_source.fetch_bars,
market,
code,
category,
start,
count,
adjust,
is_index,
) )
except Exception as exc: # noqa: BLE001 — 兜底失败不改变原错误路径 except Exception as exc: # noqa: BLE001 — 兜底失败不改变原错误路径
_logger.warning("/bars baostock 兜底异常 (%s%s): %s", market, code, exc) _logger.warning("/bars baostock 兜底异常 (%s%s): %s", market, code, exc)
return None, None return None, None
if df is None or df.empty: if df is None or df.empty:
return None, None return None, None
_logger.info("/bars 已启用 baostock 兜底 (%s%s %s%d 根)", market, code, category, len(df)) _logger.info(
"/bars 已启用 baostock 兜底 (%s%s %s%s%d 根)",
market,
code,
category,
",指数" if is_index else "",
len(df),
)
return df, "baostock" return df, "baostock"
@@ -339,7 +364,9 @@ async def security_bars(
_logger.warning("/bars 标准 TdxClient 获取失败 (%s%s): %s", market, code, exc) _logger.warning("/bars 标准 TdxClient 获取失败 (%s%s): %s", market, code, exc)
if df is None or df.empty: if df is None or df.empty:
bdf, bsource = await _baostock_last_resort(market, code, category, start, count, adjust) # 周期先归一成枚举名("4"→DAY):baostock 频率查表只认名称,
# 数字串直接透传会让兜底静默失效
bdf, bsource = await _baostock_last_resort(market, code, cat.name, start, count, adjust)
if bdf is not None: if bdf is not None:
df, source = bdf, bsource df, source = bdf, bsource
@@ -391,7 +418,9 @@ async def index_bars(
_logger.warning("/bars/index TdxClient 获取失败 (%s%s): %s", market, code, exc) _logger.warning("/bars/index TdxClient 获取失败 (%s%s): %s", market, code, exc)
if df is None or df.empty: if df is None or df.empty:
bdf, bsource = await _baostock_last_resort(market, code, category, start, count, "QFQ") bdf, bsource = await _baostock_last_resort(
market, code, category_from_str(category).name, start, count, "QFQ", is_index=True
)
if bdf is not None: if bdf is not None:
df, source = bdf, bsource df, source = bdf, bsource
+14 -10
View File
@@ -12,6 +12,7 @@ import pandas as pd
from fastapi import APIRouter, Depends, Query from fastapi import APIRouter, Depends, Query
from easy_tdx.mac.enums import Adjust, Period from easy_tdx.mac.enums import Adjust, Period
from easy_tdx.realtime.session import SHANGHAI_TZ, is_trading_time
from easy_tdx.web.convert import ( from easy_tdx.web.convert import (
board_sort_from_str, board_sort_from_str,
board_type_from_str, board_type_from_str,
@@ -20,7 +21,7 @@ from easy_tdx.web.convert import (
sort_type_from_str, sort_type_from_str,
) )
from easy_tdx.web.deps import get_mac_client from easy_tdx.web.deps import get_mac_client
from easy_tdx.web.schemas import DataFrameResponse, DictResponse from easy_tdx.web.schemas import DataFrameResponse, DictResponse, _json_safe
_logger = logging.getLogger(__name__) _logger = logging.getLogger(__name__)
@@ -37,9 +38,9 @@ _OVERVIEW_METRIC_FIELDS: dict[str, str] = {
"YTD": "chg_ytd", "YTD": "chg_ytd",
} }
_OVERVIEW_TTL = 15.0 _OVERVIEW_TTL = 15.0
# (board_type, metrics) -> (monotonic 截止时间, payload)。无锁:并发重复拉取 # (board_type, metrics, count) -> (monotonic 截止时间, payload)。无锁:并发
# 无害(AsyncMacClient 连接内本就串行),省去跨事件循环的锁生命周期问题。 # 重复拉取无害(AsyncMacClient 连接内本就串行),省去跨事件循环的锁生命周期问题。
_overview_cache: dict[tuple[str, tuple[str, ...]], tuple[float, dict[str, Any]]] = {} _overview_cache: dict[tuple[str, tuple[str, ...], int], tuple[float, dict[str, Any]]] = {}
# 可在单测中 monkeypatch 以控制 TTL 判定 # 可在单测中 monkeypatch 以控制 TTL 判定
_now = time.monotonic _now = time.monotonic
@@ -181,7 +182,7 @@ async def board_overview(
valid = ", ".join(_OVERVIEW_METRIC_FIELDS) valid = ", ".join(_OVERVIEW_METRIC_FIELDS)
raise ValueError(f"无效指标 '{','.join(invalid)}',可选值: {valid}") raise ValueError(f"无效指标 '{','.join(invalid)}',可选值: {valid}")
cache_key = (bt.name, tuple(sort_names)) cache_key = (bt.name, tuple(sort_names), count) # count 影响 payload 行数,必须入键
cached = _overview_cache.get(cache_key) cached = _overview_cache.get(cache_key)
if cached is not None and _now() < cached[0]: if cached is not None and _now() < cached[0]:
return DictResponse.from_dict(cached[1]) return DictResponse.from_dict(cached[1])
@@ -231,7 +232,12 @@ async def board_overview(
row.setdefault(field, None) row.setdefault(field, None)
rows.append(row) rows.append(row)
payload = {"board_type": bt.name, "ts": int(time.time()), "count": len(rows), "rows": rows} # 坏值(NaN/inf,如无足够历史板块的 sort_value)在入缓存前清洗成 null:
# 带 NaN 的 payload 一旦入缓存,15s TTL 内每次命中都会在 JSON 序列化时
# 500Starlette allow_nan=False)。
payload = _json_safe(
{"board_type": bt.name, "ts": int(time.time()), "count": len(rows), "rows": rows}
)
_overview_cache[cache_key] = (_now() + _OVERVIEW_TTL, payload) _overview_cache[cache_key] = (_now() + _OVERVIEW_TTL, payload)
return DictResponse.from_dict(payload) return DictResponse.from_dict(payload)
@@ -263,8 +269,8 @@ _hotspot_builds: dict[str, dict[str, Any]] = {}
def _today_str() -> str: def _today_str() -> str:
"""当日日历日(缓存失效键;单测可 monkeypatch)。""" """当日日历日(沪市时区,与主机时区无关;缓存失效键;单测可 monkeypatch)。"""
return datetime.now().strftime("%Y-%m-%d") return datetime.now(SHANGHAI_TZ).strftime("%Y-%m-%d")
async def _hotspot_build(board_key: str, bt: Any, client: Any) -> None: async def _hotspot_build(board_key: str, bt: Any, client: Any) -> None:
@@ -500,8 +506,6 @@ async def board_hotspot(
) )
rows_out = rows_out[:_HOTSPOT_MAX_ROWS] rows_out = rows_out[:_HOTSPOT_MAX_ROWS]
from easy_tdx.realtime.session import is_trading_time
payload: dict[str, Any] = { payload: dict[str, Any] = {
"status": "ready", "status": "ready",
"board_type": bt.name, "board_type": bt.name,
+14 -7
View File
@@ -131,7 +131,11 @@ async def run_formula_backtest_async(
lambda: _run_formula_backtest(df, snapshot), lambda: _run_formula_backtest(df, snapshot),
description=f"公式回测 | {snapshot.symbol or '内联数据'}", description=f"公式回测 | {snapshot.symbol or '内联数据'}",
) )
return {"task_id": task_id, "status": "running"} try:
status = runner.get(task_id).status
except KeyError: # 极端:状态尚未可查时按提交默认态上报
status = "running"
return {"task_id": task_id, "status": status}
@router.post("/formula/screen/run/async", status_code=202) @router.post("/formula/screen/run/async", status_code=202)
@@ -167,7 +171,11 @@ async def run_formula_screen_async(
lambda: _run_formula_screen(bars, compiled, snapshot.signal_col), lambda: _run_formula_screen(bars, compiled, snapshot.signal_col),
description=f"公式选股 | {len(bars)}只标的", description=f"公式选股 | {len(bars)}只标的",
) )
return {"task_id": task_id, "status": "running"} try:
status = runner.get(task_id).status
except KeyError: # 极端:状态尚未可查时按提交默认态上报
status = "running"
return {"task_id": task_id, "status": status}
# ── 内部实现 ─────────────────────────────────────────────────────────────────── # ── 内部实现 ───────────────────────────────────────────────────────────────────
@@ -187,12 +195,11 @@ async def _resolve_df(client: Any, req: FormulaComputeRequest) -> Any:
df["datetime"] = pd.to_datetime(df["datetime"], errors="coerce") df["datetime"] = pd.to_datetime(df["datetime"], errors="coerce")
return df return df
if req.symbol is not None: if req.symbol is not None:
from easy_tdx.web.convert import category_from_str, market_from_str # 经 _fetch_bars_paged 按 800/页翻页取全量(协议单次上限 800,
# 单次调用 count>800 会被服务器静默截断,指标计算窗口悄悄变短)
from easy_tdx.web.routers.backtest import _fetch_bars_paged
market_str, code = req.symbol.split(":", 1) df = await _fetch_bars_paged(client, req.symbol, req.category, req.count)
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: if df is None or len(df) == 0:
raise ValueError(f"标的 {req.symbol} 未取到 K 线数据") raise ValueError(f"标的 {req.symbol} 未取到 K 线数据")
return df return df
+4 -2
View File
@@ -226,8 +226,10 @@ async def llm_chat_async(req: LlmChatRequest) -> TaskSubmitResponse:
runner = get_runner() runner = get_runner()
task_id = runner.submit(_run, description=desc) task_id = runner.submit(_run, description=desc)
state = runner.get(task_id) try:
status: Any = state.status if state.status in ("pending", "running") else "running" status = runner.get(task_id).status
except KeyError: # 极端:状态尚未可查时按提交默认态上报
status = "running"
return TaskSubmitResponse(task_id=task_id, status=status) return TaskSubmitResponse(task_id=task_id, status=status)
+17 -16
View File
@@ -21,11 +21,13 @@ from easy_tdx.web.schemas import (
router = APIRouter(tags=["market"]) router = APIRouter(tags=["market"])
# 涨停生态结果缓存(vipdoc 盘中随通达信客户端落盘更新,60s 足够新鲜) # 涨停生态结果缓存(vipdoc 盘中随通达信客户端落盘更新,60s 足够新鲜)
_limitup_cache: tuple[float, dict[str, Any]] | None = None # 键 = effective vipdoc(显式参数 > 已存设置 > None):不同数据目录的扫描
# 结果必须互不串台。
_limitup_cache: dict[str | None, tuple[float, dict[str, Any]]] = {}
_LIMITUP_TTL = 60.0 _LIMITUP_TTL = 60.0
# 涨停逐日历史缓存(历史数据不变,10 分钟;按 days 分键 # 涨停逐日历史缓存(历史数据不变,10 分钟;键 = (days, effective vipdoc)
_limitup_history_cache: dict[int, tuple[float, dict[str, Any]]] = {} _limitup_history_cache: dict[tuple[int, str | None], tuple[float, dict[str, Any]]] = {}
def _df_response(df: Any) -> DataFrameResponse: def _df_response(df: Any) -> DataFrameResponse:
@@ -148,8 +150,8 @@ async def set_vipdoc_setting(req: dict[str, Any]) -> dict[str, Any]:
get_app_settings_store().delete(_VIPDOC_KEY) get_app_settings_store().delete(_VIPDOC_KEY)
resolved = None resolved = None
# 路径变更后旧扫描结果作废 # 路径变更后旧扫描结果作废
global _limitup_cache, _limitup_history_cache global _limitup_cache
_limitup_cache = None _limitup_cache.clear()
_limitup_history_cache.clear() _limitup_history_cache.clear()
return {"stored": path, "resolved": resolved} return {"stored": path, "resolved": resolved}
@@ -165,12 +167,11 @@ async def limitup_ecology(
涨停判定按代码段主板 10% 5% 疑似 ST 标记创业板/科创板 20% 涨停判定按代码段主板 10% 5% 疑似 ST 标记创业板/科创板 20%
.day 文件无名称name 由前端经批量报价补齐 .day 文件无名称name 由前端经批量报价补齐
""" """
global _limitup_cache
now = time.monotonic()
if _limitup_cache is not None and now - _limitup_cache[0] < _LIMITUP_TTL:
return DictResponse.from_dict(_limitup_cache[1])
effective = _effective_vipdoc(vipdoc) effective = _effective_vipdoc(vipdoc)
now = time.monotonic()
cached = _limitup_cache.get(effective)
if cached is not None and now - cached[0] < _LIMITUP_TTL:
return DictResponse.from_dict(cached[1])
def _scan() -> dict[str, Any]: def _scan() -> dict[str, Any]:
from easy_tdx.screen.limitup import compute_limitup_ecology from easy_tdx.screen.limitup import compute_limitup_ecology
@@ -194,7 +195,7 @@ async def limitup_ecology(
} }
payload = await asyncio.to_thread(_scan) payload = await asyncio.to_thread(_scan)
_limitup_cache = (now, payload) _limitup_cache[effective] = (now, payload)
return DictResponse.from_dict(payload) return DictResponse.from_dict(payload)
@@ -244,20 +245,20 @@ async def limitup_history(
全市场扫描约需数十秒结果缓存 10 分钟日期覆盖受 vipdoc 数据范围限制 全市场扫描约需数十秒结果缓存 10 分钟日期覆盖受 vipdoc 数据范围限制
""" """
global _limitup_history_cache effective = _effective_vipdoc(vipdoc)
now = time.monotonic() now = time.monotonic()
cached = _limitup_history_cache.get(days) cached = _limitup_history_cache.get((days, effective))
if cached is not None and now - cached[0] < 600: if cached is not None and now - cached[0] < 600:
return DictResponse.from_dict(cached[1]) return DictResponse.from_dict(cached[1])
def _scan() -> dict[str, Any]: def _scan() -> dict[str, Any]:
from easy_tdx.screen.limitup import compute_limitup_history from easy_tdx.screen.limitup import compute_limitup_history
rows = compute_limitup_history(_effective_vipdoc(vipdoc), days=days) rows = compute_limitup_history(effective, days=days)
return {"count": len(rows), "days": rows} return {"count": len(rows), "days": rows}
payload = await asyncio.to_thread(_scan) payload = await asyncio.to_thread(_scan)
_limitup_history_cache[days] = (now, payload) _limitup_history_cache[(days, effective)] = (now, payload)
return DictResponse.from_dict(payload) return DictResponse.from_dict(payload)
+15 -4
View File
@@ -8,15 +8,20 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
from typing import Annotated
from fastapi import APIRouter, Request from fastapi import APIRouter, Request
from pydantic import BaseModel from pydantic import BaseModel, Field, StringConstraints
from easy_tdx.config import get_best_host, get_known_hosts, get_port, save_best_host from easy_tdx.config import get_best_host, get_known_hosts, get_port, save_best_host
from easy_tdx.transport.sync import ping_all from easy_tdx.transport.sync import ping_all
router = APIRouter(tags=["server"]) router = APIRouter(tags=["server"])
# 单项 host 长度上限(域名合法最大 253 字符);列表项数上限防被当作
# 无限制的内网扫描跳板。
HostStr = Annotated[str, StringConstraints(max_length=253)]
# --------------------------------------------------------------------------- # # --------------------------------------------------------------------------- #
# Schemas # Schemas
@@ -41,10 +46,16 @@ class HostListResponse(BaseModel):
class ServerTestRequest(BaseModel): class ServerTestRequest(BaseModel):
"""POST /server/test 的请求。""" """POST /server/test 的请求
hosts: list[str] | None = None # None = 测全部候选 ``timeout`` 0.5~30sto_thread 内的同步 ping 无中断手段无上界的
timeout: float = 5.0 超时会长期占住线程池线程``hosts`` 50 单项 253 字符
"""
hosts: list[HostStr] | None = Field(
default=None, max_length=50, description="待测主机列表;None = 测全部候选"
)
timeout: float = Field(default=5.0, ge=0.5, le=30.0, description="单主机连接超时(秒)")
class ServerSwitchRequest(BaseModel): class ServerSwitchRequest(BaseModel):
+9 -2
View File
@@ -3,18 +3,22 @@
from __future__ import annotations from __future__ import annotations
from fastapi import APIRouter, HTTPException, Query from fastapi import APIRouter, HTTPException, Query
from fastapi import Path as PathParam
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from easy_tdx.web.watchlist_store import get_watchlist_store from easy_tdx.web.watchlist_store import get_watchlist_store
router = APIRouter(tags=["watchlist"]) router = APIRouter(tags=["watchlist"])
# 6 位数字代码(自选会被 QuoteStreamer 拿去轮询,非数字代码产生无效请求)
_CODE_PATTERN = r"^\d{6}$"
class WatchItemAdd(BaseModel): class WatchItemAdd(BaseModel):
"""加入自选请求。name 由前端从行情数据带过来。""" """加入自选请求。name 由前端从行情数据带过来。"""
market: str = Field(..., pattern=r"^(SZ|SH|BJ)$") market: str = Field(..., pattern=r"^(SZ|SH|BJ)$")
code: str = Field(..., min_length=6, max_length=6) code: str = Field(..., pattern=_CODE_PATTERN)
name: str = Field("", max_length=64) name: str = Field("", max_length=64)
group: str = Field("默认", max_length=32) group: str = Field("默认", max_length=32)
@@ -41,7 +45,10 @@ async def add_watch_item(req: WatchItemAdd) -> dict[str, object]:
@router.delete("/watchlist/{market}/{code}", response_model=dict[str, object]) @router.delete("/watchlist/{market}/{code}", response_model=dict[str, object])
async def remove_watch_item(market: str, code: str) -> dict[str, object]: async def remove_watch_item(
market: str,
code: str = PathParam(..., pattern=_CODE_PATTERN, description="6位数字代码"),
) -> dict[str, object]:
"""移除自选。""" """移除自选。"""
if market.upper() not in {"SZ", "SH", "BJ"}: if market.upper() not in {"SZ", "SH", "BJ"}:
raise HTTPException(status_code=400, detail=f"非法市场: {market}") raise HTTPException(status_code=400, detail=f"非法市场: {market}")
+29 -6
View File
@@ -2,11 +2,38 @@
from __future__ import annotations from __future__ import annotations
import math
from enum import IntEnum from enum import IntEnum
from typing import Any from typing import Any
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
def _json_safe(v: Any) -> Any:
"""递归把值清洗为 JSON 原生类型:NaN/±inf → None、datetime → ISO 串、
numpy 标量 Python 原生容器逐项处理
Starlette JSONResponse ``allow_nan=False`` 序列化任何 NaN/inf
漏出去都会让整个响应 500v1.32 实测/board-mac/overview 某行
sort_value=NaN 全端点 500 且带毒 payload 15s 缓存所有
DictResponse / 缓存写入路径都应先过本函数
"""
# bool 是 int 子类,须先判
if v is None or isinstance(v, bool | str | int):
return v
if isinstance(v, float):
return None if (math.isnan(v) or math.isinf(v)) else v
if hasattr(v, "isoformat"): # datetime/date/pd.Timestamp
return v.isoformat()
if hasattr(v, "item"): # numpy 标量(含 np.float32 NaN
return _json_safe(v.item())
if isinstance(v, dict):
return {str(k): _json_safe(val) for k, val in v.items()}
if isinstance(v, list | tuple):
return [_json_safe(item) for item in v]
return v
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Enums — mirror easy_tdx.models.enums but as string-based for REST clarity # Enums — mirror easy_tdx.models.enums but as string-based for REST clarity
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -143,19 +170,15 @@ class DictResponse(BaseModel):
@classmethod @classmethod
def from_dict(cls, d: dict[str, Any]) -> DictResponse: def from_dict(cls, d: dict[str, Any]) -> DictResponse:
"""序列化 dict,将其中的 DataFrame 转 records 格式""" """序列化 dictDataFrame 转 records,值递归清洗(NaN/inf → null 等)"""
import pandas as pd import pandas as pd
cleaned: dict[str, Any] = {} cleaned: dict[str, Any] = {}
for k, v in d.items(): for k, v in d.items():
if isinstance(v, pd.DataFrame): if isinstance(v, pd.DataFrame):
cleaned[k] = DataFrameResponse.from_dataframe(v).data cleaned[k] = DataFrameResponse.from_dataframe(v).data
elif hasattr(v, "isoformat"):
cleaned[k] = v.isoformat()
elif hasattr(v, "item"):
cleaned[k] = v.item()
else: else:
cleaned[k] = v cleaned[k] = _json_safe(v)
return cls(data=cleaned) return cls(data=cleaned)
+6 -4
View File
@@ -19,7 +19,7 @@ import logging
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
from easy_tdx.realtime.session import is_trading_time from easy_tdx.realtime.session import SHANGHAI_TZ, is_trading_time
from easy_tdx.web.sentiment_store import SentimentStore, get_sentiment_store from easy_tdx.web.sentiment_store import SentimentStore, get_sentiment_store
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -81,7 +81,9 @@ class SentimentSampler:
if df is None or df.empty: if df is None or df.empty:
raise RuntimeError("get_market_stat 返回空数据") raise RuntimeError("get_market_stat 返回空数据")
row = df.iloc[0] row = df.iloc[0]
now = datetime.now() # (date, minute) 键按沪市时区取"现在":与 is_trading_time 的时段判断
# 同一参照系,主机时区非 UTC+8(海外服务器)时不会整体错位
now = datetime.now(SHANGHAI_TZ)
self._store.insert( self._store.insert(
{ {
"date": now.year * 10000 + now.month * 100 + now.day, "date": now.year * 10000 + now.month * 100 + now.day,
@@ -141,7 +143,7 @@ class FundFlowSampler:
logger.info("FundFlowSampler 启动(间隔 %ss,交易日 14:45 后每日一条)", self._interval) logger.info("FundFlowSampler 启动(间隔 %ss,交易日 14:45 后每日一条)", self._interval)
while True: while True:
try: try:
now = datetime.now() now = datetime.now(SHANGHAI_TZ)
if is_trading_time(now) and (now.hour * 100 + now.minute) >= 1445: if is_trading_time(now) and (now.hour * 100 + now.minute) >= 1445:
await self._sample_once() await self._sample_once()
except asyncio.CancelledError: except asyncio.CancelledError:
@@ -153,7 +155,7 @@ class FundFlowSampler:
async def _sample_once(self) -> None: async def _sample_once(self) -> None:
from easy_tdx.mac.enums import BoardType from easy_tdx.mac.enums import BoardType
today = int(datetime.now().strftime("%Y%m%d")) today = int(datetime.now(SHANGHAI_TZ).strftime("%Y%m%d"))
if self._store.latest_fund_date() == today: if self._store.latest_fund_date() == today:
return # 当日已采样 return # 当日已采样
df = await self._client.get_board_ranking( df = await self._client.get_board_ranking(
+27 -3
View File
@@ -29,6 +29,18 @@ def _config_dir() -> Path:
return Path(os.environ.get("EASY_TDX_CONFIG_DIR", str(Path.home() / ".easy_tdx"))) return Path(os.environ.get("EASY_TDX_CONFIG_DIR", str(Path.home() / ".easy_tdx")))
def _real_or_zero(v: Any) -> float:
"""数值兜底:None/NaN/inf → 0.0REAL NOT NULL 列不吃 NaN——SQLite 会把
NaN 绑定成 NULL 而触发约束冲突整条写入失败"""
try:
f = float(v)
except (TypeError, ValueError):
return 0.0
if f != f or f in (float("inf"), float("-inf")):
return 0.0
return f
class SentimentStore: class SentimentStore:
"""情绪采样 SQLite 存储。""" """情绪采样 SQLite 存储。"""
@@ -116,7 +128,11 @@ class SentimentStore:
conn.close() conn.close()
def upsert_fund_day(self, date: int, boards: list[dict[str, Any]]) -> None: def upsert_fund_day(self, date: int, boards: list[dict[str, Any]]) -> None:
"""覆盖写入某日行业主力净流入排行(rank 按列表顺序 1 起)。""" """覆盖写入某日行业主力净流入排行(rank 按列表顺序 1 起)。
``main_net`` NaN上游行情缺失时 pandas 的空值口径会被 SQLite 存成
NULL 而触发 NOT NULL 约束冲突 整日采样失败这里统一落成 0.0
"""
with _write_lock: with _write_lock:
conn = self._connect() conn = self._connect()
try: try:
@@ -124,7 +140,13 @@ class SentimentStore:
conn.executemany( conn.executemany(
"INSERT INTO board_fund (date, rank, code, name, main_net) VALUES (?,?,?,?,?)", "INSERT INTO board_fund (date, rank, code, name, main_net) VALUES (?,?,?,?,?)",
[ [
(int(date), i + 1, str(b["code"]), str(b["name"]), float(b["main_net"])) (
int(date),
i + 1,
str(b["code"]),
str(b["name"]),
_real_or_zero(b["main_net"]),
)
for i, b in enumerate(boards) for i, b in enumerate(boards)
], ],
) )
@@ -155,7 +177,9 @@ class SentimentStore:
"rank": int(r["rank"]), "rank": int(r["rank"]),
"code": str(r["code"]), "code": str(r["code"]),
"name": str(r["name"]), "name": str(r["name"]),
"main_net": float(r["main_net"]), # 手工编辑/旧库可能存有 NULLSQLite 把 NaN 存成 NULL),
# float(None) 会 TypeError,兜成 0.0
"main_net": float(r["main_net"] or 0.0),
} }
) )
return list(grouped.values()) return list(grouped.values())
+15 -11
View File
@@ -18,9 +18,9 @@
- ``_run`` 写状态时**不假设** ``self._tasks[task_id]`` 仍在表中并发淘汰 - ``_run`` 写状态时**不假设** ``self._tasks[task_id]`` 仍在表中并发淘汰
可能在任务运行期间移除其条目``move_to_end`` try/except 容忍状态 可能在任务运行期间移除其条目``move_to_end`` try/except 容忍状态
写到本地 ``state`` 引用即使被淘汰也无害GC 回收 写到本地 ``state`` 引用即使被淘汰也无害GC 回收
- ``_evict_if_needed_locked`` 跳过 ``running`` 状态的任务正在执行的任务 - ``_evict_if_needed_locked`` 只淘汰 ``done/failed`` 终态条目pending 尚未
恰好是 OrderedDict 头部完成时才 move_to_end盲目 FIFO 淘汰会优先 起跑淘汰会产生 worker 跳过 + 查询侧恢复成永久 pending "幽灵任务"
杀掉在途任务淘汰改用最旧的 non-running 条目 running 在途无终态可淘汰时宁可持续超限
""" """
from __future__ import annotations from __future__ import annotations
@@ -156,11 +156,15 @@ class BacktestTaskRunner:
with self._lock: with self._lock:
memory_items = list(self._tasks.values()) memory_items = list(self._tasks.values())
seen = {s.task_id for s in memory_items} seen = {s.task_id for s in memory_items}
# 磁盘侧多取一些(覆盖内存 LRU 已淘汰的),再合并排序 # 磁盘侧多取一些(覆盖内存 LRU 已淘汰的),再合并排序;列表页不需要
# result,懒加载跳过 result_json 的 SELECT 与解析(几百条大结果时
# 可观省内存与事件循环停顿)
try: try:
disk_items = [ disk_items = [
self._dict_to_state(d) self._dict_to_state(d)
for d in get_task_store().list_recent(limit=limit + len(seen)) for d in get_task_store().list_recent(
limit=limit + len(seen), include_results=False
)
if d["task_id"] not in seen if d["task_id"] not in seen
] ]
except Exception: # noqa: BLE001 — 持久化故障不阻断列表查询 except Exception: # noqa: BLE001 — 持久化故障不阻断列表查询
@@ -284,20 +288,20 @@ class BacktestTaskRunner:
) )
def _evict_if_needed_locked(self) -> None: def _evict_if_needed_locked(self) -> None:
"""超过上限时丢弃最旧的 non-running 任务(调用方需持锁)。 """超过上限时丢弃最旧的终态(done/failed任务(调用方需持锁)。
running 任务不会被淘汰它们恰在 OrderedDict 头部但盲淘汰会杀在途任务 pending/running 一律不淘汰淘汰尚未起跑的 pending 会造成"幽灵任务"
只淘汰 pending/done/failed 中最旧者 worker 线程随后取不到状态直接跳过磁盘遗留的 pending 行会被查询
恢复成永远 pending全为在途任务时宁可不淘汰超限跳过
""" """
while len(self._tasks) > self._max_results: while len(self._tasks) > self._max_results:
# 找第一个 non-running 条目淘汰;若无则停止(全在 running,不强制淘汰)
evict_id: str | None = None evict_id: str | None = None
for tid, st in self._tasks.items(): for tid, st in self._tasks.items():
if st.status != "running": if st.status in ("done", "failed"):
evict_id = tid evict_id = tid
break break
if evict_id is None: if evict_id is None:
break # 全部 running,暂时无法淘汰 break # 无终态条目可淘汰
self._tasks.pop(evict_id, None) self._tasks.pop(evict_id, None)
+18 -6
View File
@@ -196,14 +196,26 @@ class TaskStore:
conn.close() conn.close()
return self._row_to_dict(row) if row is not None else None return self._row_to_dict(row) if row is not None else None
def list_recent(self, limit: int = 20) -> list[dict[str, Any]]: def list_recent(self, limit: int = 20, *, include_results: bool = True) -> list[dict[str, Any]]:
"""按 created_at 倒序列出最近 N 条任务摘要(含 result,供详情直取)。""" """按 created_at 倒序列出最近 N 条任务摘要
Args:
limit: 最多返回条数
include_results: True默认= result详情直取False =
SELECT/解析 ``result_json``列表页只展示摘要用它避免把
几百条大结果 JSON 拖进内存逐条解析result 字段为 None
"""
conn = self._connect() conn = self._connect()
try: try:
cols = (
"*"
if include_results
else "task_id, status, description, created_at, started_at, "
"finished_at, error, NULL"
)
cur = conn.execute( cur = conn.execute(
""" f"""
SELECT task_id, status, description, created_at, started_at, SELECT {cols}
finished_at, error, result_json
FROM backtest_tasks FROM backtest_tasks
ORDER BY created_at DESC, task_id DESC ORDER BY created_at DESC, task_id DESC
LIMIT ? LIMIT ?
@@ -314,7 +326,7 @@ class _NullTaskStore(TaskStore):
def load(self, task_id: str) -> dict[str, Any] | None: # noqa: ARG002 def load(self, task_id: str) -> dict[str, Any] | None: # noqa: ARG002
return None return None
def list_recent(self, limit: int = 20) -> list[dict[str, Any]]: # noqa: ARG002 def list_recent(self, limit: int = 20, *, include_results: bool = True) -> list[dict[str, Any]]: # noqa: ARG002
return [] return []
def delete(self, task_id: str) -> bool: # noqa: ARG002 def delete(self, task_id: str) -> bool: # noqa: ARG002
+303
View File
@@ -160,6 +160,309 @@ class TestClient:
assert result["ok"] is False and "API Key" in result["error"] assert result["ok"] is False and "API Key" in result["error"]
class TestApiUrlSchemeGuard:
"""api_url SSRF 防线:仅允许 http/https、禁止携带 userinfo。
背景_post_json urllib 直连用户可配的 api_url scheme 白名单时
``file:///...`` 可读本地文件llm.json 内含明文 key且格式异常分支会
回显响应前 300 字节``ftp://`` 与内网 http 可被当跳板
"""
def test_file_scheme_rejected(self):
with pytest.raises(ValueError, match="http"):
resolve_config(
LlmConfig(provider="custom", api_url="file:///C:/Users/x/llm.json", model="m")
)
def test_ftp_scheme_rejected(self):
with pytest.raises(ValueError, match="http"):
resolve_config(LlmConfig(provider="custom", api_url="ftp://internal-host/x", model="m"))
def test_missing_scheme_rejected(self):
with pytest.raises(ValueError, match="http"):
resolve_config(LlmConfig(provider="custom", api_url="api.deepseek.com/v1", model="m"))
def test_userinfo_rejected(self):
with pytest.raises(ValueError, match="user:pass"):
resolve_config(
LlmConfig(provider="custom", api_url="https://user:pass@api.x.com/v1", model="m")
)
def test_http_https_case_insensitive_allowed(self):
r = resolve_config(LlmConfig(provider="custom", api_url="HTTPS://Api.X.com/v1", model="m"))
assert r.api_url == "HTTPS://Api.X.com/v1"
r2 = resolve_config(
LlmConfig(provider="custom", api_url="http://gw.local:8000/v1", model="m")
)
assert r2.api_url == "http://gw.local:8000/v1"
def test_preset_urls_still_resolve(self, config_dir):
save_config(LlmConfig(provider="deepseek", api_key="sk-x-1234567890"))
r = resolve_config()
assert r.api_url == "https://api.deepseek.com/v1"
class TestHttpPostHardening:
"""HTTP 层加固:错误不回显原始 body、响应体大小上限。"""
def _raise_http_error(self, body: bytes, code: int = 401):
import io
import urllib.error
def fake_urlopen(req, timeout):
raise urllib.error.HTTPError(
req.full_url, code, "Unauthorized", hdrs=None, fp=io.BytesIO(body)
)
return fake_urlopen
def test_http_error_extracts_provider_message_only(self, monkeypatch):
"""错误响应只回显 provider 的 error.message,不回显原始 body 其他内容。"""
import json as _json
body = _json.dumps(
{"error": {"message": "Invalid API key", "internal_hint": "SECRET-STACK"}}
).encode()
monkeypatch.setattr(llm_mod.urllib.request, "urlopen", self._raise_http_error(body))
with pytest.raises(LlmError) as ei:
llm_mod._post_json("https://x/v1/chat/completions", {}, {"m": 1}, 5.0)
assert ei.value.status == 401
assert "Invalid API key" in str(ei.value)
assert "SECRET-STACK" not in str(ei.value)
def test_http_error_non_json_body_is_generic(self, monkeypatch):
"""非 JSON 错误页不给原始内容,只给通用 HTTP 状态描述。"""
body = b"<html><h1>gateway exploded with internal detail</h1></html>"
monkeypatch.setattr(llm_mod.urllib.request, "urlopen", self._raise_http_error(body))
with pytest.raises(LlmError) as ei:
llm_mod._post_json("https://x/v1/chat/completions", {}, {"m": 1}, 5.0)
assert "gateway exploded" not in str(ei.value)
assert "401" in str(ei.value)
def test_http_error_string_error_field_still_shown(self, monkeypatch):
"""error 为字符串的网关(如 {"error":"bad key"})仍展示该消息。"""
monkeypatch.setattr(
llm_mod.urllib.request, "urlopen", self._raise_http_error(b'{"error":"bad key"}')
)
with pytest.raises(LlmError, match="bad key") as ei:
llm_mod._post_json("https://x/v1/chat/completions", {}, {"m": 1}, 5.0)
assert ei.value.status == 401
def test_response_body_size_capped(self, config_dir, monkeypatch):
"""超过 2MB 的响应体中止解析(防异常网关撑爆内存),报可操作错误。"""
class _FakeResp:
def __init__(self, payload: bytes) -> None:
self._buf = payload
def read(self, n: int = -1) -> bytes:
if n < 0:
data, self._buf = self._buf, b""
return data
data, self._buf = self._buf[:n], self._buf[n:]
return data
def __enter__(self) -> _FakeResp:
return self
def __exit__(self, *exc: object) -> bool:
return False
big = b"x" * (llm_mod._MAX_RESPONSE_BYTES + 1)
monkeypatch.setattr(llm_mod.urllib.request, "urlopen", lambda req, timeout: _FakeResp(big))
with pytest.raises(LlmError, match="过大|上限"):
llm_mod._post_json("https://x/v1/chat/completions", {}, {"m": 1}, 5.0)
def test_normal_response_within_cap_parses(self, config_dir, monkeypatch):
class _FakeResp:
def __init__(self, payload: bytes) -> None:
self._buf = payload
def read(self, n: int = -1) -> bytes:
if n < 0:
data, self._buf = self._buf, b""
return data
data, self._buf = self._buf[:n], self._buf[n:]
return data
def __enter__(self) -> _FakeResp:
return self
def __exit__(self, *exc: object) -> bool:
return False
payload = b'{"choices": [{"message": {"content": "OK"}}]}'
monkeypatch.setattr(
llm_mod.urllib.request, "urlopen", lambda req, timeout: _FakeResp(payload)
)
data = llm_mod._post_json("https://x/v1/chat/completions", {}, {"m": 1}, 5.0)
assert data["choices"][0]["message"]["content"] == "OK"
class TestSaveConfigAtomic:
def test_replace_failure_preserves_old_file(self, config_dir, monkeypatch):
"""os.replace 失败(磁盘满等)时旧配置原样保留,不留临时文件。"""
save_config(LlmConfig(provider="deepseek", api_key="sk-old-1234567890"))
def boom(src, dst):
raise OSError("disk full")
monkeypatch.setattr(llm_mod.os, "replace", boom)
with pytest.raises(OSError):
save_config(LlmConfig(provider="kimi", api_key="sk-new-9999999999"))
assert load_config().api_key == "sk-old-1234567890" # 旧配置未被破坏
leftovers = [p.name for p in config_dir.iterdir() if p.name != "llm.json"]
assert leftovers == [] # 失败的临时文件已清理
class TestLoadConfigFieldDefense:
"""手工编辑 llm.json 的脏字段不得打挂 load_config(全部 /llm/* 依赖它)。"""
def test_null_fields_fall_back_to_defaults(self, config_dir):
import json as _json
(config_dir / "llm.json").write_text(
_json.dumps(
{
"provider": None,
"api_url": None,
"api_key": None,
"model": None,
"temperature": None,
"max_tokens": None,
"timeout": None,
"system_prompt": None,
}
),
encoding="utf-8",
)
cfg = load_config() # 旧码:float(None) TypeError
assert cfg.provider == "deepseek"
assert cfg.api_url == "" and cfg.api_key == "" and cfg.model == ""
assert cfg.temperature == 0.3
assert cfg.max_tokens == 16000
assert cfg.timeout == 180.0
assert cfg.system_prompt == LlmConfig.system_prompt
def test_wrong_types_fall_back_with_warning(self, config_dir, caplog):
import json as _json
(config_dir / "llm.json").write_text(
_json.dumps(
{
"temperature": "abc",
"max_tokens": "fast",
"timeout": [],
"provider": 123,
"system_prompt": 456,
}
),
encoding="utf-8",
)
with caplog.at_level("WARNING", logger="easy_tdx.ai.llm"):
cfg = load_config() # 旧码:float("abc") ValueError
assert cfg.temperature == 0.3
assert cfg.max_tokens == 16000
assert cfg.timeout == 180.0
assert cfg.provider == "deepseek" # 非字符串 provider 回退默认
assert cfg.system_prompt == LlmConfig.system_prompt
assert any("temperature" in r.message for r in caplog.records)
def test_non_finite_and_out_of_range_fall_back(self, config_dir):
import json as _json
(config_dir / "llm.json").write_text(
_json.dumps({"temperature": 1e999, "timeout": -5, "max_tokens": 0}), # 1e999→inf
encoding="utf-8",
)
cfg = load_config() # 旧码:inf temperature 会一路写进请求 payload
assert cfg.temperature == 0.3
assert cfg.timeout == 180.0
assert cfg.max_tokens == 16000
def test_string_numbers_leniently_coerced(self, config_dir):
import json as _json
(config_dir / "llm.json").write_text(
_json.dumps({"temperature": "0.7", "max_tokens": "8192.9", "timeout": "60"}),
encoding="utf-8",
)
cfg = load_config()
assert cfg.temperature == 0.7
assert cfg.max_tokens == 8192
assert cfg.timeout == 60.0
class TestAnthropicRobustness:
"""anthropic 协议与 openai 口径对齐:绝不静默返回空正文。"""
def _client(self) -> LlmClient:
return LlmClient(LlmConfig(provider="claude", api_key="sk-ant-123456789"))
def test_thinking_only_blocks_raise_actionable(self, config_dir, monkeypatch):
"""仅 thinking 块(max_tokens 被思考耗尽)→ 可操作错误,而非空串成功。"""
def fake_post(url, headers, payload, timeout):
return {
"content": [{"type": "thinking", "thinking": "思考" * 200}],
"stop_reason": "max_tokens",
}
monkeypatch.setattr(llm_mod, "_post_json", fake_post)
with pytest.raises(LlmError, match="思考链"):
asyncio.run(self._client().chat("hi"))
def test_content_as_plain_string_accepted(self, config_dir, monkeypatch):
"""部分网关把 content 放字符串而非块列表——正常取正文。"""
def fake_post(url, headers, payload, timeout):
return {"content": "纯字符串回复"}
monkeypatch.setattr(llm_mod, "_post_json", fake_post)
assert asyncio.run(self._client().chat("hi")) == "纯字符串回复"
def test_mixed_blocks_text_extracted(self, config_dir, monkeypatch):
def fake_post(url, headers, payload, timeout):
return {
"content": [
{"type": "thinking", "thinking": "思考"},
{"type": "text", "text": "正文"},
],
"stop_reason": "end_turn",
}
monkeypatch.setattr(llm_mod, "_post_json", fake_post)
assert asyncio.run(self._client().chat("hi")) == "正文"
def test_missing_content_raises_llm_error(self, config_dir, monkeypatch):
"""content 缺失 → LlmError(旧码 AttributeError 裸 500)。"""
def fake_post(url, headers, payload, timeout):
return {"stop_reason": "end_turn"}
monkeypatch.setattr(llm_mod, "_post_json", fake_post)
with pytest.raises(LlmError, match="格式异常"):
asyncio.run(self._client().chat("hi"))
def test_empty_blocks_generic_error_without_raw_echo(self, config_dir, monkeypatch):
def fake_post(url, headers, payload, timeout):
return {"content": [{"type": "tool_use", "id": "tool_1", "secret": "S3CR3T"}]}
monkeypatch.setattr(llm_mod, "_post_json", fake_post)
with pytest.raises(LlmError, match="content 为空") as ei:
asyncio.run(self._client().chat("hi"))
assert "S3CR3T" not in str(ei.value) # 不回显原始响应体
def test_empty_string_content_with_max_tokens_stop(self, config_dir, monkeypatch):
def fake_post(url, headers, payload, timeout):
return {"content": "", "stop_reason": "max_tokens"}
monkeypatch.setattr(llm_mod, "_post_json", fake_post)
with pytest.raises(LlmError, match="截断"):
asyncio.run(self._client().chat("hi"))
def test_provider_presets_cover_major_vendors(): def test_provider_presets_cover_major_vendors():
vendors = [ vendors = [
"deepseek", "deepseek",
+64
View File
@@ -141,3 +141,67 @@ class TestPortfolioCLIFlags:
assert "--evaluate" in result.output assert "--evaluate" in result.output
assert "--wf" in result.output assert "--wf" in result.output
assert "--auto-fees" in result.output assert "--auto-fees" in result.output
class TestIgnoredFlagWarnings:
"""静默忽略的旗标组合必须显式告警(stderr),不得无提示吞掉。"""
def test_combo_with_wf_warns(self):
"""--combo-strategies + --wf:旧码静默忽略,新码应告警并继续报策略文件错误。"""
from easy_tdx.backtest.cli import backtest
runner = CliRunner()
result = runner.invoke(
backtest,
[
"SZ",
"000001",
"--combo-strategies",
"nope_a.py,nope_b.py",
"--wf",
],
)
assert "已忽略" in result.output
assert result.exit_code != 0 # 随后仍因策略文件不存在报错(无网络依赖)
def test_combo_with_evaluate_warns(self):
from easy_tdx.backtest.cli import backtest
runner = CliRunner()
result = runner.invoke(
backtest,
[
"SZ",
"000001",
"--combo-strategies",
"nope_a.py,nope_b.py",
"--evaluate",
],
)
assert "已忽略" in result.output
def test_single_with_wf_no_warning(self):
"""单策略 + --wf 是合法组合,不告警(在策略校验失败前无「已忽略」字样)。"""
from easy_tdx.backtest.cli import backtest
runner = CliRunner()
result = runner.invoke(backtest, ["SZ", "000001", "--wf"])
assert "已忽略" not in result.output
def test_optimize_all_with_param_warns(self):
"""--all + --param:旧码静默忽略自定义网格,新码应告警。"""
from easy_tdx.backtest.cli import optimize
runner = CliRunner()
# 用非法市场名在联网取数前中断,仅验证告警已发出
result = runner.invoke(optimize, ["XX", "000001", "--all", "--param", "fast=5,10"])
assert "忽略" in result.output
def test_optimize_single_with_param_no_warning(self):
from easy_tdx.backtest.cli import optimize
runner = CliRunner()
result = runner.invoke(
optimize, ["XX", "000001", "--strategy", "ma_cross", "--param", "fast=5,10"]
)
assert "已忽略" not in result.output
@@ -197,6 +197,31 @@ def test_evaluate_strategy_auto_fees_for_etf():
assert report["config"]["auto_fees"] is True assert report["config"]["auto_fees"] is True
def test_evaluate_portfolio_auto_fees_fitness_per_stock_symbol():
"""组合体检逐段应按各标的品种解析费率(auto_fees 与主回测同口径)。
回归evaluate_portfolio per-stock FitnessEngine 漏传 symbol
auto_fees 不生效ETF 组合的三段体检被按股票口径错收印花税卖方 0.001
与组合回测主路径PortfolioBacktestEngine 逐标的 resolve_fee_model不一致
期望报告里的三段体检与正确传入 symbol 的独立体检逐段一致
"""
from easy_tdx.backtest.benchmark import evaluate_portfolio
from easy_tdx.backtest.portfolio_engine import StockData
df = _df(240, drift=0.002)
stocks = [StockData("159915", "SZ", df)] # ETF:法定免印花税
report = evaluate_portfolio(
_CycleTrader, stocks, total_cash=100_000, auto_fees=True, n_windows=3
)
expected = FitnessEngine(symbol="SZ159915", strategy=_CycleTrader, auto_fees=True).evaluate(df)
actual_returns = [seg["total_return"] for seg in report["fitness"]["segments"]]
expected_returns = [seg.total_return for seg in expected.segments]
assert len(actual_returns) == 3
for actual, exp in zip(actual_returns, expected_returns):
assert actual == pytest.approx(exp, abs=1e-9)
# ── evaluate_portfoliov1.31 组合级一条龙)─────────────────────────────────── # ── evaluate_portfoliov1.31 组合级一条龙)───────────────────────────────────
def _stocks_for_portfolio() -> list[Any]: def _stocks_for_portfolio() -> list[Any]:
from easy_tdx.backtest.portfolio_engine import StockData from easy_tdx.backtest.portfolio_engine import StockData
+90
View File
@@ -157,3 +157,93 @@ def test_momentum_score_helper():
score = momentum_score(10)(df) score = momentum_score(10)(df)
assert score > 0 assert score > 0
assert momentum_score(10)(_stock(5)) == 0.0 # 数据不足 → 0 assert momentum_score(10)(_stock(5)) == 0.0 # 数据不足 → 0
# ── 回归:停牌/初始调仓/历史不足(审查修复) ─────────────────────────────────
def _bar_frame(dates: pd.DatetimeIndex, prices: list[float]) -> pd.DataFrame:
closes = np.asarray(prices, dtype=float)
return pd.DataFrame(
{
"datetime": dates[: len(closes)],
"open": closes * 0.999,
"high": closes * 1.01,
"low": closes * 0.98,
"close": closes,
"vol": 1e6,
}
)
def test_rotation_suspension_defers_fill_to_resume_open():
"""停牌日挂单顺延:成交日=复牌日、成交价=复牌开盘(旧码在停牌日按停牌前价格成交)。"""
dates = pd.date_range("2024-01-01", periods=12, freq="D")
a = _bar_frame(dates, [10 + 0.05 * i for i in range(12)])
# B01-01..01-07 有 bar01-07 收盘崩盘跌出排名),01-08 停牌(下标 7 无 bar),
# 01-09 复牌开盘 -30%(下标 8 = 4.2
b_prices = [10, 10.1, 10.2, 10.3, 10.4, 10.5, 6.0, 4.9, 4.2, 4.3, 4.3, 4.3]
b = pd.DataFrame(
[
{
"datetime": dates[i],
"open": p * 0.999,
"high": p * 1.01,
"low": p * 0.98,
"close": p,
"vol": 1e6,
}
for i, p in enumerate(b_prices)
if i != 7 # 01-08 停牌,无 bar
]
)
engine = RotationEngine(
{"SH:600001": a, "SZ:000002": b},
momentum_score(2),
slots=1,
refresh="daily",
keep_rank=1,
)
res = engine.run()
sells_b = [t for t in res.trades if t["symbol"] == "SZ:000002" and t["direction"] == "SELL"]
assert len(sells_b) == 1
sell = sells_b[0]
assert sell["datetime"] == "2024-01-09" # 旧码记 2024-01-08(停牌日)
assert sell["price"] == pytest.approx(4.2 * 0.999) # 旧码记 6.0 * 0.999(停牌前开盘)
# 复牌前净值按最后已知收盘估值,不应把持仓价值清零
eq_by_date = {r["datetime"]: r for r in res.equity_curve}
assert eq_by_date["2024-01-08"]["position_value"] > 0
def test_rotation_day0_counts_as_first_rebalance():
"""day0 即为首个调仓日(排名只用 ≤day0 数据,次日开盘执行),不再人为空仓一天。"""
pool = _pool({f"SH:60000{i}": 0.002 for i in range(5)}, n=40)
res = RotationEngine(pool, momentum_score(5), slots=3, refresh="weekly").run()
# 旧码首个调仓日是下一 ISO 周 2024-01-08
assert res.rebalance_dates[0] == "2024-01-01"
# next_open 语义:任何成交不早于第二个交易日(day0 信号次日执行)
if res.trades:
assert min(t["datetime"] for t in res.trades) > res.rebalance_dates[0]
def test_rotation_new_listing_not_bought_on_zero_score():
"""历史不足(<5 根)从买入候选剔除:次新股 0 分不得排在负动量标的之前被买入。"""
n = 60
dates = pd.date_range("2024-01-01", periods=n, freq="B")
declining = 100.0 * np.cumprod(np.full(n, 1.0 - 0.005))
a = _bar_frame(dates, list(declining)) # 长历史持续阴跌,动量为负
b = _bar_frame(dates, [10.0, 10.0, 10.0]) # 末段才上市,全程 idx<5
res = RotationEngine(
{"SH:600001": a, "SZ:000002": b},
momentum_score(5),
slots=1,
refresh="daily",
).run()
buys_b = [t for t in res.trades if t["symbol"] == "SZ:000002" and t["direction"] == "BUY"]
assert buys_b == [] # 旧码 B 以 0 分登顶被买入
# A 作为唯一有效候选被正常买入
assert any(t["symbol"] == "SH:600001" and t["direction"] == "BUY" for t in res.trades)
+97
View File
@@ -174,3 +174,100 @@ def test_wf_auto_fes_passed_through():
assert wf_engine._engine_kwargs["auto_fees"] is True assert wf_engine._engine_kwargs["auto_fees"] is True
wf = wf_engine.run(_trend_df(300)) wf = wf_engine.run(_trend_df(300))
assert len(wf.windows) == 3 assert len(wf.windows) == 3
# ── 回归:窗口绩效口径 / 聚合方向 / 切窗下限 / 失败日志 / int 日期 ────────────
def test_wf_window_metrics_exclude_context_bars():
"""上下文只做指标预热:窗口绩效指标不随 context_bars 变化。
旧码把 context 恒定现金段一并喂给 PerformanceAnalyzersharpe/年化/波动
被稀释同窗 total_return 相同而 sharpe 相差近一倍
"""
df = _trend_df(500)
wf0 = WalkForwardEngine(_BuyFirstBar, n_windows=5, warmup_ratio=0.3, context_bars=0).run(df)
wf60 = WalkForwardEngine(_BuyFirstBar, n_windows=5, warmup_ratio=0.3, context_bars=60).run(df)
assert len(wf0.windows) == len(wf60.windows) == 5
for w0, w60 in zip(wf0.windows, wf60.windows):
assert w0.total_return == pytest.approx(w60.total_return)
assert w0.sharpe == pytest.approx(w60.sharpe)
assert w0.max_drawdown == pytest.approx(w60.max_drawdown)
assert w0.performance["annual_return"] == pytest.approx(w60.performance["annual_return"])
assert w0.performance["volatility"] == pytest.approx(w60.performance["volatility"])
def test_wf_worst_drawdown_is_max_not_min():
"""worst_drawdown 应取各窗最深回撤(max);旧码 min 取成最浅回撤。"""
from easy_tdx.backtest.walkforward import WalkForwardResult, WalkForwardWindow
result = WalkForwardResult(n_windows=3, warmup_ratio=0.3)
for i, dd in enumerate((0.05, 0.40, 0.11)):
result.windows.append(
WalkForwardWindow(
index=i,
start="2024-01-01",
end="2024-02-01",
bars=20,
total_return=0.01,
sharpe=1.0,
max_drawdown=dd,
total_trades=2,
win_rate=0.5,
)
)
WalkForwardEngine._aggregate(result)
assert result.worst_drawdown == pytest.approx(0.40)
def test_wf_windows_below_min_bars_skipped():
"""单窗实际 bar 数 < 20 时跳过(与 docstring「每窗 ≥ 20 根」口径一致)。"""
wf = WalkForwardEngine(_BuyFirstBar, n_windows=9).run(_trend_df(220))
assert wf.windows == []
class _BoomStrategy(Strategy):
"""init 即抛错:单窗失败应记 warning 而非静默跳过。"""
def init(self) -> None:
raise RuntimeError("boom")
def next(self) -> None:
pass
def test_wf_window_failure_logs_warning(caplog):
"""单窗回测异常记 warning(含窗号与异常摘要),不拖垮整组。"""
import logging
with caplog.at_level(logging.WARNING, logger="easy_tdx.backtest.walkforward"):
wf = WalkForwardEngine(_BoomStrategy, n_windows=3).run(_trend_df(300))
assert wf.windows == []
msgs = [r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING]
assert any("第 0 窗" in m and "boom" in m for m in msgs), msgs
def test_wf_int_yyyymmdd_date_column_window_labels():
"""datetime 为 int YYYYMMDDTDX 日线原样)时窗口起止日期正确。
旧码 pd.Timestamp(int) 按纳秒换算窗口日期全变 1970-01-01
"""
n = 300
dates = pd.date_range("2023-01-02", periods=n, freq="B")
close = 10.0 * np.linspace(1.0, 2.0, n)
df = pd.DataFrame(
{
"datetime": dates.strftime("%Y%m%d").astype(int),
"open": close * 0.999,
"high": close * 1.01,
"low": close * 0.99,
"close": close,
"vol": 1000.0,
}
)
wf = WalkForwardEngine(_BuyFirstBar, n_windows=3, context_bars=10).run(df)
assert len(wf.windows) == 3
eval_start = int(n * 0.3)
assert wf.windows[0].start == dates[eval_start].strftime("%Y-%m-%d")
assert wf.windows[0].end == dates[eval_start + (n - eval_start) // 3 - 1].strftime("%Y-%m-%d")
assert not wf.windows[0].start.startswith("1970")
+86 -3
View File
@@ -46,16 +46,28 @@ def _fake_rows(n: int, end: str = "2026-09-04") -> list[list[str]]:
return [[d, "10.0", "11.0", "9.5", "10.5", "100000", "1050000.0", "1"] for d in dates] return [[d, "10.0", "11.0", "9.5", "10.5", "100000", "1050000.0", "1"] for d in dates]
def _weekly_rows(n: int, end: str = "2026-09-04") -> list[list[str]]:
"""n 个周线行(无 tradestatus 列,与真实 W/M 返回一致)。"""
dates = pd.date_range(end=end, periods=n, freq="W-FRI").strftime("%Y-%m-%d")
return [[d, "10.0", "11.0", "9.5", "10.5", "500000", "5250000.0"] for d in dates]
def _install_fake_bs( def _install_fake_bs(
rows: list[list[str]] | None, rows: list[list[str]] | None,
captured: dict, captured: dict,
*, *,
query_error: bool = False, query_error: bool = False,
login_error: bool = False,
) -> types.ModuleType: ) -> types.ModuleType:
mod = types.ModuleType("baostock") mod = types.ModuleType("baostock")
def _login(): # type: ignore[no-untyped-def] def _login(): # type: ignore[no-untyped-def]
captured["login"] = captured.get("login", 0) + 1 captured["login"] = captured.get("login", 0) + 1
if login_error:
result = _FakeLoginResult()
result.error_code = "10001"
result.error_msg = "用户登录失败"
return result
return _FakeLoginResult() return _FakeLoginResult()
mod.login = _login # type: ignore[attr-defined] mod.login = _login # type: ignore[attr-defined]
@@ -166,12 +178,83 @@ def test_unsupported_inputs(fake_bs):
assert "calls" not in fake_bs assert "calls" not in fake_bs
def test_query_error_returns_none(fake_bs): def test_query_error_raises_runtimeerror_and_logs(fake_bs, caplog):
"""baostock 查询失败:返回 None 且不向上抛(兜底失败不改变原错误路径)。""" """baostock 查询失败error_code≠0):记 warning 并抛 RuntimeError。
回归旧实现吞掉所有异常静默返回 None`--source baostock` 显式使用时
故障被伪装成"无数据"sync skipped 而非 failedauto 兜底路径
web/routers/bars.py except Exception 包裹调用不受影响
"""
_install_fake_bs([], fake_bs, query_error=True) _install_fake_bs([], fake_bs, query_error=True)
from easy_tdx.sources import baostock as bs_source from easy_tdx.sources import baostock as bs_source
assert bs_source.fetch_bars("SH", "600519", "DAY", 0, 5, "QFQ") is None with pytest.raises(RuntimeError, match="baostock 拉取失败"):
bs_source.fetch_bars("SH", "600519", "DAY", 0, 5, "QFQ")
assert "baostock 拉取失败" in caplog.text
def test_login_failure_raises_runtimeerror(fake_bs, monkeypatch: pytest.MonkeyPatch):
"""baostock 登录失败:同样 warning + RuntimeError(不再静默)。"""
_install_fake_bs([], fake_bs, login_error=True)
from easy_tdx.sources import baostock as bs_source
with pytest.raises(RuntimeError, match="拉取失败"):
bs_source.fetch_bars("SZ", "000001", "DAY", 0, 5, "QFQ")
def test_weekly_monthly_fields_exclude_tradestatus(fake_bs):
"""W/M 请求不传 tradestatusbaostock 实测 error_code=10004012 报错),
日线保留"""
from easy_tdx.sources import baostock as bs_source
weekly_rows = _weekly_rows(6)
_install_fake_bs(weekly_rows, fake_bs)
df = bs_source.fetch_bars("SH", "600519", "WEEK", 0, 5, "QFQ")
assert df is not None and len(df) == 5
assert fake_bs["frequency"] == "w"
assert "tradestatus" not in fake_bs["fields"]
assert list(df.columns) == ["date", "open", "close", "high", "low", "vol", "amount"]
_install_fake_bs(_weekly_rows(6), fake_bs)
df = bs_source.fetch_bars("SH", "600519", "MONTH", 0, 5, "QFQ")
assert df is not None and len(df) == 5
assert fake_bs["frequency"] == "m"
assert "tradestatus" not in fake_bs["fields"]
# 日线仍保留 tradestatus(停牌剔除依赖它)
_install_fake_bs(_fake_rows(6), fake_bs)
bs_source.fetch_bars("SH", "600519", "DAY", 0, 5, "QFQ")
assert "tradestatus" in fake_bs["fields"]
def test_weekly_suspension_dropped_by_volume(fake_bs):
"""WEEK 无 tradestatus 列时,停牌/无成交周(volume=0)按 vol>0 兜底剔除。"""
from easy_tdx.sources import baostock as bs_source
rows = _weekly_rows(6)
rows[2][5] = "0" # volume=0 的停牌周
_install_fake_bs(rows, fake_bs)
df = bs_source.fetch_bars("SZ", "000001", "WEEK", 0, 10, "QFQ")
assert df is not None and len(df) == 5
assert (df["vol"] > 0).all()
def test_index_volume_converted_to_lots(fake_bs):
"""is_index=True:指数 vol 股→手(÷100),对齐 /bars/index 契约。
实测 sh.000001 2026-09-04baostock volume=53,728,616,100
÷100 = 537,286,161 TDX 指数日线口径为手
"""
from easy_tdx.sources import baostock as bs_source
df = bs_source.fetch_bars("SH", "000001", "DAY", 0, 5, "NONE", is_index=True)
assert df is not None
assert (df["vol"] == 1000.0).all() # 100000 股 ÷100 = 1000 手
# 默认(个股路径)不换算
_install_fake_bs(_fake_rows(6), fake_bs)
df = bs_source.fetch_bars("SH", "600519", "DAY", 0, 5, "QFQ")
assert (df["vol"] == 100000).all()
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
+36
View File
@@ -379,3 +379,39 @@ def test_hotspot_correlation_building_passthrough():
assert body["status"] in ("building", "error", "ready") # 单机假客户端极快时可能已完成 assert body["status"] in ("building", "error", "ready") # 单机假客户端极快时可能已完成
if body["status"] == "building": if body["status"] == "building":
assert 0.0 <= body["progress"] <= 1.0 assert 0.0 <= body["progress"] <= 1.0
# ── 时区统一(v1.32.6):日历日一律取沪市时区,与主机时区无关 ────────────────
# 模块导入时捕获真实实现(autouse fixture 会把 _today_str 换成钉死的 lambda
_board_mac_mod = pytest.importorskip("easy_tdx.web.routers.board_mac")
_REAL_TODAY_STR = _board_mac_mod._today_str
def test_today_str_uses_shanghai_tz(monkeypatch):
"""_today_str 必须用 SHANGHAI_TZ 取"今天"(旧实现用主机本地时区)。
海外机器 UTC-5上北京时间 09-06 02:00 时本地还是 09-05
旧实现会把热点矩阵的"今日"判定错一天
"""
from datetime import datetime
pytest.importorskip("fastapi")
from easy_tdx.realtime.session import SHANGHAI_TZ
from easy_tdx.web.routers import board_mac
# 恢复被 autouse fixture 钉住的真实现
monkeypatch.setattr(board_mac, "_today_str", _REAL_TODAY_STR)
captured: dict = {}
class _FakeDatetime:
@classmethod
def now(cls, tz=None):
captured["tz"] = tz
return datetime(2026, 9, 6, 2, 0, tzinfo=tz) if tz else datetime(2026, 9, 6, 2, 0)
monkeypatch.setattr(board_mac, "datetime", _FakeDatetime)
assert board_mac._today_str() == "2026-09-06"
assert captured["tz"] is SHANGHAI_TZ
+72
View File
@@ -219,3 +219,75 @@ def test_overview_zero_pre_close_change_pct_null():
row = resp.json()["data"]["rows"][0] row = resp.json()["data"]["rows"][0]
assert row["change_pct"] is None assert row["change_pct"] is None
assert row["leader_change_pct"] is None assert row["leader_change_pct"] is None
def test_overview_cache_key_includes_count():
"""缓存键须含 count:不同 count 的请求在 TTL 内不互相命中。
旧实现缓存键只有 (board_type, metrics)先到的小 count 请求会把大 count
的响应"污染"成少数行15s TTL
"""
pytest.importorskip("fastapi")
from fastapi.testclient import TestClient
frames = {
"CHANGE_PCT": _board_df(
[
_board_row("881001", "软件服务", 5000.0, 4900.0),
_board_row("881002", "半导体", 3000.0, 2950.0),
]
),
}
class _CountingFake(_FakeOverviewMacClient):
"""尊重 count 参数(与真实客户端一致地截断行数)。"""
async def get_board_list(self, board_type=None, count=10000, sort_column=None):
df = await super().get_board_list(
board_type=board_type, count=count, sort_column=sort_column
)
return df.head(count) if df is not None else df
fake = _CountingFake(frames)
with TestClient(_overview_app(fake)) as client:
r_small = client.get("/api/v1/board-mac/overview", params={"board_type": "HY", "count": 1})
assert r_small.status_code == 200
assert r_small.json()["data"]["count"] == 1
r_big = client.get("/api/v1/board-mac/overview", params={"board_type": "HY", "count": 2})
assert r_big.status_code == 200
# 不允许命中 count=1 的缓存
assert r_big.json()["data"]["count"] == 2
assert fake.calls.count("CHANGE_PCT") == 2 # 两个 count 各拉一次
def test_overview_nan_payload_cleaned_before_cache():
"""坏值(NaN)行不产生 500,且写入缓存前已清洗(缓存里不留 NaN)。
旧实现sort_value=NaN payload NaN Starlette allow_nan=False
序列化 500且带毒 payload 先入 15s 缓存TTL 内持续 500
"""
pytest.importorskip("fastapi")
from fastapi.testclient import TestClient
from easy_tdx.web.routers import board_mac
nan = float("nan")
frames = {
"CHANGE_PCT": _board_df([_board_row("881001", "软件服务", 5000.0, 4900.0)]),
"SPEED": _board_df([_board_row("881001", "软件服务", 5000.0, 4900.0, sort_value=nan)]),
}
fake = _FakeOverviewMacClient(frames)
with TestClient(_overview_app(fake)) as client:
r1 = _get_overview(client)
assert r1.status_code == 200
assert r1.json()["data"]["rows"][0]["speed"] is None
# 坏 payload 不得入缓存:缓存里的 speed 应已是 None
cached = board_mac._overview_cache[("HY", ("SPEED", "CHANGE_20D"), 2000)][1]
assert cached["rows"][0]["speed"] is None
r2 = _get_overview(client) # 命中缓存也不再 500
assert r2.status_code == 200
assert r2.json()["data"]["rows"][0]["speed"] is None
assert fake.calls.count("SPEED") == 1
+32
View File
@@ -177,6 +177,38 @@ def test_parse_xml_error_page_raises() -> None:
parse_xml("404 page,非 XML 内容") parse_xml("404 page,非 XML 内容")
def test_parse_xml_structure_change_raises() -> None:
"""XML 合法但 0 个 <data> 节点且存在其他子结构:抛 CcpmError(官网改版信号)。
回归旧实现静默返回空表改版后 CLI/Web 层展示"无数据"而非报错
改版长期无人察觉正常发布日必有 <data>真实无数据日走 302
"""
from easy_tdx.ccpm import CcpmError, parse_xml
changed = (
'<?xml version="1.0" encoding="UTF-8"?>'
"<positionRank><record><instrumentid>IF2609</instrumentid></record></positionRank>"
)
with pytest.raises(CcpmError, match="结构可能已变更"):
parse_xml(changed)
def test_parse_xml_empty_root_is_no_data() -> None:
"""空 <positionRank/>(无任何子节点):按当日无数据处理,返回空表。"""
from easy_tdx.ccpm import parse_xml
assert parse_xml('<?xml version="1.0" encoding="UTF-8"?><positionRank></positionRank>') == []
def test_get_rank_structure_change_raises(isolated_config, monkeypatch) -> None:
"""结构变更经 get_rank 透传为 CcpmError(不缓存、不返回空表伪装成功)。"""
from easy_tdx.ccpm import CcpmClient, CcpmError
_mock_fetch(monkeypatch, sample="<positionRank><unknown>1</unknown></positionRank>")
with pytest.raises(CcpmError, match="结构可能已变更"):
CcpmClient().get_rank("IF", "2026-09-02")
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# 日期归一化 # 日期归一化
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
+251
View File
@@ -0,0 +1,251 @@
"""CLI 参数校验与退出码测试(cmd_warehouse / cmd_formula#审查修复)。
覆盖
- ``市场:代码`` 解析辅助缺冒号/空段 click.BadParameter而非裸 ValueError
- ``warehouse sync``failed>0 exit 1对齐 ccpm 口径summary source 标注
``--period`` Choice 限定``--source baostock`` 不支持分钟周期时参数层报错
- ``warehouse check``先做只支持一个标的校验再解析单标的缺冒号也报错
- ``formula screen``缺冒号标的前置报错不再裸 traceback
"""
from __future__ import annotations
import json
from typing import Any
import pytest
from click.testing import CliRunner
from easy_tdx.cli.cmd_formula import _parse_symbol as _parse_symbol_formula
from easy_tdx.cli.cmd_warehouse import (
_BAOSTOCK_PERIODS,
_PERIOD_CHOICES,
warehouse_check,
warehouse_sync,
)
from easy_tdx.cli.cmd_warehouse import (
_parse_symbol as _parse_symbol_warehouse,
)
class TestParseSymbol:
@pytest.mark.parametrize("parse", [_parse_symbol_formula, _parse_symbol_warehouse])
def test_valid(self, parse):
assert parse("SH:600519") == ("SH", "600519")
assert parse(" sz:000001 ") == ("SZ", "000001")
@pytest.mark.parametrize("parse", [_parse_symbol_formula, _parse_symbol_warehouse])
@pytest.mark.parametrize(
"bad", ["SH600519", "SH:", ":600519", ":", "SH 600519", "SH:600519:extra"]
)
def test_malformed_raises_bad_parameter(self, parse, bad):
# 旧码:sym.split(":", 1) 裸 ValueError"SH:600519:extra" 旧码能过但语义错,也收紧)
from click import BadParameter
with pytest.raises(BadParameter, match="市场:代码"):
parse(bad)
class _FakeWarehouse:
"""context-manager 形假的 KlineWarehousecmd 只当透传对象用)。"""
def __init__(self, *a: Any, **k: Any) -> None:
pass
def __enter__(self) -> _FakeWarehouse:
return self
def __exit__(self, *exc: object) -> bool:
return False
def health_check(self, market: str | None = None, code: str | None = None) -> dict[str, Any]:
return {"issues": [], "market": market, "code": code}
class _FakeSyncer:
"""可编程结果假 WarehouseSyncercmd_warehouse 从 easy_tdx.warehouse 导入它)。"""
result: dict[str, Any] = {}
def __init__(self, *a: Any, **k: Any) -> None:
pass
def sync(self, symbols: Any, period: str, progress: Any = None) -> dict[str, Any]:
if progress is not None:
progress(1, len(symbols), str(symbols[0]))
return dict(self.result)
@pytest.fixture()
def patched_warehouse(monkeypatch):
"""打桩 cmd_warehouse 的全部外部依赖(仓库 / TDX 客户端 / 同步器)。"""
import easy_tdx.cli.cmd_warehouse as cw
import easy_tdx.cli.conn as conn_mod
import easy_tdx.warehouse as wh_pkg
class _FakeMacClient:
def __enter__(self) -> _FakeMacClient:
return self
def __exit__(self, *exc: object) -> bool:
return False
monkeypatch.setattr(cw, "_require_warehouse", lambda db_path: _FakeWarehouse())
monkeypatch.setattr(conn_mod, "get_mac_client", lambda: _FakeMacClient())
monkeypatch.setattr(wh_pkg, "WarehouseSyncer", _FakeSyncer)
return cw
class TestWarehouseSync:
def _invoke(self, *args: str):
return CliRunner().invoke(warehouse_sync, list(args), catch_exceptions=False)
def test_failed_symbols_exit_1(self, patched_warehouse):
"""有标的失败 → exit 1(旧码:failed 只进 summary,命令仍 exit 0)。"""
_FakeSyncer.result = {
"total": 2,
"ok": 1,
"added": 3,
"updated": 0,
"skipped": 0,
"failed": 1,
"details": [
{"symbol": "SH:600519", "added": 3, "updated": 0, "skipped": 0, "error": None},
{"symbol": "SZ:000001", "added": 0, "updated": 0, "skipped": 0, "error": "boom"},
],
}
result = self._invoke("--symbols", "SH:600519,SZ:000001", "--source", "tdx")
assert result.exit_code == 1, result.output
def test_all_ok_exit_0_and_source_in_summary(self, patched_warehouse):
"""全部成功 → exit 0summary JSON 带 source 标注(与 /bars 响应呼应)。"""
_FakeSyncer.result = {
"total": 1,
"ok": 1,
"added": 3,
"updated": 0,
"skipped": 0,
"failed": 0,
"details": [
{"symbol": "SH:600519", "added": 3, "updated": 0, "skipped": 0, "error": None},
],
}
result = self._invoke("--symbols", "SH:600519", "--source", "tdx")
assert result.exit_code == 0, result.output
payload = json.loads(result.stdout) # stdout 仅 summary JSON(进度/错误在 stderr
assert payload["source"] == "tdx"
assert payload["ok"] == 1
def test_malformed_symbol_no_bare_traceback(self, patched_warehouse):
"""缺冒号标的 → 友好 BadParameter(exit 2),不发网络请求不裸崩。"""
result = CliRunner().invoke(warehouse_sync, ["--symbols", "SH600519", "--source", "tdx"])
assert result.exit_code == 2
assert "市场:代码" in result.output
def test_period_choice_rejects_unknown(self, patched_warehouse):
result = CliRunner().invoke(warehouse_sync, ["--symbols", "SH:600519", "--period", "WEEKN"])
assert result.exit_code == 2
def test_baostock_rejects_intraday_period(self, patched_warehouse):
"""--source baostock + 分钟周期 → 参数层直接报错(旧码:静默空转 exit 0)。"""
for period in _PERIOD_CHOICES:
if period not in _BAOSTOCK_PERIODS:
result = CliRunner().invoke(
warehouse_sync,
["--symbols", "SH:600519", "--source", "baostock", "--period", period],
)
assert result.exit_code == 2, (period, result.output)
assert "baostock" in result.output
def test_baostock_accepts_daily(self, patched_warehouse, monkeypatch):
"""--source baostock + DAILY 正常放行(不走 TDX 客户端)。"""
_FakeSyncer.result = {
"total": 1,
"ok": 1,
"added": 0,
"updated": 0,
"skipped": 0,
"failed": 0,
"details": [],
}
# baostock 路径不经过 get_mac_client——若被调用说明走错分支
import easy_tdx.cli.conn as conn_mod
def _no_tdx():
raise AssertionError("baostock source 不应触碰 TDX 客户端")
monkeypatch.setattr(conn_mod, "get_mac_client", _no_tdx)
result = self._invoke("--symbols", "SH:600519", "--source", "baostock")
assert result.exit_code == 0, result.output
assert json.loads(result.stdout)["source"] == "baostock"
class TestWarehouseCheck:
def test_multiple_symbols_rejected_before_parse(self, monkeypatch):
"""多标的先报「只支持一个」,不再先 split 崩溃(旧码顺序颠倒)。"""
import easy_tdx.cli.cmd_warehouse as cw
monkeypatch.setattr(cw, "_require_warehouse", lambda db_path: _FakeWarehouse())
result = CliRunner().invoke(
warehouse_check,
["--symbols", "SH600519,SZ:000001"], # 旧码:含冒号绕过校验 → split 裸崩
catch_exceptions=False,
)
assert result.exit_code == 1
assert "只支持一个标的" in result.output
def test_single_symbol_missing_colon_rejected(self, monkeypatch):
import easy_tdx.cli.cmd_warehouse as cw
monkeypatch.setattr(cw, "_require_warehouse", lambda db_path: _FakeWarehouse())
result = CliRunner().invoke(
warehouse_check,
["--symbols", "SH600519"], # 旧码:split(":", 1) 裸 ValueError
catch_exceptions=False,
)
assert result.exit_code == 2
assert "市场:代码" in result.output
def test_single_valid_symbol_passes_market_code(self, monkeypatch):
import easy_tdx.cli.cmd_warehouse as cw
monkeypatch.setattr(cw, "_require_warehouse", lambda db_path: _FakeWarehouse())
result = CliRunner().invoke(
warehouse_check,
["--symbols", "SH:600519"],
catch_exceptions=False,
)
assert result.exit_code == 0, result.output
payload = json.loads(result.output)
assert payload["market"] == "SH" and payload["code"] == "600519"
def test_issues_still_exit_0_by_design(self, monkeypatch):
"""自检发现 issues → 正常输出并 exit 0(自检结果本身是正常输出,保持原口径)。"""
import easy_tdx.cli.cmd_warehouse as cw
class _Wh(_FakeWarehouse):
def health_check(self, market=None, code=None):
return {"issues": ["gap"]}
monkeypatch.setattr(cw, "_require_warehouse", lambda db_path: _Wh())
result = CliRunner().invoke(
warehouse_check, ["--symbols", "SH:600519"], catch_exceptions=False
)
assert result.exit_code == 0
assert "gap" in result.output
class TestFormulaScreenSymbolValidation:
def test_malformed_symbol_fails_fast(self):
"""缺冒号标的前置报错(旧码:循环里裸 ValueError traceback)。"""
from easy_tdx.cli.cmd_formula import formula_screen
result = CliRunner().invoke(
formula_screen,
["--symbols", "SH:600519,SH600036", "--formula", "金叉: CROSS(MA(C,5), MA(C,20));"],
catch_exceptions=False,
)
assert result.exit_code == 2
assert "市场:代码" in result.output
assert "SH600036" in result.output # 报错指出坏标的
+78
View File
@@ -212,3 +212,81 @@ def test_compiled_formula_is_dataclass_safe():
assert np.allclose( assert np.allclose(
f.compute(_df(20)).columns[""], f2.compute(_df(20)).columns[""], equal_nan=True f.compute(_df(20)).columns[""], f2.compute(_df(20)).columns[""], equal_nan=True
) )
# ── 回归:归类收严 / FILTER 副作用 / 递归上限 / REF 负移位(审查修复) ────────
def test_ratio_output_classified_as_value_not_signal():
"""0~1 区间的价格比率是数值列,不是信号列(旧码按 [0,1] 值域兜底误判)。"""
res = compile_formula("比率: C / HHV(C, 20);").compute(_df(40))
assert res.signals == []
assert res.values == ["比率"]
def test_normalized_oscillator_classified_as_value():
"""归一化振荡器(RSI/100)是数值列,不是信号列。"""
res = compile_formula("强度: RSI(C, 14) / 100;").compute(_df(40))
assert res.signals == []
assert res.values == ["强度"]
def test_binary_zero_one_values_still_signal():
"""真正的 0/1 两值输出仍兜底归信号列。"""
res = compile_formula("X: IF(C > MA(C, 5), 1, 0);").compute(_df(40))
assert res.signals == ["X"]
def test_filter_does_not_pollute_series():
"""FILTER(C, N) 不改写输入序列:同公式后续 MA(C, 2) 与未过滤一致。"""
df = _df(30)
res = compile_formula("A: FILTER(C, 2); B: MA(C, 2);").compute(df)
close = pd.to_numeric(df["close"]).to_numpy(dtype=float)
expected = pd.Series(close).rolling(2).mean().to_numpy()
np.testing.assert_allclose(res.columns["B"], expected, equal_nan=True)
def test_deep_paren_nesting_formula_error():
"""超深括号嵌套抛 FormulaError(嵌套过深),而非 RecursionError 逃逸。"""
text = "X: " + "(" * 5000 + "C" + ")" * 5000 + ";"
with pytest.raises(FormulaError, match="嵌套过深"):
compile_formula(text)
def test_deep_unary_chain_formula_error():
"""超长一元运算符链同样受深度上限保护。"""
with pytest.raises(FormulaError, match="嵌套过深"):
compile_formula("X: " + "!" * 5000 + "C;")
def test_moderate_nesting_still_compiles():
"""常规嵌套深度不受上限影响。"""
res = compile_formula("X: -(-(-(C + 1) * 2) + 3);").compute(_df(10))
assert res.columns["X"].shape == (10,)
def test_ref_negative_shift_banned():
"""REF 负移位(未来函数)显式 FormulaError,不再依赖 float 类型巧合。"""
with pytest.raises(FormulaError, match="负移位"):
compile_formula("X: REF(C, -1);").compute(_df(30))
def test_ref_negative_via_expression_banned():
"""负移位经表达式算出(如 0-1)同样被禁。"""
with pytest.raises(FormulaError, match="负移位"):
compile_formula("X: REF(C, 0 - 1);").compute(_df(30))
def test_ref_positive_still_works():
res = compile_formula("X: REF(C, 1);").compute(_df(30))
assert np.isnan(res.columns["X"][0])
assert res.columns["X"][1] == pytest.approx(float(pd.to_numeric(_df(30)["close"]).iloc[0]))
def test_mytt_internal_negative_ref_unaffected():
"""MyTT 库内直调(ICHIMOKU 迟行带)不经公式白名单,负移位仍可用。"""
from easy_tdx.MyTT import REF
close = np.arange(5, dtype=float)
out = REF(close, -1)
assert out[0] == pytest.approx(1.0)
+10
View File
@@ -191,3 +191,13 @@ def test_rest_formula_screen_async_task():
# symbols 路径需要行情连接——离线环境预期 400/500(无 mock client # symbols 路径需要行情连接——离线环境预期 400/500(无 mock client
# 这里只验证请求校验(symbols 非空)不炸 # 这里只验证请求校验(symbols 非空)不炸
assert r.status_code in (400, 500, 202) assert r.status_code in (400, 500, 202)
def test_pick_signal_columns_ignores_ratio_value_column():
"""0~1 值域的比率列归类为数值输出后,不再被自动挑成买卖信号列。"""
_, result = attach_formula_columns(
_df(60), compile_formula("比率: C / HHV(C, 20);\n强弱: C > MA(C, 5);")
)
buy, sell = pick_signal_columns(result)
assert buy == "强弱" # 旧码 signals 含「比率」且排在首位,被误选为买入列
assert sell is None
+38
View File
@@ -205,6 +205,44 @@ def test_combined_metrics_dd_duration_unclosed_counts_to_end():
assert m.max_dd_duration == 20 assert m.max_dd_duration == 20
def test_combined_metrics_missing_drawdown_pct_row_continues_state():
"""缺行 drawdown_pct 视为状态延续,不得当作创新高截断水下期。
回归compute_combined_metrics drawdown_pct 分支曾把缺行/None 0
=创新高把真实水下段从中间截断max_dd_duration 被低估3 2
"""
eq = [
{"total": 100.0, "datetime": "2024-01-01", "drawdown_pct": 0.0},
{"total": 50.0, "datetime": "2024-01-02", "drawdown_pct": 0.5},
{"total": 50.0, "datetime": "2024-01-03"}, # 缺 drawdown_pct → 沿用水下
{"total": 50.0, "datetime": "2024-01-04", "drawdown_pct": 0.5},
]
m = compute_combined_metrics(eq)
# idx0 之后一直未创新高(缺行延续 idx1 的水下状态)→ 计到末点 = 3
assert m.max_dd_duration == 3
def test_combined_metrics_nan_drawdown_pct_continues_state():
"""NaN drawdown_pct 沿用上一根状态:峰值后的 NaN 仍按峰值处理。
回归NaN 行既不算峰值也不截断last_peak 停在上一根真峰值导致
水下期被多算一根2 1
"""
eq = [
{"total": 100.0, "datetime": "2024-01-01", "drawdown_pct": 0.0},
{"total": 110.0, "datetime": "2024-01-02", "drawdown_pct": 0.0},
{"total": 110.0, "datetime": "2024-01-03", "drawdown_pct": float("nan")},
{
"total": 109.0,
"datetime": "2024-01-04",
"drawdown_pct": (110.0 - 109.0) / 110.0,
},
]
m = compute_combined_metrics(eq)
# NaN 行沿用 idx1 的峰值状态 → 最后一次创新高为 idx2 → 水下 1 根
assert m.max_dd_duration == 1
def test_combined_metrics_insufficient_points(): def test_combined_metrics_insufficient_points():
m = compute_combined_metrics([{"total": 100.0}]) m = compute_combined_metrics([{"total": 100.0}])
assert m.n_points == 1 assert m.n_points == 1
+100
View File
@@ -128,6 +128,67 @@ def test_limitup_empty_vipdoc(tmp_path):
assert eco.summary()["limit_up_count"] == 0 assert eco.summary()["limit_up_count"] == 0
# ── 涨跌停价舍入(回归:浮点 floor(x*100+0.5) 在半分边界错 1 分)──────────────
def test_limit_price_matches_exchange_rounding_all_range():
"""_limit_price 与交易所 ROUND_HALF_UP 对 1.00~600.00 全价位零差异。
旧实现float 乘后 floor ±10% 67/318 个价位±5% 90/884
价位算低 1 prev=1.15涨停价应 1.27旧算 1.26
"""
from decimal import ROUND_HALF_UP, Decimal
from easy_tdx.screen.limitup import _limit_price
for pct in (10, 5, 20, -10, -5, -20):
for cents in range(100, 60001):
prev = Decimal(cents).scaleb(-2)
expected = (
int(
(Decimal(cents) * (100 + pct) / 100).quantize(
Decimal("1"), rounding=ROUND_HALF_UP
)
)
/ 100
)
got = _limit_price(float(prev), pct)
assert got == expected, (prev, pct, got, expected)
def test_exchange_boundary_prices_detected(tmp_path):
"""半分边界价位的真实涨跌停不因浮点舍入漏判。
选点依据.day 读回raw×0.01的浮点误差会抵消部分边界33.0536.36
2.652.39 是经读回仿真验证后旧实现floor 浮点版仍漏判的价位
"""
from easy_tdx.screen.limitup import compute_limitup_ecology
# prev=33.05 → 交易所涨停价 36.36(旧实现误算 36.35 → 漏判涨停)
_write_stock(tmp_path, "sh", "600901", [33.05, 36.36])
# prev=2.65 → 交易所跌停价 2.39(旧实现误算 2.38 → 漏判跌停)
_write_stock(tmp_path, "sz", "000902", [2.65, 2.39])
eco = compute_limitup_ecology(tmp_path)
up = {e.code: e for e in eco.limit_up}
down = {e.code: e for e in eco.limit_down}
assert "600901" in up, f"33.05→36.36 应判涨停,实际 limit_up={up}"
assert up["600901"].streak == 1
assert "000902" in down, f"2.65→2.39 应判跌停,实际 limit_down={down}"
def test_history_boundary_prices_counted(tmp_path):
"""历史回补同样按交易所口径计涨跌停(33.05→36.36 / 2.65→2.39)。"""
from easy_tdx.screen.limitup import compute_limitup_history
_write_stock(tmp_path, "sh", "600901", [33.05, 36.36])
_write_stock(tmp_path, "sz", "000902", [2.65, 2.39])
rows = {r["date"]: r for r in compute_limitup_history(tmp_path, days=5)}
assert rows[20260802]["limit_up"] == 1
assert rows[20260802]["limit_down"] == 1
def test_limitup_endpoint_and_cache(vipdoc, monkeypatch): def test_limitup_endpoint_and_cache(vipdoc, monkeypatch):
"""端点返回 DictResponse 包装;60s 内命中缓存(扫描只跑一次)。""" """端点返回 DictResponse 包装;60s 内命中缓存(扫描只跑一次)。"""
pytest.importorskip("fastapi") pytest.importorskip("fastapi")
@@ -164,3 +225,42 @@ def test_limitup_endpoint_and_cache(vipdoc, monkeypatch):
assert r2.json()["data"] == d1 assert r2.json()["data"] == d1
assert calls["n"] == 1 # 第二次命中缓存 assert calls["n"] == 1 # 第二次命中缓存
def test_limitup_endpoint_cache_key_includes_vipdoc(vipdoc, tmp_path, monkeypatch):
"""缓存键须含 vipdoc:不同 vipdoc 的请求在 TTL 内不互相命中。
旧实现 _limitup_cache 是单值缓存先到的 vipdoc=A 结果会被 vipdoc=B
的请求在 60s TTL 内复用
"""
pytest.importorskip("fastapi")
from fastapi import FastAPI
from fastapi.testclient import TestClient
from easy_tdx.web.errors import register_exception_handlers
from easy_tdx.web.routers import market as market_mod
other = tmp_path / "vipdoc_other"
(other / "sh" / "lday").mkdir(parents=True)
(other / "sh" / "lday" / "sh600100.day").write_bytes(
_day(20260801, 9.95, 11.0, 9.9, 11.0) + _day(20260802, 11.0, 12.1, 10.9, 12.1)
)
app = FastAPI()
register_exception_handlers(app)
app.include_router(market_mod.router, prefix="/api/v1")
app.state.tdx_client = object()
with TestClient(app) as client:
r1 = client.get("/api/v1/limitup-ecology", params={"vipdoc": str(vipdoc)})
assert r1.status_code == 200
r2 = client.get("/api/v1/limitup-ecology", params={"vipdoc": str(other)})
assert r2.status_code == 200
# 同 vipdoc 的第二次请求才命中缓存;不同 vipdoc 必须各自扫描
with TestClient(app) as client:
client.get("/api/v1/limitup-ecology", params={"vipdoc": str(vipdoc)})
client.get("/api/v1/limitup-ecology", params={"vipdoc": str(vipdoc)})
d = client.get("/api/v1/limitup-ecology", params={"vipdoc": str(other)}).json()["data"]
# other 目录只有 600100 一只 2 连板,不含 vipdoc 目录的 3 连板数据
assert d["summary"]["limit_up_count"] == 1
+26
View File
@@ -13,6 +13,7 @@ from __future__ import annotations
import numpy as np import numpy as np
import pandas as pd import pandas as pd
import pytest
from easy_tdx.backtest.multi_strategy_engine import ( from easy_tdx.backtest.multi_strategy_engine import (
MultiStrategyEngine, MultiStrategyEngine,
@@ -206,6 +207,31 @@ class TestMultiStrategyEngine:
# 合并曲线长度应至少覆盖两个范围的最晚结束日(并集) # 合并曲线长度应至少覆盖两个范围的最晚结束日(并集)
assert len(result.combined_equity) >= 60 assert len(result.combined_equity) >= 60
def test_total_return_capital_weighted_with_disjoint_dates(self) -> None:
"""晚起步槽位建仓前应按初始资金趴账(合并曲线首值=总投入资金)。
回归旧实现对日期并集的前导缺口填 0B 槽位起步前贡献 0 而非其
分得的 50 合并曲线首值 = 50 < 总资金 100 total_return 被虚增
正确口径前导缺口用每列首个有效值=初始资金回填bfill
"""
df_a = _make_df(60, seed=1, start="2024-01-01")
df_b = _make_df(60, seed=2, start="2024-03-01")
slots = [
StrategySlot("A", "SH:601088", SimpleBuyStrategy(), df_a),
StrategySlot("B", "SZ:000001", SimpleBuyStrategy(), df_b),
]
result = MultiStrategyEngine(slots, total_cash=1_000_000).run()
# 合并曲线首值 = 总投入资金(旧实现 = 500000,缺晚起步槽位的资金)
assert result.combined_equity["total"].iloc[0] == pytest.approx(1_000_000.0)
# total_return == 各槽位资金加权真实收益
weighted = sum(
0.5 * res.performance.get("total_return", 0.0)
for res in result.individual_results.values()
)
assert result.total_performance["total_return"] == pytest.approx(weighted, abs=1e-9)
def test_empty_strategies_returns_empty_result(self) -> None: def test_empty_strategies_returns_empty_result(self) -> None:
"""空策略列表应返回空结果,不抛异常。""" """空策略列表应返回空结果,不抛异常。"""
engine = MultiStrategyEngine([], total_cash=1_000_000) engine = MultiStrategyEngine([], total_cash=1_000_000)
+41
View File
@@ -557,3 +557,44 @@ class TestBBPandBBW:
quiet = 100 + rng.standard_normal(60) * 0.1 quiet = 100 + rng.standard_normal(60) * 0.1
wild = 100 + rng.standard_normal(60) * 5.0 wild = 100 + rng.standard_normal(60) * 5.0
assert MyTT.BBW(wild)[-1] > MyTT.BBW(quiet)[-1] assert MyTT.BBW(wild)[-1] > MyTT.BBW(quiet)[-1]
class TestFilter:
"""FILTER 无副作用(审查修复:曾原地改写输入序列)。"""
def test_filter_does_not_mutate_input(self):
x = np.array([1.0, 0.0, 1.0, 1.0, 1.0, 0.0, 1.0])
snapshot = x.copy()
out = MyTT.FILTER(x, 2)
np.testing.assert_array_equal(x, snapshot) # 旧码把 x 原地置零,失败
# x[0]=1 触发 → 后 2 根置零;x[3]=1 触发 → 后 2 根置零;x[6]=1
np.testing.assert_array_equal(out, [1.0, 0.0, 0.0, 1.0, 0.0, 0.0, 1.0])
def test_filter_accepts_bool_input(self):
s = np.array([True, False, True, True])
out = MyTT.FILTER(s, 1)
np.testing.assert_array_equal(s, [True, False, True, True])
assert out.dtype == bool
# s[0]=True 触发 → s[1] 置零;s[2]=True 触发 → s[3] 置零(s[2] 本身保留)
np.testing.assert_array_equal(out, [True, False, True, False])
def test_filter_formula_series_not_polluted(self):
"""公式场景:FILTER(C, N) 后 C 仍是原收盘序列。"""
from easy_tdx.formula import compile_formula
n = 30
close = 10.0 * np.linspace(1.0, 2.0, n)
df = pd.DataFrame(
{
"datetime": pd.date_range("2024-01-01", periods=n),
"open": close,
"high": close,
"low": close,
"close": close,
"vol": np.ones(n),
}
)
res = compile_formula("A: FILTER(C, 2); B: MA(C, 2);").compute(df)
np.testing.assert_allclose(
res.columns["B"], pd.Series(close).rolling(2).mean().to_numpy(), equal_nan=True
)
+45
View File
@@ -290,6 +290,51 @@ class TestPortfolioFullMetrics:
) )
assert result.total_performance["total_return"] == pytest.approx(weighted, abs=1e-9) assert result.total_performance["total_return"] == pytest.approx(weighted, abs=1e-9)
def test_total_return_capital_weighted_with_uneven_start_dates(self) -> None:
"""晚上市标的建仓前应按初始资金趴账(合并曲线首值=总投入资金)。
回归旧实现 ``_build_combined_equity`` 对日期并集的前导缺口填 0
晚上市标的上市前贡献 0 而非其分得的初始资金合并曲线首值 < 总投入
total_return 被系统性虚增100 + 60 根起步的组合实测虚增约 10
正确口径与组合 Walk-Forward ffill().bfill() 一致前导缺口用
每列首个有效值回填资金在组合起点即已分配建仓前趴账
"""
n = 100
close = np.linspace(10.0, 12.0, n)
dates = pd.bdate_range("2024-01-02", periods=n)
def _mk(cnt: int) -> pd.DataFrame:
c = close[-cnt:]
return pd.DataFrame(
{
"datetime": dates[-cnt:],
"open": c,
"high": c,
"low": c,
"close": c,
"vol": 1e6,
"amount": c * 1e6,
}
)
stocks = [
StockData("000001", "SZ", _mk(n)), # 全程 100 根
StockData("600000", "SH", _mk(40)), # 同涨势、晚 60 根起步
]
result = PortfolioBacktestEngine(
strategy=SimpleBuyStrategy, stocks=stocks, total_cash=200000
).run()
# 1) 合并曲线首值 = 总投入资金(旧实现 = 100000,缺晚上市标的的资金)
assert result.combined_equity["total"].iloc[0] == pytest.approx(200000.0)
# 2) total_return == 各标的资金加权真实收益
weighted = sum(
0.5 * res.performance.get("total_return", 0.0)
for res in result.individual_results.values()
)
assert result.total_performance["total_return"] == pytest.approx(weighted, abs=1e-9)
def test_to_dict_contains_trades(self) -> None: def test_to_dict_contains_trades(self) -> None:
"""to_dict 应包含组合层成交表(REST/AI 解读消费)。""" """to_dict 应包含组合层成交表(REST/AI 解读消费)。"""
stocks = [StockData("000001", "SZ", _make_df(100, seed=42))] stocks = [StockData("000001", "SZ", _make_df(100, seed=42))]
+27
View File
@@ -182,3 +182,30 @@ def test_multi_strategy_wf_empty_slots() -> None:
wf = MultiStrategyWalkForwardEngine(strategies=[], n_windows=3).run() wf = MultiStrategyWalkForwardEngine(strategies=[], n_windows=3).run()
assert wf.windows == [] assert wf.windows == []
def test_combo_slot_failure_logs_warning(caplog):
"""单槽位回测异常记 warning(含槽位标识),不拖垮整窗(旧码静默 continue)。"""
import logging
from easy_tdx.backtest.multi_strategy_engine import StrategySlot
from easy_tdx.backtest.walkforward import MultiStrategyWalkForwardEngine
class Boom(Strategy):
def init(self) -> None:
raise RuntimeError("slot-boom")
def next(self) -> None:
pass
slots = [
StrategySlot(
label="正常", symbol="SH:601088", strategy=PeriodicStrategy(), df=_make_df(400, seed=42)
),
StrategySlot(label="炸裂", symbol="SZ:000001", strategy=Boom(), df=_make_df(400, seed=99)),
]
with caplog.at_level(logging.WARNING, logger="easy_tdx.backtest.walkforward"):
wf = MultiStrategyWalkForwardEngine(strategies=slots, n_windows=3).run()
assert len(wf.windows) == 3 # 正常槽位照常出窗
assert wf.total_trades > 0
assert any("炸裂" in r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING)
+7
View File
@@ -116,6 +116,13 @@ def test_detect_gap_uses_chinext_threshold() -> None:
assert detect_ex_dividend_gaps(df2, "300750") == ["2010-01-06"] assert detect_ex_dividend_gaps(df2, "300750") == ["2010-01-06"]
def test_detect_gap_ignores_nonfinite_ratio() -> None:
"""前收缺失(NaN)导致的非有限比率不计为除权跳空(首根前收缺失常见)。"""
df = _kline([float("nan")] + [10.0] * 5, opens=[5.0] + [10.0] * 5)
# 旧实现把 ratio=NaN 也当跳空 → 误报 ["2010-01-02"]
assert detect_ex_dividend_gaps(df, "600000") == []
# --------------------------------------------------------------------------- # # --------------------------------------------------------------------------- #
# 已知案例回归(合成) # 已知案例回归(合成)
# --------------------------------------------------------------------------- # # --------------------------------------------------------------------------- #
+20
View File
@@ -332,3 +332,23 @@ class TestStopFlag:
await asyncio.wait_for(feed._run_sync_loop(client, None), timeout=2.0) await asyncio.wait_for(feed._run_sync_loop(client, None), timeout=2.0)
assert client.calls == [] assert client.calls == []
async def test_restart_after_stop_runs_again(self) -> None:
"""start→stop→start:停止请求一次性消费,实例可再次启动。
回归旧实现 _stop_requested 只置位不复位stop 后再次 run_async
会静默立即返回假启动同一实例永久失效
"""
bus = EventBus()
client = AsyncMockClient([_sample_quotes_df()])
feed = RealtimeDataFeed(bus=bus, symbols=[(0, "000001")], sessions=(), interval=0.1)
await feed.run_async(client, max_iterations=1)
assert len(client.calls) == 1
feed.stop()
await feed.run_async(client, max_iterations=1) # 消费停止请求:启动即退出
assert len(client.calls) == 1
await feed.run_async(client, max_iterations=1) # 再次启动应正常轮询
assert len(client.calls) == 2
+180 -4
View File
@@ -7,6 +7,7 @@ sentiment_store 用 EASY_TDX_CONFIG_DIR 指向临时目录;limitup 历史复
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import pathlib
import pytest import pytest
@@ -110,7 +111,11 @@ def test_sampler_inserts_store_rows(store):
@pytest.fixture @pytest.fixture
def vipdoc_factory(tmp_path): def vipdoc_factory(tmp_path):
"""{文件名: {dates, closes}} 合成 vipdoc 目录的工厂。""" """{文件名: {dates, closes}} 合成 vipdoc 目录的工厂。
``factory(specs, root=None)``root 缺省写 tmp_path同一测试需要多个
独立 vipdoc 目录时传不同 root
"""
from easy_tdx.offline.daily_bar import _DAILY_FMT from easy_tdx.offline.daily_bar import _DAILY_FMT
def _day(date: int, close: float) -> bytes: def _day(date: int, close: float) -> bytes:
@@ -125,14 +130,15 @@ def vipdoc_factory(tmp_path):
0, 0,
) )
def factory(specs: dict[str, dict]) -> object: def factory(specs: dict[str, dict], root=None) -> object:
base = pathlib.Path(root) if root is not None else tmp_path
for filename, spec in specs.items(): for filename, spec in specs.items():
exchange = filename[:2] exchange = filename[:2]
lday = tmp_path / exchange / "lday" lday = base / exchange / "lday"
lday.mkdir(parents=True, exist_ok=True) lday.mkdir(parents=True, exist_ok=True)
data = b"".join(_day(d, c) for d, c in zip(spec["dates"], spec["closes"])) data = b"".join(_day(d, c) for d, c in zip(spec["dates"], spec["closes"]))
(lday / f"{filename}.day").write_bytes(data) (lday / f"{filename}.day").write_bytes(data)
return tmp_path return base
return factory return factory
@@ -198,3 +204,173 @@ def test_limitup_history_endpoint_cache(vipdoc_factory, monkeypatch):
assert body["days"][0] == {"date": 20260802, "limit_up": 1, "limit_down": 0} assert body["days"][0] == {"date": 20260802, "limit_up": 1, "limit_down": 0}
client.get("/api/v1/market/limitup-history", params={"days": 10, "vipdoc": str(v)}) client.get("/api/v1/market/limitup-history", params={"days": 10, "vipdoc": str(v)})
assert calls["n"] == 1 # 缓存命中 assert calls["n"] == 1 # 缓存命中
def test_board_fund_history_null_main_net_no_type_error(tmp_path):
"""历史遗留的 main_net NULL 行不得让 /market/board-fund/history 抛 TypeError。
正式 schema main_net REAL NOT NULLNaN 会被 SQLite 存成 NULL
NOT NULL 拒绝但手工编辑/旧版本库可能存在 NULL 读侧须兜底
旧实现float(None) TypeError
"""
import sqlite3
db = tmp_path / "legacy_sentiment.db"
conn = sqlite3.connect(db)
conn.executescript(
"""
CREATE TABLE samples (date INTEGER NOT NULL, minute INTEGER NOT NULL, ts INTEGER NOT NULL,
up_count INTEGER NOT NULL, down_count INTEGER NOT NULL, neutral_count INTEGER NOT NULL,
total_count INTEGER NOT NULL, limit_up_count INTEGER NOT NULL,
limit_down_count INTEGER NOT NULL, total_amount REAL NOT NULL,
PRIMARY KEY (date, minute));
CREATE TABLE board_fund (date INTEGER NOT NULL, rank INTEGER NOT NULL,
code TEXT NOT NULL, name TEXT NOT NULL, main_net REAL,
PRIMARY KEY (date, rank));
"""
)
conn.execute(
"INSERT INTO board_fund (date, rank, code, name, main_net) VALUES (?,?,?,?,?)",
(20260901, 1, "881001", "银行", None),
)
conn.commit()
conn.close()
from easy_tdx.web.sentiment_store import SentimentStore
store = SentimentStore(db_path=db)
assert store.list_fund_days(5) == [
{
"date": 20260901,
"boards": [{"rank": 1, "code": "881001", "name": "银行", "main_net": 0.0}],
}
]
def test_upsert_fund_day_nan_main_net_stored_as_zero(store):
"""写入口径:NaN 主力净流入落库为 0.0REAL NOT NULL 列不吃 NaN)。"""
store.upsert_fund_day(
20260902,
[{"code": "881001", "name": "银行", "main_net": float("nan")}],
)
days = store.list_fund_days(5)
assert days[0]["boards"][0]["main_net"] == 0.0
def test_limitup_history_cache_key_includes_vipdoc(vipdoc_factory, monkeypatch, tmp_path):
"""limitup-history 缓存键须含 (days, vipdoc),不同 vipdoc 不互相命中。"""
pytest.importorskip("fastapi")
from fastapi import FastAPI
from fastapi.testclient import TestClient
from easy_tdx.screen import limitup as limitup_mod
from easy_tdx.web.errors import register_exception_handlers
from easy_tdx.web.routers import market as market_mod
v1 = vipdoc_factory({"sh600100": {"dates": [20260801, 20260802], "closes": [10.0, 11.0]}})
v2_dir = tmp_path / "vipdoc_v2"
(v2_dir / "sh" / "lday").mkdir(parents=True)
v2 = vipdoc_factory(
{
"sh600100": {"dates": [20260801, 20260802], "closes": [10.0, 11.0]},
"sz000200": {"dates": [20260801, 20260802], "closes": [10.0, 9.0]},
},
root=v2_dir,
)
calls = {"n": 0}
real = limitup_mod.compute_limitup_history
def counting(*a, **kw):
calls["n"] += 1
return real(*a, **kw)
monkeypatch.setattr(limitup_mod, "compute_limitup_history", counting)
app = FastAPI()
register_exception_handlers(app)
app.include_router(market_mod.router, prefix="/api/v1")
app.state.tdx_client = object()
with TestClient(app) as client:
r1 = client.get("/api/v1/market/limitup-history", params={"days": 10, "vipdoc": str(v1)})
assert r1.status_code == 200
r2 = client.get("/api/v1/market/limitup-history", params={"days": 10, "vipdoc": str(v2)})
assert r2.status_code == 200
# v2 多一只跌停股 → 结果必须不同(不允许命中 v1 的缓存)
assert r2.json()["data"]["days"][0]["limit_down"] == 1
assert calls["n"] == 2
# ── 采样器时区统一(v1.32.6):日期/分钟键一律取沪市时区 ─────────────────────
class _RecorderDatetime:
"""记录 now(tz) 实参的 datetime 替身。"""
captured: dict = {}
@classmethod
def now(cls, tz=None):
cls.captured["tz"] = tz
from datetime import datetime as _dt
return _dt(2026, 9, 6, 2, 30, tzinfo=tz) if tz else _dt(2026, 9, 6, 2, 30)
def test_sentiment_sampler_uses_shanghai_tz(store, monkeypatch):
"""SentimentSampler._sample_once 的 (date, minute) 键必须取沪市时区。"""
import asyncio
import pandas as pd
from easy_tdx.realtime.session import SHANGHAI_TZ
from easy_tdx.web import sentiment_sampler as ss_mod
async def fake_stat():
return pd.DataFrame(
[
{
"up_count": 2000,
"down_count": 2000,
"neutral_count": 100,
"total_count": 4100,
"limit_up_count": 50,
"limit_down_count": 10,
"total_amount": 8e11,
}
]
)
monkeypatch.setattr(ss_mod, "datetime", _RecorderDatetime)
sampler = ss_mod.SentimentSampler(fake_stat, store=store)
asyncio.run(sampler._sample_once())
assert _RecorderDatetime.captured["tz"] is SHANGHAI_TZ
rows = store.day_samples(20260906)
assert len(rows) == 1 and rows[0]["minute"] == 230
def test_fund_flow_sampler_uses_shanghai_tz(store, monkeypatch):
"""FundFlowSampler._sample_once 的采样日期必须取沪市时区。"""
import asyncio
import pandas as pd
from easy_tdx.realtime.session import SHANGHAI_TZ
from easy_tdx.web import sentiment_sampler as ss_mod
class _FakeMac:
async def get_board_ranking(self, **kw):
return pd.DataFrame(
[
{"code": "881001", "name": "银行", "main_net_amount": 1.2e9},
]
)
monkeypatch.setattr(ss_mod, "datetime", _RecorderDatetime)
sampler = ss_mod.FundFlowSampler(_FakeMac(), store=store)
asyncio.run(sampler._sample_once())
assert _RecorderDatetime.captured["tz"] is SHANGHAI_TZ
days = store.list_fund_days(5)
assert days[0]["date"] == 20260906
+75
View File
@@ -257,3 +257,78 @@ def test_task_list_includes_persisted_history_after_new_app(persisted_env):
client = _client() client = _client()
tasks = client.get("/api/v1/backtest/tasks?limit=50").json()["tasks"] tasks = client.get("/api/v1/backtest/tasks?limit=50").json()["tasks"]
assert any(t["task_id"] == task_id for t in tasks) assert any(t["task_id"] == task_id for t in tasks)
# ── v1.32.6 修复:列表查询懒加载 result_json + 淘汰只移终态 ────────────────────
def test_list_recent_without_results_skips_result_json(persisted_env):
"""include_results=False 时不 SELECT/解析 result_json(列表页瘦身)。
旧签名无该参数列表页会把每条任务的完整结果 JSON 拖进内存解析
"""
import sqlite3
from easy_tdx.web.task_store import TaskStore
store = TaskStore()
store.save(
task_id="big",
status="done",
created_at=2.0,
result={
"performance": {"total_return": 0.25},
"equity_curve": [{"i": i} for i in range(500)],
},
)
# 把 result_json 打坏:include_results=False 路径根本不读它 → 不受影响
with sqlite3.connect(store.path) as conn:
conn.execute("UPDATE backtest_tasks SET result_json = '{not-json' WHERE task_id='big'")
conn.commit()
rows = store.list_recent(limit=10, include_results=False)
assert len(rows) == 1
assert rows[0]["task_id"] == "big"
assert rows[0]["status"] == "done"
assert rows[0]["result"] is None
# 详情(load)仍取全量并走损坏降级路径
d = store.load("big")
assert d is not None and d["result"] is None
def test_eviction_only_removes_terminal_states():
"""超限淘汰只移 done/failedpending/running 不淘汰(消灭 pending 幽灵)。
旧实现"淘汰最旧 non-running"会把最早提交尚未起跑的 pending 条目淘汰
worker 随后取不到状态直接跳过 磁盘遗留 pending 行被恢复成永久
pending 的幽灵任务
"""
from easy_tdx.web.task_runner import BacktestTaskRunner, TaskState
runner = BacktestTaskRunner(max_workers=1, max_results=2)
with runner._lock:
# 插入顺序即 LRU 序:pending 最旧、done 最新
runner._tasks["p1"] = TaskState(task_id="p1", status="pending")
runner._tasks["r1"] = TaskState(task_id="r1", status="running")
runner._tasks["d1"] = TaskState(task_id="d1", status="done")
runner._evict_if_needed_locked()
assert "p1" in runner._tasks, "pending 不得被淘汰"
assert "r1" in runner._tasks, "running 不得被淘汰"
assert "d1" not in runner._tasks, "超限时应淘汰最旧的终态条目"
def test_eviction_pending_still_evicted_by_old_logic_regression_guard():
"""回归对照:若恢复旧逻辑(淘汰首个 non-running),pending 会先被选中。
本测试钉死新语义全部条目为 pending 时宁可不淘汰超限跳过
"""
from easy_tdx.web.task_runner import BacktestTaskRunner, TaskState
runner = BacktestTaskRunner(max_workers=1, max_results=1)
with runner._lock:
runner._tasks["p1"] = TaskState(task_id="p1", status="pending")
runner._tasks["p2"] = TaskState(task_id="p2", status="pending")
runner._evict_if_needed_locked()
assert len(runner._tasks) == 2 # 无终态可淘汰 → 跳过,不丢任务
+21 -14
View File
@@ -17,6 +17,8 @@
import struct import struct
from datetime import time from datetime import time
import pytest
from easy_tdx import UNUSUAL_TYPE_NAMES from easy_tdx import UNUSUAL_TYPE_NAMES
from easy_tdx.mac.commands.unusual import UnusualCmd, _describe_unusual from easy_tdx.mac.commands.unusual import UnusualCmd, _describe_unusual
@@ -48,15 +50,20 @@ class TestDescribeUnusualKnownTypes:
def test_type_0x04(self): def test_type_0x04(self):
# 真实样本:605365 立达信 2026-09-01 09:35:11 # 真实样本:605365 立达信 2026-09-01 09:35:11
desc, val = _describe_unusual(0x04, bytes.fromhex("00b8d73d3d0000000000000000")) desc, val = _describe_unusual(0x04, bytes.fromhex("00b8d73d3d0000000000000000"), 9)
assert desc == "加速拉升" assert desc == "加速拉升"
assert val == "4.63%" assert val == "4.63%"
def test_unknown_type_fallback(self): def test_unknown_type_fallback(self):
desc, val = _describe_unusual(0x42, bytes.fromhex("00" * 13)) desc, val = _describe_unusual(0x42, bytes.fromhex("00" * 13), 9)
assert desc == "异动类型0x42" assert desc == "异动类型0x42"
assert val == "" assert val == ""
def test_hour_is_required(self):
"""hour 必传:缺省 9 会让 15:00 的 0x15 被误标「竞价」(回归防护)。"""
with pytest.raises(TypeError):
_describe_unusual(0x15, bytes.fromhex("00" * 13))
class TestType0x15: class TestType0x15:
"""0x15 竞价/尾盘异动(Issue #62)。 """0x15 竞价/尾盘异动(Issue #62)。
@@ -67,19 +74,19 @@ class TestType0x15:
def test_auction_drop(self): def test_auction_drop(self):
# 真实样本:600551 时代出版 09:25:00v1=0x03 竞价下跌 # 真实样本:600551 时代出版 09:25:00v1=0x03 竞价下跌
desc, val = _describe_unusual(0x15, bytes.fromhex("030c9846bc003e1d4700000000")) desc, val = _describe_unusual(0x15, bytes.fromhex("030c9846bc003e1d4700000000"), 9)
assert desc == "竞价下跌" assert desc == "竞价下跌"
assert val == "-1.21%/40254手" assert val == "-1.21%/40254手"
def test_auction_rise(self): def test_auction_rise(self):
# 真实样本:600127 金健米业 09:25:00v1=0x02 竞价拉升(尾段自 10.84 冲至 12.05 # 真实样本:600127 金健米业 09:25:00v1=0x02 竞价拉升(尾段自 10.84 冲至 12.05
desc, val = _describe_unusual(0x15, bytes.fromhex("0213d2cd3d00367b4700000000")) desc, val = _describe_unusual(0x15, bytes.fromhex("0213d2cd3d00367b4700000000"), 9)
assert desc == "竞价拉升" assert desc == "竞价拉升"
assert val == "10.05%/64310手" assert val == "10.05%/64310手"
def test_auction_flat(self): def test_auction_flat(self):
# 真实样本:600410 华胜天成 09:25:01v1=0x01 竞价平稳(尾段价格未动) # 真实样本:600410 华胜天成 09:25:01v1=0x01 竞价平稳(尾段价格未动)
desc, val = _describe_unusual(0x15, bytes.fromhex("01000000000098a54500000000")) desc, val = _describe_unusual(0x15, bytes.fromhex("01000000000098a54500000000"), 9)
assert desc == "竞价平稳" assert desc == "竞价平稳"
assert val == "0.00%/5299手" assert val == "0.00%/5299手"
@@ -96,7 +103,7 @@ class TestType0x15:
assert val == "-0.77%/171634手" assert val == "-0.77%/171634手"
def test_unknown_sub_type_falls_back(self): def test_unknown_sub_type_falls_back(self):
desc, _ = _describe_unusual(0x15, struct.pack("<B2fI", 0x77, 0.0, 100.0, 0)) desc, _ = _describe_unusual(0x15, struct.pack("<B2fI", 0x77, 0.0, 100.0, 0), 9)
assert desc == "竞价异动" assert desc == "竞价异动"
desc, _ = _describe_unusual(0x15, struct.pack("<B2fI", 0x77, 0.0, 100.0, 0), 15) desc, _ = _describe_unusual(0x15, struct.pack("<B2fI", 0x77, 0.0, 100.0, 0), 15)
assert desc == "尾盘异动" assert desc == "尾盘异动"
@@ -107,19 +114,19 @@ class TestType0x16:
def test_strong_at_auction(self): def test_strong_at_auction(self):
# 真实样本:600551 时代出版 09:25:00v2=+5.82% 与当日开盘涨幅精确一致 # 真实样本:600551 时代出版 09:25:00v2=+5.82% 与当日开盘涨幅精确一致
desc, val = _describe_unusual(0x16, bytes.fromhex("010f506e3dcb846e3d00000000")) desc, val = _describe_unusual(0x16, bytes.fromhex("010f506e3dcb846e3d00000000"), 9)
assert desc == "盘中强势" assert desc == "盘中强势"
assert val == "5.82%" assert val == "5.82%"
def test_weak_at_auction(self): def test_weak_at_auction(self):
# 真实样本:600683 京投发展 09:25:01v1=0xFF(弱势 1 级),v2=-6.40% # 真实样本:600683 京投发展 09:25:01v1=0xFF(弱势 1 级),v2=-6.40%
desc, val = _describe_unusual(0x16, bytes.fromhex("ffc71d83bd690383bd00000000")) desc, val = _describe_unusual(0x16, bytes.fromhex("ffc71d83bd690383bd00000000"), 9)
assert desc == "盘中弱势" assert desc == "盘中弱势"
assert val == "-6.40%" assert val == "-6.40%"
def test_new_stock_no_limit(self): def test_new_stock_no_limit(self):
# 真实样本:601123 N马矿 09:25:00,新股无涨跌幅限制,v2=+245.86% # 真实样本:601123 N马矿 09:25:00,新股无涨跌幅限制,v2=+245.86%
desc, val = _describe_unusual(0x16, bytes.fromhex("03775a1d404a5b1d4000000000")) desc, val = _describe_unusual(0x16, bytes.fromhex("03775a1d404a5b1d4000000000"), 9)
assert desc == "盘中强势" assert desc == "盘中强势"
assert val == "245.86%" assert val == "245.86%"
@@ -134,13 +141,13 @@ class TestType0x13:
def test_auction_test_buy(self): def test_auction_test_buy(self):
# 真实样本:603980 09:15:14,申报价 8.71 高于昨收 7.92(往上试) # 真实样本:603980 09:15:14,申报价 8.71 高于昨收 7.92(往上试)
desc, val = _describe_unusual(0x13, bytes.fromhex("00295c0b41006c354600000000")) desc, val = _describe_unusual(0x13, bytes.fromhex("00295c0b41006c354600000000"), 9)
assert desc == "竞价试买" assert desc == "竞价试买"
assert val == "8.71/11611手" assert val == "8.71/11611手"
def test_auction_test_sell(self): def test_auction_test_sell(self):
# 真实样本:603900 09:15:17,申报价 6.46 低于昨收 7.17(往下试) # 真实样本:603900 09:15:17,申报价 6.46 低于昨收 7.17(往下试)
desc, val = _describe_unusual(0x13, bytes.fromhex("0152b8ce400000c94300000000")) desc, val = _describe_unusual(0x13, bytes.fromhex("0152b8ce400000c94300000000"), 9)
assert desc == "竞价试卖" assert desc == "竞价试卖"
assert val == "6.46/402手" assert val == "6.46/402手"
@@ -150,13 +157,13 @@ class TestType0x1D0x1E:
def test_fast_rise(self): def test_fast_rise(self):
# 真实样本:605365 立达信 09:35:08 # 真实样本:605365 立达信 09:35:08
desc, val = _describe_unusual(0x1D, bytes.fromhex("009d50843c0000000000000000")) desc, val = _describe_unusual(0x1D, bytes.fromhex("009d50843c0000000000000000"), 9)
assert desc == "急速拉升" assert desc == "急速拉升"
assert val == "1.62%" assert val == "1.62%"
def test_fast_fall(self): def test_fast_fall(self):
# 真实样本:601123 N马矿 09:35:03 # 真实样本:601123 N马矿 09:35:03
desc, val = _describe_unusual(0x1E, bytes.fromhex("019cd393bc0000000000000000")) desc, val = _describe_unusual(0x1E, bytes.fromhex("019cd393bc0000000000000000"), 9)
assert desc == "急速下跌" assert desc == "急速下跌"
assert val == "-1.80%" assert val == "-1.80%"
@@ -211,7 +218,7 @@ class TestTypeNames:
"""映射表中的类型不应落入"异动类型0x??"兜底分支。""" """映射表中的类型不应落入"异动类型0x??"兜底分支。"""
zeros = bytes.fromhex("00" * 13) zeros = bytes.fromhex("00" * 13)
for utype in UNUSUAL_TYPE_NAMES: for utype in UNUSUAL_TYPE_NAMES:
desc, _ = _describe_unusual(utype, zeros) desc, _ = _describe_unusual(utype, zeros, 9)
assert not desc.startswith("异动类型"), f"0x{utype:02X} 未实现解析分支" assert not desc.startswith("异动类型"), f"0x{utype:02X} 未实现解析分支"
def test_top_level_export(self): def test_top_level_export(self):
+1 -1
View File
@@ -62,7 +62,7 @@ def test_vipdoc_settings_endpoints(settings_env, tmp_path):
r = client.put("/api/v1/settings/vipdoc", json={"path": str(real_dir)}) r = client.put("/api/v1/settings/vipdoc", json={"path": str(real_dir)})
assert r.status_code == 200 assert r.status_code == 200
assert r.json()["stored"] == str(real_dir) assert r.json()["stored"] == str(real_dir)
assert market_mod._limitup_cache is None assert market_mod._limitup_cache == {}
# GET 回读 # GET 回读
assert client.get("/api/v1/settings/vipdoc").json()["stored"] == str(real_dir) assert client.get("/api/v1/settings/vipdoc").json()["stored"] == str(real_dir)
+237
View File
@@ -2,6 +2,9 @@
from __future__ import annotations from __future__ import annotations
import datetime as _dt
import logging
import numpy as np import numpy as np
import pandas as pd import pandas as pd
import pytest import pytest
@@ -176,6 +179,131 @@ def test_promote_provisional(wh):
assert len(wh.query("SH", "600519")) == 3 # 转正后可见 assert len(wh.query("SH", "600519")) == 3 # 转正后可见
def test_promote_provisional_scoped_to_market_code_and_before(wh):
"""scoped 转正:只转正指定标的且 datetime <= before 的 provisional 行。"""
def _one(d: str) -> pd.DataFrame:
return pd.DataFrame(
{
"datetime": pd.date_range(d, periods=1),
"open": 10.0,
"high": 10.1,
"low": 9.9,
"close": 10.0,
"vol": 100.0,
"amount": 1000.0,
}
)
wh.upsert_bars("SH", "600519", _one("2024-01-05"), status="provisional")
wh.upsert_bars("SH", "600519", _one("2024-06-01"), status="provisional")
wh.upsert_bars("SZ", "000001", _one("2024-01-05"), status="provisional")
n = wh.promote_provisional(market="SH", code="600519", before=pd.Timestamp("2024-03-01"))
assert n == 1
out = wh.query("SH", "600519") # 默认查询只含 completed
assert len(out) == 1
assert pd.Timestamp(out["datetime"].iloc[0]) == pd.Timestamp("2024-01-05")
all_rows = wh.query("SH", "600519", include_provisional=True)
assert len(all_rows) == 2 # 2024-06-01 行超出 before,保持 provisional
def _fake_clock(
store_mod, # noqa: ANN001 — monkeypatch 目标模块(未用)
*,
shanghai: tuple[int, int, int, int],
local: tuple[int, int, int, int],
):
"""伪造 store 模块时钟:now(tz)=沪时区正确墙钟;now()=本地误判墙钟。
模拟UTC 主机沪市已 18:00当日 bar 应为 completed本地 naive
时钟却还是 10:00旧实现会误标 provisional
"""
class _FixedDT(_dt.datetime):
@classmethod
def now(cls, tz=None): # type: ignore[override]
if tz is not None:
y, m, d, hh = shanghai
return _dt.datetime(y, m, d, hh, 0, tzinfo=tz)
y, m, d, hh = local
return _dt.datetime(y, m, d, hh, 0)
return _FixedDT
def test_provisional_uses_shanghai_clock_not_local(wh, monkeypatch):
"""provisional 判定按沪市墙钟:沪市 18:00(收盘后)当日 bar 必须 completed。
回归旧实现用系统本地 now()UTC 主机上沪市收盘时本地才 10:00
当日 bar 被误标 provisional默认查询隐藏当天数据
"""
import easy_tdx.warehouse.store as store_mod
monkeypatch.setattr(
store_mod,
"datetime",
_fake_clock(store_mod, shanghai=(2026, 9, 7, 18), local=(2026, 9, 7, 10)),
)
dates = pd.date_range(
pd.Timestamp("2026-09-07") - pd.Timedelta(days=10), periods=11, freq="D"
).tolist() # 2026-08-28 .. 2026-09-07(末根 = 沪市「当日」)
df = pd.DataFrame(
{
"datetime": dates,
"open": 10.0,
"high": 10.1,
"low": 9.9,
"close": 10.0,
"vol": 100.0,
"amount": 1000.0,
}
)
wh.upsert_bars("SH", "600519", df)
# 沪市已收盘:全部 11 根都应为 completed(旧实现:当日根 provisional
assert len(wh.query("SH", "600519")) == 11
def test_promote_provisional_uses_shanghai_date(wh, monkeypatch):
"""无参转正的「今日」边界按沪市日期:沪市已过 0 点即转正昨日临时行。"""
import easy_tdx.warehouse.store as store_mod
monkeypatch.setattr(
store_mod,
"datetime",
_fake_clock(store_mod, shanghai=(2026, 9, 8, 0), local=(2026, 9, 7, 16)),
)
old = pd.DataFrame(
{
"datetime": pd.date_range("2026-09-07", periods=1),
"open": 10.0,
"high": 10.1,
"low": 9.9,
"close": 10.0,
"vol": 100.0,
"amount": 1000.0,
}
)
wh.upsert_bars("SH", "600519", old, status="provisional")
# 沪市日期已是 9/8 → 9/7 的临时行应转正(旧实现按本地 9/7 → n=0)
assert wh.promote_provisional() == 1
assert len(wh.query("SH", "600519")) == 1
def test_open_conflict_clear_error(tmp_path, monkeypatch):
"""仓库文件被其他进程占用:给可操作的中文错误而非裸 duckdb 异常。"""
import duckdb as duckdb_mod
def _raise(*args, **kwargs): # type: ignore[no-untyped-def]
raise duckdb_mod.IOException("Could not set lock on file")
monkeypatch.setattr(duckdb_mod, "connect", _raise)
with pytest.raises(RuntimeError, match="占用"):
KlineWarehouse(tmp_path / "lock.duckdb")
# ── 健康自检 ───────────────────────────────────────────────────────────────── # ── 健康自检 ─────────────────────────────────────────────────────────────────
@@ -291,6 +419,115 @@ def test_sync_progress_callback(tmp_path):
warehouse.close() warehouse.close()
class _ScriptedClient:
"""按调用序返回预置 DataFrame 的假客户端(末帧可重复)。"""
def __init__(self, frames: list[pd.DataFrame]) -> None:
self._frames = frames
self.calls: list[int] = []
def get_stock_kline(self, market, code, period="DAILY", start=0, count=800, adjust="NONE"):
self.calls.append(count)
idx = min(len(self.calls) - 1, len(self._frames) - 1)
return self._frames[idx].copy()
def test_sync_refetch_full_when_tail_gap(tmp_path, caplog):
"""增量尾部覆盖不到上次同步点(首 bar 晚于 existing_last)→ 全量重拉补缺。
回归旧实现固定只拉 tail_bars 超过 15 个交易日未同步的标的
中间日期永不补齐且无任何告警
"""
warehouse = KlineWarehouse(tmp_path / "gap.duckdb")
try:
source_full = _bars(130) # 2024-01-01 起 130 个工作日
initial = source_full.iloc[:100] # 首同步窗口(末根 idx99
stale_tail = source_full.iloc[115:] # 增量窗口:首根 idx115 > idx99 → 有缺口
client = _ScriptedClient([initial, stale_tail, source_full])
syncer = WarehouseSyncer(client, warehouse, max_bars=800, tail_bars=15)
with caplog.at_level(logging.WARNING, logger="easy_tdx.warehouse.sync"):
syncer.sync(["SH:600519"])
syncer.sync(["SH:600519"])
assert client.calls == [800, 15, 800] # 第二次 sync 触发了全量重拉
rows = warehouse.query("SH", "600519")
assert len(rows) == 130 # 无缺口
bridge = pd.Timestamp(source_full["datetime"].iloc[100])
dts = pd.to_datetime(rows["datetime"])
assert (dts == bridge).any() # 缺口桥接 bar 已补上
assert "缺口" in caplog.text
finally:
warehouse.close()
def test_sync_failure_keeps_provisional(tmp_path):
"""拉取失败:不转正 provisional,盘中临时值不会被洗成 completed。"""
warehouse = KlineWarehouse(tmp_path / "keep.duckdb")
try:
old = pd.DataFrame(
{
"datetime": pd.date_range("2024-01-01", periods=3),
"open": 10.0,
"high": 10.1,
"low": 9.9,
"close": 10.0,
"vol": 100.0,
"amount": 1000.0,
}
)
warehouse.upsert_bars("SH", "600519", old, status="provisional")
class _BadClient:
def get_stock_kline(self, *a, **kw):
raise ConnectionError("断网")
s = WarehouseSyncer(_BadClient(), warehouse).sync(["SH:600519"])
assert s["failed"] == 1
# 仍为 provisional:默认查询不可见(旧实现 sync 前盲转正 → 可见)
assert len(warehouse.query("SH", "600519")) == 0
assert len(warehouse.query("SH", "600519", include_provisional=True)) == 3
finally:
warehouse.close()
def test_sync_promotes_only_up_to_fetched_max(tmp_path):
"""转正上界 = 本次成功拉到的最大 datetime:未覆盖到的行保持 provisional。"""
warehouse = KlineWarehouse(tmp_path / "bound.duckdb")
try:
def _one(d: str) -> pd.DataFrame:
return pd.DataFrame(
{
"datetime": pd.date_range(d, periods=1),
"open": 10.0,
"high": 10.1,
"low": 9.9,
"close": 10.0,
"vol": 100.0,
"amount": 1000.0,
}
)
warehouse.upsert_bars("SH", "600519", _one("2024-01-05"), status="provisional")
warehouse.upsert_bars("SH", "600519", _one("2024-06-01"), status="provisional")
fetched = _bars(11, start="2024-01-10") # 最大 datetime 2024-01-24
client = _ScriptedClient([fetched])
WarehouseSyncer(client, warehouse, tail_bars=15).sync(["SH:600519"])
completed = warehouse.query("SH", "600519")
# 01-05 行 <= 拉取上界 → 已转正;06-01 行超出上界 → 保持 provisional
assert len(completed) == 12
all_rows = warehouse.query("SH", "600519", include_provisional=True)
assert len(all_rows) == 13
stale = all_rows[all_rows["status"] == "provisional"]
assert len(stale) == 1
assert pd.Timestamp(stale["datetime"].iloc[0]) == pd.Timestamp("2024-06-01")
finally:
warehouse.close()
def test_missing_duckdb_helpful_error(tmp_path, monkeypatch): def test_missing_duckdb_helpful_error(tmp_path, monkeypatch):
"""duckdb 未安装时给出安装指引(模拟 ImportError)。""" """duckdb 未安装时给出安装指引(模拟 ImportError)。"""
import builtins import builtins
+49
View File
@@ -153,3 +153,52 @@ def test_is_trading_hours() -> None:
assert _is_trading_hours(datetime(2026, 9, 1, 10, 0, tzinfo=tz)) is True # 周二盘中 assert _is_trading_hours(datetime(2026, 9, 1, 10, 0, tzinfo=tz)) is True # 周二盘中
assert _is_trading_hours(datetime(2026, 9, 1, 3, 0, tzinfo=tz)) is False # 凌晨 assert _is_trading_hours(datetime(2026, 9, 1, 3, 0, tzinfo=tz)) is False # 凌晨
assert _is_trading_hours(datetime(2026, 9, 5, 10, 0, tzinfo=tz)) is False # 周六 assert _is_trading_hours(datetime(2026, 9, 5, 10, 0, tzinfo=tz)) is False # 周六
# ── /watchlist 端点 code 格式校验(v1.32.6)─────────────────────────────────
def _watch_app(monkeypatch, tmp_path):
from fastapi import FastAPI
from easy_tdx.web import watchlist_store as ws
from easy_tdx.web.errors import register_exception_handlers
from easy_tdx.web.routers import watchlist as watchlist_mod
monkeypatch.setenv("EASY_TDX_CONFIG_DIR", str(tmp_path / "cfg"))
ws._store = None
app = FastAPI()
register_exception_handlers(app)
app.include_router(watchlist_mod.router, prefix="/api/v1")
return app
def test_watchlist_add_rejects_non_numeric_code(monkeypatch, tmp_path):
"""code 非 6 位数字 → 422(旧实现可把 'abcdef' 存进自选并喂给轮询器)。"""
pytest.importorskip("fastapi")
from fastapi.testclient import TestClient
app = _watch_app(monkeypatch, tmp_path)
with TestClient(app) as client:
bad = client.post("/api/v1/watchlist", json={"market": "SZ", "code": "abcdef", "name": "x"})
assert bad.status_code == 422
short = client.post(
"/api/v1/watchlist", json={"market": "SZ", "code": "00001", "name": "x"}
)
assert short.status_code == 422
ok = client.post(
"/api/v1/watchlist", json={"market": "SZ", "code": "000001", "name": "平安银行"}
)
assert ok.status_code == 200
def test_watchlist_remove_validates_code_format(monkeypatch, tmp_path):
"""remove 路径 code 非 6 位数字 → 422,不触达存储。"""
pytest.importorskip("fastapi")
from fastapi.testclient import TestClient
app = _watch_app(monkeypatch, tmp_path)
with TestClient(app) as client:
resp = client.delete("/api/v1/watchlist/SZ/abc123")
assert resp.status_code == 422
+30
View File
@@ -599,3 +599,33 @@ def test_create_app_no_ui_mode():
app_ui = create_app() app_ui = create_app()
mounts_ui = [r for r in app_ui.routes if type(r).__name__ == "Mount"] mounts_ui = [r for r in app_ui.routes if type(r).__name__ == "Mount"]
assert any(getattr(m, "name", "") == "web-ui" for m in mounts_ui) assert any(getattr(m, "name", "") == "web-ui" for m in mounts_ui)
# ── /server/test 输入约束(v1.32.6timeout 上界 + hosts 限长)───────────────
def test_server_test_request_constraints():
"""timeout 限 0.5~30shosts ≤50 项且单项 ≤253 字符。
旧实现 timeout 无上界1e9 会把 to_thread 线程挂死hosts 不限长
可当内网扫描跳板
"""
pytest.importorskip("fastapi")
from pydantic import ValidationError
from easy_tdx.web.routers.server import ServerTestRequest
assert ServerTestRequest(hosts=None, timeout=5.0).timeout == 5.0
assert ServerTestRequest(hosts=["127.0.0.1"], timeout=0.5).timeout == 0.5
with pytest.raises(ValidationError):
ServerTestRequest(timeout=31.0) # 超上界
with pytest.raises(ValidationError):
ServerTestRequest(timeout=0.1) # 低于下界
with pytest.raises(ValidationError):
ServerTestRequest(hosts=[f"h{i}" for i in range(51)]) # 超 50 项
with pytest.raises(ValidationError):
ServerTestRequest(hosts=["x" * 254]) # 单项超 253 字符
# 边界可用
ok = ServerTestRequest(hosts=["h" * 253] * 50, timeout=30.0)
assert len(ok.hosts) == 50
+296 -21
View File
@@ -349,37 +349,31 @@ def test_task_runner_captures_failure():
def test_task_runner_lru_eviction(): def test_task_runner_lru_eviction():
"""过上限应丢弃最旧的非 running 任务。 """限淘汰只移终态任务(v1.32.6pending/running 不淘汰,防幽灵任务
注意淘汰发生在 submit 淘汰对象是当时最旧的非 running 任务 确定性设计max_workers=1 串行先提交 5 个并等全部 done提交瞬间的
max_workers=1 串行执行时哪个任务被淘汰取决于提交速度 vs 执行速度 淘汰因全是 pending 而跳过这正是新语义随后提交第 6 个触发淘汰
的竞态快机器上 t0 还在 running 会被跳过慢机器上 t0 已完成会被淘汰 此时 5 个全是终态 LRU 淘到只剩 max_results=3t3/t4/t5
所以本测试不断言特定 task_id 被淘汰而是验证
(1) 存活的 non-running 任务数 max_results
(2) 最后提交的任务一定存活它是最近的不可能被 LRU 淘汰
(3) 至少有 2 个任务被淘汰5 提交 - 3 上限 = 2
""" """
from easy_tdx.web.task_runner import BacktestTaskRunner from easy_tdx.web.task_runner import BacktestTaskRunner
runner = BacktestTaskRunner(max_workers=1, max_results=3) runner = BacktestTaskRunner(max_workers=1, max_results=3)
ids = [runner.submit(lambda: {"i": i}, description=f"t{i}") for i in range(5)] ids = [runner.submit(lambda: {"i": i}, description=f"t{i}") for i in range(5)]
# 等待存活的任务全部完成(被淘汰的 peek 返回 None,跳过) # 等 5 个任务全部完成
for _ in range(200): for _ in range(200):
alive = [tid for tid in ids if runner.peek(tid) is not None] states = [runner.peek(tid) for tid in ids]
if all(runner.peek(tid).status in ("done", "failed") for tid in alive): if all(s is not None and s.status in ("done", "failed") for s in states):
break break
time.sleep(0.02) time.sleep(0.02)
# 最后提交的任务一定存活(LRU 最近,不可能被淘汰 # 全部完成后提交第 6 个 → 淘汰最旧的 3 个 donet0/t1/t2
assert runner.peek(ids[4]) is not None, "最后提交的任务不应被淘汰" ids.append(runner.submit(lambda: {"i": 5}, description="t5"))
assert runner.peek(ids[5]) is not None
# 至少淘汰 2 个(5 提交 - max_results 3 = 2 assert runner.peek(ids[0]) is None, "最旧的 done 应被淘汰"
surviving = [tid for tid in ids if runner.peek(tid) is not None] assert runner.peek(ids[1]) is None
evicted = [tid for tid in ids if runner.peek(tid) is None] assert runner.peek(ids[2]) is None
assert len(evicted) >= 2, f"应至少淘汰 2 个任务,实际淘汰 {len(evicted)}" for tid in ids[3:]:
assert runner.peek(tid) is not None, "最近的任务不应被淘汰"
# 存活任务数不超过 max_resultsrunning 完成后)
assert len(surviving) <= 3, f"存活任务 {len(surviving)} 超过上限 3"
runner.shutdown() runner.shutdown()
@@ -1365,3 +1359,284 @@ def test_multi_strategy_evaluate_endpoint(client, monkeypatch):
assert report["grade"]["scenario"] == "portfolio" assert report["grade"]["scenario"] == "portfolio"
assert report["fitness"]["total_checks"] == 8 assert report["fitness"]["total_checks"] == 8
assert report["config"]["slots"] == ["双均线交叉@SH:601088", "MACD 金叉@SZ:000001"] assert report["config"]["slots"] == ["双均线交叉@SH:601088", "MACD 金叉@SZ:000001"]
# ── submit 响应真实状态(v1.32.6 修复:不再把 done/failed 谎报为 running)──────
class _InstantDoneRunner:
"""submit 即同步跑完的假 runner(模拟"拿到 future 前任务已完成")。"""
def __init__(self, status: str = "done"):
self.status = status
def submit(self, func, *, description=""):
func() # 同步执行完毕
return "tid-done"
def get(self, task_id):
from easy_tdx.web.task_runner import TaskState
return TaskState(
task_id=task_id,
status=self.status, # type: ignore[arg-type]
result={
"performance": {},
"equity_curve": [],
"trades": [],
"positions": [],
"config": {},
},
finished_at=1.0,
started_at=0.0,
created_at=0.0,
)
def test_task_submit_response_accepts_terminal_status():
"""TaskSubmitResponse.status 应接受 done/failed(旧 Literal 只许 pending/running)。"""
from easy_tdx.web.backtest_schemas import TaskSubmitResponse
assert TaskSubmitResponse(task_id="x", status="done").status == "done"
assert TaskSubmitResponse(task_id="x", status="failed").status == "failed"
assert TaskSubmitResponse(task_id="x", status="pending").status == "pending"
def test_async_submit_reports_real_done_status(monkeypatch):
"""极快任务已 done 时,202 响应应透传真实状态 "done"(旧码谎报 "running")。"""
from fastapi import FastAPI
from fastapi.testclient import TestClient
from easy_tdx.web.errors import register_exception_handlers
from easy_tdx.web.routers import backtest as backtest_mod
monkeypatch.setattr(backtest_mod, "get_runner", lambda: _InstantDoneRunner("done"))
app = FastAPI()
register_exception_handlers(app)
app.include_router(backtest_mod.router, prefix="/api/v1")
app.state.tdx_client = object()
with TestClient(app) as tc:
resp = tc.post(
"/api/v1/backtest/run/async",
json={
"strategy": "ma_cross",
"params": {"fast": 3, "slow": 6},
"ohlcv": [
{
"datetime": f"2024-01-{d:02d}",
"open": 10.0,
"high": 10.5,
"low": 9.5,
"close": 10.0,
"vol": 1000.0,
"amount": 10000.0,
}
for d in range(1, 6)
],
},
)
assert resp.status_code == 202, resp.text
assert resp.json()["status"] == "done"
def test_async_submit_reports_failed_status(monkeypatch):
"""任务同步失败时响应应报 "failed" 而非 "running""""
from fastapi import FastAPI
from fastapi.testclient import TestClient
from easy_tdx.web.errors import register_exception_handlers
from easy_tdx.web.routers import backtest as backtest_mod
class _FailingRunner(_InstantDoneRunner):
def submit(self, func, *, description=""):
try:
func()
except Exception:
pass
return "tid-fail"
monkeypatch.setattr(backtest_mod, "get_runner", lambda: _FailingRunner("failed"))
app = FastAPI()
register_exception_handlers(app)
app.include_router(backtest_mod.router, prefix="/api/v1")
app.state.tdx_client = object()
with TestClient(app) as tc:
resp = tc.post(
"/api/v1/backtest/run/async",
json={
"strategy": "no_such_strategy",
"ohlcv": [
{
"datetime": "2024-01-01",
"open": 10.0,
"high": 10.5,
"low": 9.5,
"close": 10.0,
"vol": 1000.0,
"amount": 10000.0,
},
{
"datetime": "2024-01-02",
"open": 10.0,
"high": 10.5,
"low": 9.5,
"close": 10.0,
"vol": 1000.0,
"amount": 10000.0,
},
],
},
)
assert resp.status_code == 202, resp.text
assert resp.json()["status"] == "failed"
def test_formula_submit_reports_real_status(monkeypatch):
"""formula 回测提交响应透传真实状态(旧码硬编码 "running")。"""
from fastapi import FastAPI
from fastapi.testclient import TestClient
from easy_tdx.web.errors import register_exception_handlers
from easy_tdx.web.routers import formula as formula_mod
monkeypatch.setattr(formula_mod, "get_runner", lambda: _InstantDoneRunner("done"))
app = FastAPI()
register_exception_handlers(app)
app.include_router(formula_mod.router, prefix="/api/v1")
app.state.tdx_client = object()
ohlcv = [
{
"datetime": f"2024-01-{d:02d}",
"open": 10.0,
"high": 10.5,
"low": 9.5,
"close": 10.0,
"vol": 1000.0,
}
for d in range(1, 6)
]
with TestClient(app) as tc:
resp = tc.post(
"/api/v1/formula/backtest/run/async",
json={"text": "CROSS(C, MA(C, 3));", "ohlcv": ohlcv},
)
assert resp.status_code == 202, resp.text
assert resp.json()["status"] == "done"
# ── optimize 费率口径(stamp_tax / min_commission / auto_fees 透传)───────────
def test_optimize_request_accepts_fee_fields():
"""OptimizeBacktestRequest 应支持 stamp_tax/min_commission/auto_fees(镜像单标的)。"""
from easy_tdx.web.backtest_schemas import OptimizeBacktestRequest
req = OptimizeBacktestRequest(
strategy="ma_cross",
param_grid={"fast": [3]},
symbol="SZ:000001",
stamp_tax=0.0,
min_commission=1.0,
auto_fees=True,
)
assert req.stamp_tax == 0.0
assert req.min_commission == 1.0
assert req.auto_fees is True
with pytest.raises(ValueError):
OptimizeBacktestRequest(
strategy="ma_cross",
param_grid={"fast": [3]},
symbol="SZ:000001",
stamp_tax=0.5, # > le=0.01
)
def test_optimize_auto_fees_etf_matches_explicit_fee_backtest(sample_ohlcv):
"""ETF + auto_fees:寻优结果的买入持有基准与"显式 ETF 费率"口径一致。
旧实现不透传 stamp_tax/min_commission恒按股票默认 0.001/5.0品种
口径无法生效用可转债佣金 0.0002 / 最低佣金 1.0 / 免印花税验证
寻优结果的买入持有基准与"显式可转债费率"同口径且与旧股票默认口径
不同_BuyAndHold 不卖出印花税不进收益差异来自佣金/最低佣金
"""
from easy_tdx.backtest.benchmark import run_buy_hold_benchmark
from easy_tdx.web.backtest_schemas import OptimizeBacktestRequest
from easy_tdx.web.routers.backtest import _run_optimize
df = pd.DataFrame(sample_ohlcv)
df["datetime"] = pd.to_datetime(df["datetime"])
req = OptimizeBacktestRequest(
strategy="ma_cross",
param_grid={"fast": [3], "slow": [10]},
cash=1_000_000.0,
symbol="SH:110059", # 可转债:佣金 0.0002 / 最低佣金 1.0 / 免印花税
auto_fees=True,
)
out = _run_optimize(df, req)
# 期望口径:可转债费率
expected = run_buy_hold_benchmark(
df,
cash=req.cash,
commission=0.0002,
min_commission=1.0,
slippage=req.slippage,
execution=req.execution,
)
legacy = run_buy_hold_benchmark(
df,
cash=req.cash,
commission=req.commission, # 0.0003(股票默认)
min_commission=5.0,
slippage=req.slippage,
execution=req.execution,
)
assert out["buy_hold"] is not None
assert out["buy_hold"]["total_return"] == pytest.approx(expected["total_return"])
# 若与旧股票默认口径相同则说明 auto_fees 没生效
assert out["buy_hold"]["total_return"] != pytest.approx(legacy["total_return"])
def test_optimize_passes_resolved_fees_to_optimizer(sample_ohlcv, monkeypatch):
"""auto_fees 解析出的费率应透传给 ParamGridOptimizer(显式值优先)。"""
import easy_tdx.backtest.optimizer as opt_mod
from easy_tdx.web.backtest_schemas import OptimizeBacktestRequest
from easy_tdx.web.routers.backtest import _run_optimize
captured: dict = {}
class SpyOptimizer(opt_mod.ParamGridOptimizer):
def __init__(self, *args, **kwargs):
captured.update(kwargs)
super().__init__(*args, **kwargs)
monkeypatch.setattr(opt_mod, "ParamGridOptimizer", SpyOptimizer)
df = pd.DataFrame(sample_ohlcv)
df["datetime"] = pd.to_datetime(df["datetime"])
req = OptimizeBacktestRequest(
strategy="ma_cross",
param_grid={"fast": [3], "slow": [10]},
symbol="SH:510300",
auto_fees=True,
)
_run_optimize(df, req)
# ETF 口径:印花税解析为 0
assert captured.get("stamp_tax") == 0.0
assert captured.get("min_commission") == 5.0
# 显式非默认 stamp_tax 优先于品种默认
captured.clear()
req2 = OptimizeBacktestRequest(
strategy="ma_cross",
param_grid={"fast": [3], "slow": [10]},
symbol="SH:510300",
auto_fees=True,
stamp_tax=0.005,
)
_run_optimize(df, req2)
assert captured.get("stamp_tax") == 0.005
+219
View File
@@ -0,0 +1,219 @@
"""/bars、/bars/index 的 baostock 兜底集成测试(v1.32.6 修复项)。
覆盖
- 数字周期字符串category="4"归一后也能走兜底旧实现直接透传原串
baostock 频率查表落空 兜底静默失效维持原错误
- 指数兜底必须传 is_index=Truebaostock 指数 vol 个股路径不传
- fetch_bars 真故障抛 RuntimeError 时按"兜底不可用"处理维持原 TDX 错误
"""
from __future__ import annotations
import sys
import types
import pandas as pd
import pytest
pytest.importorskip("fastapi")
from fastapi.testclient import TestClient # noqa: E402
# ── 测试替身 ─────────────────────────────────────────────────────────────────
class _RaisingMac:
async def get_stock_kline(self, *args, **kwargs): # noqa: ANN002, ANN003
raise RuntimeError("MAC 连接失败")
class _RaisingTdx:
async def get_security_bars(self, *args, **kwargs): # noqa: ANN002, ANN003
raise RuntimeError("标准协议连接失败")
async def get_index_bars(self, *args, **kwargs): # noqa: ANN002, ANN003
raise RuntimeError("标准协议连接失败")
def _bars_app(mac_client, tdx_client):
from fastapi import FastAPI
from easy_tdx.web.errors import register_exception_handlers
from easy_tdx.web.routers import bars
app = FastAPI()
register_exception_handlers(app)
app.include_router(bars.router, prefix="/api/v1")
app.state.tdx_client = tdx_client
app.state.mac_client = mac_client
return app
def _fallback_df(n: int = 5) -> pd.DataFrame:
dates = pd.bdate_range(end="2026-09-04", periods=n)
return pd.DataFrame(
{
"date": dates.normalize(),
"open": [10.0] * n,
"close": [10.5] * n,
"high": [11.0] * n,
"low": [9.5] * n,
"vol": [100000.0] * n,
"amount": [1050000.0] * n,
}
)
def _install_fake_bs_module(monkeypatch: pytest.MonkeyPatch, rows: int = 10) -> dict:
"""装一个最小可用的 baostock 模块替身,返回 captured 观测点。"""
from easy_tdx.sources import baostock as bs_source
captured: dict = {}
def _login():
lg = types.SimpleNamespace()
lg.error_code = "0"
lg.error_msg = "ok"
return lg
def query_history_k_data_plus(**kwargs): # noqa: ANN003
captured.update(kwargs)
captured["calls"] = captured.get("calls", 0) + 1
data = [
[f"2026-08-{d:02d}", "10.0", "10.5", "11.0", "9.5", "100000", "1050000", "1"]
for d in range(1, rows + 1)
]
rs = types.SimpleNamespace()
rs.error_code = "0"
rs.error_msg = "ok"
rs._rows = data
rs._i = 0
rs.next = lambda: rs._i < len(rs._rows) # type: ignore[method-assign]
rs.get_row_data = lambda: rs._rows[rs._i] # type: ignore[method-assign]
def _advance():
row = rs._rows[rs._i]
rs._i += 1
return row
rs.get_row_data = _advance # type: ignore[method-assign]
return rs
mod = types.ModuleType("baostock")
mod.login = _login # type: ignore[attr-defined]
mod.logout = lambda: None # type: ignore[attr-defined]
mod.query_history_k_data_plus = query_history_k_data_plus # type: ignore[attr-defined]
monkeypatch.setitem(sys.modules, "baostock", mod)
monkeypatch.delenv("EASY_TDX_BAOSTOCK", raising=False)
monkeypatch.setattr(bs_source, "_logged_in", False)
return captured
# ── 项11:数字周期字符串归一后再兜底 ──────────────────────────────────────────
def test_bars_numeric_category_still_falls_back(monkeypatch):
"""category="4"(=DAY 的数字形式)TDX 全败时也应命中 baostock 兜底。
旧实现把原串 "4" 透传给 fetch_bars_FREQ_BY_CATEGORY.get("4") 落空
返回 None 兜底静默失效客户端拿到 500
"""
captured = _install_fake_bs_module(monkeypatch)
with TestClient(
_bars_app(_RaisingMac(), _RaisingTdx()), raise_server_exceptions=False
) as client:
resp = client.get(
"/api/v1/bars", params={"market": "SH", "code": "600519", "category": "4"}
)
assert resp.status_code == 200, resp.text
body = resp.json()
assert body["source"] == "baostock"
assert body["count"] > 0
assert captured["frequency"] == "d" # 归一成 DAY 后映射到日线
# ── 项12:is_index 传递与异常语义 ────────────────────────────────────────────
def test_index_fallback_passes_is_index_true(monkeypatch):
"""/bars/index 兜底必须带 is_index=True(指数 vol 股→手 ÷100)。"""
from easy_tdx.sources import baostock as bs_source
calls: dict = {}
def fake_fetch(market, code, category, start, count, adjust, is_index=False):
calls["is_index"] = is_index
return _fallback_df()
monkeypatch.setattr(bs_source, "is_enabled", lambda: True)
monkeypatch.setattr(bs_source, "fetch_bars", fake_fetch)
with TestClient(_bars_app(None, _RaisingTdx())) as client:
resp = client.get(
"/api/v1/bars/index", params={"market": "SH", "code": "000001", "category": "DAY"}
)
assert resp.status_code == 200
assert resp.json()["source"] == "baostock"
assert calls["is_index"] is True
def test_bars_stock_fallback_keeps_is_index_false(monkeypatch):
"""个股路径兜底 is_index=Falsevol 保持股口径)。"""
from easy_tdx.sources import baostock as bs_source
calls: dict = {}
def fake_fetch(market, code, category, start, count, adjust, is_index=False):
calls["is_index"] = is_index
return _fallback_df()
monkeypatch.setattr(bs_source, "is_enabled", lambda: True)
monkeypatch.setattr(bs_source, "fetch_bars", fake_fetch)
with TestClient(_bars_app(_RaisingMac(), _RaisingTdx())) as client:
resp = client.get("/api/v1/bars", params={"market": "SH", "code": "600519"})
assert resp.status_code == 200
assert calls["is_index"] is False
def test_bars_fallback_exception_keeps_original_tdx_error(monkeypatch):
"""fetch_bars 真故障抛 RuntimeError → 按"兜底不可用"处理,重抛原 TDX 异常。
响应错误详情须是标准协议的失败原因而非 baostock 的失败原因baostock
的失败只记日志且不返回空数据伪装成功
"""
from easy_tdx.sources import baostock as bs_source
def boom(*args, **kwargs): # noqa: ANN002, ANN003
raise RuntimeError("baostock 拉取失败: 网络异常")
monkeypatch.setattr(bs_source, "is_enabled", lambda: True)
monkeypatch.setattr(bs_source, "fetch_bars", boom)
with TestClient(
_bars_app(_RaisingMac(), _RaisingTdx()), raise_server_exceptions=False
) as client:
resp = client.get("/api/v1/bars", params={"market": "SH", "code": "600519"})
assert resp.status_code == 500
assert "标准协议连接失败" in resp.json()["detail"]
assert "baostock" not in resp.json()["detail"]
def test_index_fallback_exception_keeps_original_tdx_error(monkeypatch):
"""/bars/index 同语义:baostock 异常不吞掉原 TDX 错误。"""
from easy_tdx.sources import baostock as bs_source
def boom(*args, **kwargs): # noqa: ANN002, ANN003
raise RuntimeError("baostock 拉取失败: 网络异常")
monkeypatch.setattr(bs_source, "is_enabled", lambda: True)
monkeypatch.setattr(bs_source, "fetch_bars", boom)
with TestClient(_bars_app(None, _RaisingTdx()), raise_server_exceptions=False) as client:
resp = client.get(
"/api/v1/bars/index", params={"market": "SH", "code": "000001", "category": "DAY"}
)
assert resp.status_code == 500
assert "标准协议连接失败" in resp.json()["detail"]
+172
View File
@@ -0,0 +1,172 @@
"""count>800 的分页取数测试(离线)。
TDX 协议单次 get_security_bars 最多返回 800 旧实现里 multiseed /
rotation / formula 的单次调用在 count>800 时被服务器静默截断本文件钉死
"分页取全量 + 页序正确 + 数据起点提前停止"行为
"""
from __future__ import annotations
import asyncio
import pandas as pd
import pytest
pytest.importorskip("fastapi")
_PAGE_CAP = 800
class _CappedBarsClient:
"""模拟 TDX 服务器:单次最多返回 _PAGE_CAP 根,start 为回看偏移。"""
def __init__(self, total_bars: int = 3000):
self.total_bars = total_bars
self.calls: list[tuple[int, int]] = [] # (start, count)
def _make_page(self, start: int, n: int) -> pd.DataFrame:
"""start 偏移处往前 n 根(升序页);越过数据起点则截断为 0 根。"""
hi = self.total_bars - start # 本页最旧一根的全局序号(0 起)
lo = max(0, hi - n)
if hi <= 0:
return pd.DataFrame()
dates = pd.date_range("2020-01-01", periods=self.total_bars, freq="B")
idx = dates[lo:hi]
return pd.DataFrame(
{
"date": idx,
"open": 10.0,
"high": 11.0,
"low": 9.0,
"close": 10.5,
"vol": 1000.0,
"amount": 10000.0,
}
)
async def get_security_bars(self, market, code, category, start, count, **kw):
self.calls.append((int(start), int(count)))
return self._make_page(int(start), int(count))
# ── 共享分页辅助 ───────────────────────────────────────────────────────────────
def test_fetch_bars_paged_requests_multiple_pages():
"""count=2000 → 3 次请求(800/800/400),拼齐 2000 根且时间升序。"""
from easy_tdx.web.routers.backtest import _fetch_bars_paged
fake = _CappedBarsClient(total_bars=3000)
df = asyncio.run(_fetch_bars_paged(fake, "SZ:000001", "DAY", 2000))
assert fake.calls == [(0, 800), (800, 800), (1600, 400)]
assert len(df) == 2000
dates = pd.to_datetime(df["date"])
assert dates.is_monotonic_increasing # 页序拼接后必须升序
def test_fetch_bars_paged_stops_at_data_start():
"""数据起点不足一页时提前停止,不多发请求。"""
from easy_tdx.web.routers.backtest import _fetch_bars_paged
fake = _CappedBarsClient(total_bars=1000)
df = asyncio.run(_fetch_bars_paged(fake, "SH:600519", "DAY", 2000))
# 第二页只回 200 根(不足一页)= 数据起点,循环不再发第三笔请求
assert fake.calls == [(0, 800), (800, 800)]
assert len(df) == 1000
def test_fetch_bars_paged_small_count_single_call():
"""count≤800 仍单页取齐(不多打请求)。"""
from easy_tdx.web.routers.backtest import _fetch_bars_paged
fake = _CappedBarsClient(total_bars=3000)
df = asyncio.run(_fetch_bars_paged(fake, "SZ:000001", "DAY", 250))
assert fake.calls == [(0, 250)]
assert len(df) == 250
def test_fetch_bars_paged_empty_returns_empty_df():
from easy_tdx.web.routers.backtest import _fetch_bars_paged
fake = _CappedBarsClient(total_bars=0)
df = asyncio.run(_fetch_bars_paged(fake, "SZ:000001", "DAY", 800))
assert df.empty
# ── multiseed / rotation 端点(取数在 handler 内完成,POST 返回即可断言)──────
def _app_with(fake_client):
from fastapi import FastAPI
from easy_tdx.web.errors import register_exception_handlers
from easy_tdx.web.routers import backtest as backtest_mod
app = FastAPI()
register_exception_handlers(app)
app.include_router(backtest_mod.router, prefix="/api/v1")
app.state.tdx_client = fake_client
app.state.mac_client = None
app.state.ex_client = None
return app
def test_multiseed_fetches_full_count_via_paging():
"""multiseed count=900>800)→ 每标的 2 次请求,不再被 800 截断。"""
pytest.importorskip("fastapi")
from fastapi.testclient import TestClient
fake = _CappedBarsClient(total_bars=3000)
with TestClient(_app_with(fake)) as client:
resp = client.post(
"/api/v1/backtest/multiseed/run/async",
json={
"strategy": "ma_cross",
"params": {"fast": 3, "slow": 6},
"stocks": ["SZ:000001", "SH:600519"],
"count": 900,
},
)
assert resp.status_code == 202, resp.text
# 2 标的 × 2 页
assert fake.calls == [(0, 800), (800, 100), (0, 800), (800, 100)]
def test_rotation_fetches_full_count_via_paging():
"""rotation count=900>800)→ 每标的 2 次请求。"""
pytest.importorskip("fastapi")
from fastapi.testclient import TestClient
fake = _CappedBarsClient(total_bars=3000)
with TestClient(_app_with(fake)) as client:
resp = client.post(
"/api/v1/backtest/rotation/run/async",
json={
"stocks": ["SZ:000001", "SH:600519"],
"count": 900,
},
)
assert resp.status_code == 202, resp.text
assert fake.calls == [(0, 800), (800, 100), (0, 800), (800, 100)]
# ── formula 取数路径 ──────────────────────────────────────────────────────────
def test_formula_resolve_df_pages_full_count():
"""formula _resolve_df symbol 路径 count=2000 → 3 页拼齐且升序。"""
from easy_tdx.web.routers.backtest import _fetch_bars_paged # noqa: F401 需已存在
from easy_tdx.web.routers.formula import FormulaComputeRequest, _resolve_df
fake = _CappedBarsClient(total_bars=3000)
df = asyncio.run(
_resolve_df(fake, FormulaComputeRequest(text="C", symbol="SZ:000001", count=2000))
)
assert fake.calls == [(0, 800), (800, 800), (1600, 400)]
assert len(df) == 2000
dates = pd.to_datetime(df["date"])
assert dates.is_monotonic_increasing
+108
View File
@@ -0,0 +1,108 @@
/**
* api.ts Node test runnerfetch DOM
*
* node --test src/__tests__/api.test.ts
*
*
* 1. fetchRankList market (market=2) 'BJ'
* fetchBoardMembers 'SZ'
* 2. runLlmChatWithPolling AbortSignalsignal
* AbortError 退
*/
import { test } from 'node:test'
import assert from 'node:assert/strict'
import { fetchRankList, runLlmChatWithPolling } from '../api.ts'
import type { TaskState } from '../types.ts'
/** 临时替换 globalThis.fetch,返回可编程响应序列。 */
function stubFetch(handler: (url: string, init?: RequestInit) => unknown): {
calls: string[]
restore: () => void
} {
const calls: string[] = []
const original = globalThis.fetch
globalThis.fetch = (async (input: RequestInfo | URL, init?: RequestInit) => {
const url = String(input instanceof URL ? input : input)
calls.push(url)
const body = handler(url, init)
return new Response(JSON.stringify(body), {
status: 200,
headers: { 'Content-Type': 'application/json' },
})
}) as typeof fetch
return {
calls,
restore: () => {
globalThis.fetch = original
},
}
}
// ── 1. fetchRankList 北交所映射 ───────────────────────────────────────────────
test('fetchRankList: market=2(北交所)归一化为 BJ', async () => {
const stub = stubFetch(() => ({
data: [
{ close: 10.5, pre_close: 9.55, market: 2, code: '920002', name: '测试北交所' },
{ close: 8.0, pre_close: 8.8, market: 0, code: '000001', name: '平安银行' },
{ close: 20.0, pre_close: 19.0, market: 1, code: '600519', name: '贵州茅台' },
],
}))
try {
const rows = await fetchRankList('DESC', 3)
assert.equal(rows[0].market, 'BJ', 'market=2 应映射为 BJ(旧码错标 SZ')
assert.equal(rows[1].market, 'SZ', 'market=0 应映射为 SZ')
assert.equal(rows[2].market, 'SH', 'market=1 应映射为 SH')
// 顺带核对涨跌幅口径:close/pre-1
assert.ok(Math.abs((rows[0].change_pct ?? 0) - ((10.5 / 9.55 - 1) * 100)) < 1e-6)
} finally {
stub.restore()
}
})
// ── 2. runLlmChatWithPolling 的 AbortSignal ──────────────────────────────────
const runningState: TaskState = { status: 'running' } as unknown as TaskState
test('runLlmChatWithPolling: signal 中止后立即退出并抛 AbortError', async () => {
const stub = stubFetch((url) => {
if (url.includes('/llm/chat/async')) return { task_id: 't1' }
return runningState
})
try {
const ctrl = new AbortController()
const promise = runLlmChatWithPolling(
'ping',
null,
() => ctrl.abort(), // 第一次轮询即中止
5, // intervalMs
30_000, // timeoutMs(远大于中止所需时间,确保超时兜底不先触发)
ctrl.signal,
)
await assert.rejects(promise, (e: unknown) => (e as Error).name === 'AbortError')
// 中止后不应继续轮询:任务查询次数应极少(≤3)
const polls = stub.calls.filter((u) => u.includes('/llm/chat/tasks/')).length
assert.ok(polls <= 3, `中止后应停止轮询,实际轮询 ${polls}`)
} finally {
stub.restore()
}
})
test('runLlmChatWithPolling: 未中止时正常返回 done', async () => {
let polls = 0
const stub = stubFetch((url) => {
if (url.includes('/llm/chat/async')) return { task_id: 't2' }
polls += 1
if (polls < 3) return runningState
return { status: 'done', result: { reply: 'ok', model: 'm', provider: 'p' } }
})
try {
const state = await runLlmChatWithPolling('ping', null, undefined, 1, 30_000)
assert.equal(state.status, 'done')
assert.equal((state.result as { reply?: string } | null)?.reply, 'ok')
} finally {
stub.restore()
}
})
+14 -1
View File
@@ -646,7 +646,9 @@ export async function fetchRankList(
const r = { ...row } as Record<string, unknown> const r = { ...row } as Record<string, unknown>
r.price = close r.price = close
r.change_pct = pre > 0 ? (close / pre - 1) * 100 : 0 r.change_pct = pre > 0 ? (close / pre - 1) * 100 : 0
r.market = Number(row.market) === 1 ? 'SH' : 'SZ' // 与 fetchBoardMembers 同口径:MAC 协议 market 1=SH / 2=BJ / 其余 SZ
const m = Number(row.market)
r.market = m === 1 ? 'SH' : m === 2 ? 'BJ' : 'SZ'
return r as RankRow return r as RankRow
}) })
} }
@@ -799,6 +801,9 @@ export async function fetchLlmChatTask(taskId: string): Promise<TaskState> {
* *
* 1-3 1.5s 0.3s * 1-3 1.5s 0.3s
* 20 LLM 600s * 20 LLM 600s
*
* @param signal abort
* name='AbortError' 退
*/ */
export async function runLlmChatWithPolling( export async function runLlmChatWithPolling(
prompt: string, prompt: string,
@@ -806,11 +811,13 @@ export async function runLlmChatWithPolling(
onPoll?: (state: TaskState) => void, onPoll?: (state: TaskState) => void,
intervalMs = 1_500, intervalMs = 1_500,
timeoutMs = 20 * 60_000, timeoutMs = 20 * 60_000,
signal?: AbortSignal,
): Promise<TaskState> { ): Promise<TaskState> {
const { task_id } = await submitLlmChatTask(prompt, context) const { task_id } = await submitLlmChatTask(prompt, context)
const start = Date.now() const start = Date.now()
// eslint-disable-next-line no-constant-condition // eslint-disable-next-line no-constant-condition
while (true) { while (true) {
if (signal?.aborted) throw abortError()
const state = await fetchLlmChatTask(task_id) const state = await fetchLlmChatTask(task_id)
onPoll?.(state) onPoll?.(state)
if (state.status === 'done' || state.status === 'failed') return state if (state.status === 'done' || state.status === 'failed') return state
@@ -818,9 +825,15 @@ export async function runLlmChatWithPolling(
throw new Error(`AI 解读任务超时(${timeoutMs / 1000}s),任务仍在后台运行,可稍后重试`) throw new Error(`AI 解读任务超时(${timeoutMs / 1000}s),任务仍在后台运行,可稍后重试`)
} }
await new Promise((r) => setTimeout(r, intervalMs)) await new Promise((r) => setTimeout(r, intervalMs))
if (signal?.aborted) throw abortError()
} }
} }
/** 构造与 fetch 中止一致的 AbortError(便于调用方按 name 识别并静默)。 */
function abortError(): Error {
return new DOMException('AI 解读已取消', 'AbortError')
}
// ── AI 解读历史 ────────────────────────────────────────────────────────────── // ── AI 解读历史 ──────────────────────────────────────────────────────────────
/** 列出 AI 解读历史(时间倒序,含 Prompt/正文/策略上下文)。 */ /** 列出 AI 解读历史(时间倒序,含 Prompt/正文/策略上下文)。 */
+25 -2
View File
@@ -2,7 +2,7 @@
// AI /Prompt + / + // AI /Prompt + / +
// Prompt // Prompt
// LLM AI // LLM AI
import { onMounted, ref, watch } from 'vue' import { onBeforeUnmount, onMounted, ref, watch } from 'vue'
import { formatError, fetchLlmConfig, runLlmChatWithPolling } from '../api' import { formatError, fetchLlmConfig, runLlmChatWithPolling } from '../api'
import type { LlmChatContext, LlmChatResult } from '../types' import type { LlmChatContext, LlmChatResult } from '../types'
@@ -28,6 +28,8 @@ const aiRunning = ref(false)
const aiElapsed = ref(0) const aiElapsed = ref(0)
const aiReply = ref('') const aiReply = ref('')
let aiTimer = 0 let aiTimer = 0
/** 卸载时中止后台轮询(runLlmChatWithPolling 检查 signal 立即退出)。 */
let abortCtrl: AbortController | null = null
onMounted(() => { onMounted(() => {
// LLM Prompt // LLM Prompt
@@ -42,6 +44,15 @@ onMounted(() => {
}) })
}) })
onBeforeUnmount(() => {
if (aiTimer) {
window.clearInterval(aiTimer)
aiTimer = 0
}
abortCtrl?.abort()
abortCtrl = null
})
watch( watch(
() => props.prompt, () => props.prompt,
() => { () => {
@@ -61,8 +72,16 @@ async function runAiInterpret() {
aiTimer = window.setInterval(() => { aiTimer = window.setInterval(() => {
aiElapsed.value += 1 aiElapsed.value += 1
}, 1000) }, 1000)
abortCtrl = new AbortController()
try { try {
const state = await runLlmChatWithPolling(props.prompt, props.context) const state = await runLlmChatWithPolling(
props.prompt,
props.context,
undefined,
undefined,
undefined,
abortCtrl.signal,
)
// TaskState.result LLM // TaskState.result LLM
const r = state.result as LlmChatResult | null const r = state.result as LlmChatResult | null
// failed // failed
@@ -75,9 +94,13 @@ async function runAiInterpret() {
aiMsg.value = `解读失败:${state.error ?? '未知错误'}(可在「AI 设置」检查配置,或复制 Prompt 手动使用)` aiMsg.value = `解读失败:${state.error ?? '未知错误'}(可在「AI 设置」检查配置,或复制 Prompt 手动使用)`
} }
} catch (e) { } catch (e) {
// DOM
if ((e as Error)?.name === 'AbortError') return
aiMsg.value = `解读失败:${formatError(e)}(可在「AI 设置」检查配置,或复制 Prompt 手动使用)` aiMsg.value = `解读失败:${formatError(e)}(可在「AI 设置」检查配置,或复制 Prompt 手动使用)`
} finally { } finally {
window.clearInterval(aiTimer) window.clearInterval(aiTimer)
aiTimer = 0
abortCtrl = null
aiRunning.value = false aiRunning.value = false
} }
} }
+10 -2
View File
@@ -139,16 +139,24 @@ const membersLoading = ref(false)
// 80 1000 // 80 1000
const MEMBER_FETCH_COUNT = 1000 const MEMBER_FETCH_COUNT = 1000
/** 请求序号守卫:升/降序快速连点、或 props.code 切换时,旧响应直接丢弃。 */
let memberSeq = 0
async function loadMembers() { async function loadMembers() {
const my = ++memberSeq
membersLoading.value = true membersLoading.value = true
membersError.value = '' membersError.value = ''
try { try {
members.value = await fetchBoardMembers(props.code, MEMBER_FETCH_COUNT, memberOrder.value) const rows = await fetchBoardMembers(props.code, MEMBER_FETCH_COUNT, memberOrder.value)
// /
if (my !== memberSeq) return
members.value = rows
} catch (e) { } catch (e) {
if (my !== memberSeq) return
members.value = [] members.value = []
membersError.value = formatError(e) membersError.value = formatError(e)
} finally { } finally {
membersLoading.value = false if (my === memberSeq) membersLoading.value = false
} }
} }
@@ -18,6 +18,8 @@ const error = ref('')
const loading = ref(false) const loading = ref(false)
let poll: number | null = null let poll: number | null = null
/** 请求序号守卫:await 后比对,过期响应(旧类型/旧窗口)直接丢弃。 */
let loadSeq = 0
function stopPoll() { function stopPoll() {
if (poll !== null) { if (poll !== null) {
@@ -27,9 +29,12 @@ function stopPoll() {
} }
async function load() { async function load() {
const my = ++loadSeq
error.value = '' error.value = ''
try { try {
const r = await fetchHotspotCorrelation(props.boardType, props.days, props.perDay) const r = await fetchHotspotCorrelation(props.boardType, props.days, props.perDay)
// resp
if (my !== loadSeq) return
if (r.status === 'building') { if (r.status === 'building') {
resp.value = null resp.value = null
loading.value = true loading.value = true
@@ -45,6 +50,7 @@ async function load() {
resp.value = r resp.value = r
render() render()
} catch (e) { } catch (e) {
if (my !== loadSeq) return
stopPoll() stopPoll()
loading.value = false loading.value = false
error.value = formatError(e) error.value = formatError(e)
+4 -2
View File
@@ -35,7 +35,9 @@ const geom = computed(() => {
</script> </script>
<template> <template>
<svg :width="W" :height="H" class="sparkline" viewBox="0 0 92 28" preserveAspectRatio="none"> <!-- viewBox 必须跟随 width/height几何坐标按 W/H 计算写死 92×28 会让
传大尺寸的调用方Dashboard 150×30右侧被 viewport 裁剪 -->
<svg :width="W" :height="H" class="sparkline" :viewBox="`0 0 ${W} ${H}`" preserveAspectRatio="none">
<template v-if="geom"> <template v-if="geom">
<line <line
v-if="geom.baseValid" v-if="geom.baseValid"
@@ -56,7 +58,7 @@ const geom = computed(() => {
vector-effect="non-scaling-stroke" vector-effect="non-scaling-stroke"
/> />
</template> </template>
<text v-else x="46" y="18" text-anchor="middle" fill="#5c6370" font-size="9">加载中</text> <text v-else :x="W / 2" :y="H - 10" text-anchor="middle" fill="#5c6370" font-size="9">加载中</text>
</svg> </svg>
</template> </template>
+1 -1
View File
@@ -12,7 +12,7 @@
* / 252 线/线 * / 252 线/线
*/ */
import type { EquityPoint } from '../types' import type { EquityPoint } from '../types.ts'
/** 年化因子(按交易日)。 */ /** 年化因子(按交易日)。 */
const TRADING_DAYS_PER_YEAR = 252 const TRADING_DAYS_PER_YEAR = 252
+2 -2
View File
@@ -4,8 +4,8 @@
* + // * + //
*/ */
import { GRADE_THRESHOLDS, type DimensionScore, type Grade, type GradeResult, type VetoHit } from './types' import { GRADE_THRESHOLDS, type DimensionScore, type Grade, type GradeResult, type VetoHit } from './types.ts'
import { THRESHOLDS, type DimensionKey } from './thresholds' import { THRESHOLDS, type DimensionKey } from './thresholds.ts'
/** /**
* 线 0100 * 线 0100
+11 -9
View File
@@ -12,10 +12,10 @@
* @see docs/superpowers/plans * @see docs/superpowers/plans
*/ */
import type { BacktestResult, EquityPoint, GridPointResult, Performance, PortfolioResult } from '../types' import type { BacktestResult, EquityPoint, GridPointResult, Performance, PortfolioResult } from '../types.ts'
import { buildResult, scoreDimension } from './engine' import { buildResult, scoreDimension } from './engine.ts'
import { computeCombinedMetrics } from './combinedMetrics' import { computeCombinedMetrics } from './combinedMetrics.ts'
import type { DimensionScore, GradeResult, VetoHit } from './types' import type { DimensionScore, GradeResult, VetoHit } from './types.ts'
// ════════════════════════════════════════════════════════════════════════════ // ════════════════════════════════════════════════════════════════════════════
// 一票否决规则(所有场景共用) // 一票否决规则(所有场景共用)
@@ -315,8 +315,10 @@ export function gradeBacktestResult(result: BacktestResult): GradeResult {
} }
// ── 重新导出常用类型和工具,便于调用方一处 import ─────────────────────────── // ── 重新导出常用类型和工具,便于调用方一处 import ───────────────────────────
export { GRADE_META, GRADE_THRESHOLDS } from './types' // 注:相对路径统一带 .ts 扩展——Node --test 直跑(type-stripping 不改写
export type { Grade, GradeResult, DimensionScore, VetoHit, GradeMeta } from './types' // import 说明符)与 Vite/vue-tscallowImportingTsExtensions)两侧都可用。
export { worseGrade, scoreToGrade } from './engine' export { GRADE_META, GRADE_THRESHOLDS } from './types.ts'
export { computeCombinedMetrics } from './combinedMetrics' export type { Grade, GradeResult, DimensionScore, VetoHit, GradeMeta } from './types.ts'
export type { CombinedMetrics } from './combinedMetrics' export { worseGrade, scoreToGrade } from './engine.ts'
export { computeCombinedMetrics } from './combinedMetrics.ts'
export type { CombinedMetrics } from './combinedMetrics.ts'
+9 -1
View File
@@ -37,18 +37,26 @@ const error = ref('')
const lastRefresh = ref('') const lastRefresh = ref('')
const stat = ref<MarketStat | null>(null) const stat = ref<MarketStat | null>(null)
/** 请求序号守卫:await 后比对,过期响应直接丢弃。 */
let overviewSeq = 0
async function loadOverview() { async function loadOverview() {
const my = ++overviewSeq
loading.value = rows.value.length === 0 loading.value = rows.value.length === 0
error.value = '' error.value = ''
try { try {
const resp = await fetchBoardOverview(activeType.value) const resp = await fetchBoardOverview(activeType.value)
// / rows
// diffFlips /绿
if (my !== overviewSeq) return
rows.value = resp.rows rows.value = resp.rows
lastRefresh.value = new Date().toLocaleTimeString('zh-CN', { hour12: false }) lastRefresh.value = new Date().toLocaleTimeString('zh-CN', { hour12: false })
diffFlips(resp.rows) diffFlips(resp.rows)
} catch (e) { } catch (e) {
if (my !== overviewSeq) return
error.value = formatError(e) error.value = formatError(e)
} finally { } finally {
loading.value = false if (my === overviewSeq) loading.value = false
} }
} }
+5 -3
View File
@@ -265,10 +265,11 @@ function distHeight(count: number): string {
} }
function distColor(i: number): string { function distColor(i: number): string {
// 1..20 -10..+9 10 绿 10 // 1..20 -10..+91..10 绿110 12..20
if (i === 0) return 'var(--down)' if (i === 0) return 'var(--down)'
if (i === BUCKETS.length - 1) return 'var(--up)' if (i === BUCKETS.length - 1) return 'var(--up)'
return i <= 10 ? 'var(--down)' : 'var(--up)' if (i === 11) return 'var(--text-dim)'
return i < 11 ? 'var(--down)' : 'var(--up)'
} }
// + 120 / // + 120 /
@@ -454,7 +455,8 @@ const boardDialog = ref<{ code: string; name: string } | null>(null)
function openDialog(code: string, name: string, marketHint?: string) { function openDialog(code: string, name: string, marketHint?: string) {
if (!code) return if (!code) return
const mkt = marketHint ?? (/^(6|9|5)/.test(code) ? 'SH' : /^(4|8|92|43)/.test(code) ? 'BJ' : 'SZ') // 920xxx 9 9 SH
const mkt = marketHint ?? (/^(4|8|92|43)/.test(code) ? 'BJ' : /^[659]/.test(code) ? 'SH' : 'SZ')
dialog.value = { market: mkt, code, name } dialog.value = { market: mkt, code, name }
} }
+11 -1
View File
@@ -26,7 +26,10 @@ const viewMode = ref<'matrix' | 'corr'>('matrix')
watch( watch(
() => props.boardType, () => props.boardType,
(t) => { (t) => {
if (t && t !== boardType.value) setType(t) // /styles FG/hotspots propsundefined
// 退 HY FG
const target = t ?? 'HY'
if (target !== boardType.value) setType(target)
}, },
) )
@@ -48,6 +51,8 @@ const loading = ref(false)
const lastRefresh = ref('') const lastRefresh = ref('')
let buildTimer = 0 let buildTimer = 0
/** 请求序号守卫:await 后比对,过期响应直接丢弃(不覆盖新类型/新参数的状态)。 */
let loadSeq = 0
function stopBuildPoll() { function stopBuildPoll() {
if (buildTimer) { if (buildTimer) {
@@ -57,8 +62,12 @@ function stopBuildPoll() {
} }
async function load(retry = false) { async function load(retry = false) {
const my = ++loadSeq
try { try {
const r = await fetchBoardHotspot(boardType.value, days.value, mode.value, PER_DAY, retry) const r = await fetchBoardHotspot(boardType.value, days.value, mode.value, PER_DAY, retry)
// // resp
//
if (my !== loadSeq) return
if (r.status === 'building') { if (r.status === 'building') {
buildError.value = '' buildError.value = ''
resp.value = null resp.value = null
@@ -78,6 +87,7 @@ async function load(retry = false) {
loading.value = false loading.value = false
lastRefresh.value = new Date().toLocaleTimeString('zh-CN', { hour12: false }) lastRefresh.value = new Date().toLocaleTimeString('zh-CN', { hour12: false })
} catch (e) { } catch (e) {
if (my !== loadSeq) return
stopBuildPoll() stopBuildPoll()
buildingProgress.value = null buildingProgress.value = null
buildError.value = formatError(e) buildError.value = formatError(e)
+10 -3
View File
@@ -24,8 +24,11 @@ const bars = ref<Bar[]>([])
const lastUpdate = ref('') const lastUpdate = ref('')
async function loadIndex(idx: number) { async function loadIndex(idx: number) {
if (barsByIndex.has(idx)) { const cached = barsByIndex.get(idx)
bars.value = barsByIndex.get(idx)! if (cached) {
bars.value = cached
error.value = '' // /
loading.value = false
return return
} }
loading.value = true loading.value = true
@@ -33,9 +36,13 @@ async function loadIndex(idx: number) {
try { try {
const meta = INDICES[idx] const meta = INDICES[idx]
const data = await fetchIndexBars(meta.market, meta.code, 550) // 2.2 const data = await fetchIndexBars(meta.market, meta.code, 550) // 2.2
if (data.length === 0) {
// "" return
error.value = `${meta.name} 日K返回空`
return
}
barsByIndex.set(idx, data) barsByIndex.set(idx, data)
bars.value = data bars.value = data
if (data.length === 0) error.value = `${meta.name} 日K返回空`
lastUpdate.value = new Date().toLocaleTimeString('zh-CN', { hour12: false }) lastUpdate.value = new Date().toLocaleTimeString('zh-CN', { hour12: false })
} catch (e) { } catch (e) {
error.value = formatError(e) error.value = formatError(e)
+9 -4
View File
@@ -230,7 +230,8 @@ async function loadVolume() {
fetchBars('SH', '000001', 'MIN_5', start), fetchBars('SH', '000001', 'MIN_5', start),
fetchBars('SZ', '399001', 'MIN_5', start), fetchBars('SZ', '399001', 'MIN_5', start),
]) ])
// 5 amount // + 5 amount
// 5
const byDate = new Map<string, Map<string, number>>() const byDate = new Map<string, Map<string, number>>()
for (const b of [...sh, ...sz]) { for (const b of [...sh, ...sz]) {
const d = b.datetime.slice(0, 10) const d = b.datetime.slice(0, 10)
@@ -343,7 +344,7 @@ async function buildDigest(): Promise<string> {
) )
} }
if (volRatio.value !== null) { if (volRatio.value !== null) {
lines.push(`量能:当日两市累计成交较近 5 日同期均值 ${fmtPctSigned(volRatio.value)}`) lines.push(`量能:沪深指数(上证指数+深证成指)当日成交较近 5 日同期均值 ${fmtPctSigned(volRatio.value)}`)
} }
try { try {
const eco = await fetchLimitUpEcology() const eco = await fetchLimitUpEcology()
@@ -471,8 +472,9 @@ onBeforeUnmount(() => {
<!-- 量能仪表盘 --> <!-- 量能仪表盘 -->
<div class="section"> <div class="section">
<div class="sec-title"> <div class="sec-title">
量能 · 两市累计成交额最近交易日{{ volDate ? ` ${volDate.slice(5)}` : '' }} vs 5 日同期均值 量能 · 沪深指数成交额上证指数 + 深证成指 MIN_5 合计最近交易日{{ volDate ? ` ${volDate.slice(5)}` : '' }} vs 5 日同期均值
<span v-if="volRatio !== null" :class="volRatio > 0 ? 'up' : 'down'">{{ fmtPctSigned(volRatio) }}</span> <span v-if="volRatio !== null" :class="volRatio > 0 ? 'up' : 'down'">{{ fmtPctSigned(volRatio) }}</span>
<span class="dim">· 非全市场口径全市场总成交见顶部今日总成交</span>
</div> </div>
<div class="card chart-card"> <div class="card chart-card">
<div ref="volEl" class="chart"></div> <div ref="volEl" class="chart"></div>
@@ -486,7 +488,10 @@ onBeforeUnmount(() => {
<div v-for="d in fundDays" :key="d.date" class="fund-row"> <div v-for="d in fundDays" :key="d.date" class="fund-row">
<span class="mono dim fund-date">{{ String(d.date).slice(4, 6) }}-{{ String(d.date).slice(6, 8) }}</span> <span class="mono dim fund-date">{{ String(d.date).slice(4, 6) }}-{{ String(d.date).slice(6, 8) }}</span>
<span v-for="b in d.boards" :key="b.code" class="fund-chip mono"> <span v-for="b in d.boards" :key="b.code" class="fund-chip mono">
{{ b.name }} <span class="up">+{{ (b.main_net / 1e8).toFixed(1) }}亿</span> {{ b.name }}
<span :class="b.main_net >= 0 ? 'up' : 'down'">
{{ b.main_net >= 0 ? '+' : '-' }}{{ (Math.abs(b.main_net) / 1e8).toFixed(1) }}亿
</span>
</span> </span>
</div> </div>
<div v-if="fundDays.length === 0" class="empty-hint dim"> <div v-if="fundDays.length === 0" class="empty-hint dim">
+10 -5
View File
@@ -68,15 +68,20 @@ async function fillMissingNames() {
// SSE + REST // SSE + REST
/** SSE 未覆盖时(自选刚加、服务重启间隙)用 REST 主动拉一次。 */ /** SSE REST
* 后端 /quotes 单次最多 80 通达信协议上限超量需分批 */
const QUOTE_BATCH = 80
async function restFallback() { async function restFallback() {
if (items.value.length === 0) return if (items.value.length === 0) return
const missing = items.value.filter((i) => !quoteStore.getQuote(i.symbol)) const missing = items.value.filter((i) => !quoteStore.getQuote(i.symbol))
if (missing.length === 0) return if (missing.length === 0) return
try { for (let i = 0; i < missing.length; i += QUOTE_BATCH) {
await fetchQuotes(missing.map((i) => ({ market: i.market, code: i.code }))) try {
} catch { await fetchQuotes(missing.slice(i, i + QUOTE_BATCH).map((it) => ({ market: it.market, code: it.code })))
// SSE } catch {
// SSE
}
} }
} }