mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 14:24:15 +08:00
Compare commits
200
Commits
c7a921cc60
...
54ef03ac7d
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
54ef03ac7d | ||
|
|
ea4d8a8278 | ||
|
|
a90c12b3a5 | ||
|
|
fb7f7650d9 | ||
|
|
facf5f3c96 | ||
|
|
8304175ffc | ||
|
|
3625de9188 | ||
|
|
77dbf1a8b8 | ||
|
|
79eb150de2 | ||
|
|
77b829a096 | ||
|
|
2b059d4b43 | ||
|
|
e5f7f625fb | ||
|
|
e956e3a6d7 | ||
|
|
40c2468cbc | ||
|
|
7755ab3a7d | ||
|
|
bb94cd7b29 | ||
|
|
5c31616260 | ||
|
|
4ce7812565 | ||
|
|
cd41bf54d3 | ||
|
|
a8f314df63 | ||
|
|
afb442f432 | ||
|
|
87f185cdc3 | ||
|
|
8f71cd8732 | ||
|
|
c252cc328e | ||
|
|
ee4f0620b2 | ||
|
|
5d0f1cba79 | ||
|
|
3f9c6da7dd | ||
|
|
0da6aea52a | ||
|
|
a714d70f83 | ||
|
|
be54e11912 | ||
|
|
5da887b737 | ||
|
|
03f8966856 | ||
|
|
72a0f51a64 | ||
|
|
1e2c7afa93 | ||
|
|
cc48d51c20 | ||
|
|
d0a14b5c1b | ||
|
|
71cec0dc35 | ||
|
|
6345eb93ab | ||
|
|
01347b72bc | ||
|
|
b531475ca8 | ||
|
|
a5e725c738 | ||
|
|
8d8b66c7f8 | ||
|
|
3db7f38446 | ||
|
|
771a95a6d1 | ||
|
|
5ffe43d42d | ||
|
|
efd9f820d4 | ||
|
|
6eb1508fd8 | ||
|
|
7dc11125bf | ||
|
|
7e99a07270 | ||
|
|
fe1410532f | ||
|
|
5467b0dca4 | ||
|
|
464a5d9d50 | ||
|
|
99f478b257 | ||
|
|
970296e82e | ||
|
|
eebe533986 | ||
|
|
1b90f57501 | ||
|
|
68005c92e6 | ||
|
|
fd5d883c8f | ||
|
|
d2931aaadb | ||
|
|
252872e84f | ||
|
|
4bbc7d07a3 | ||
|
|
a49c8b3b15 | ||
|
|
7a13713903 | ||
|
|
ea27376e03 | ||
|
|
e349f2ba9b | ||
|
|
a0925cb0c8 | ||
|
|
2266ecc1d8 | ||
|
|
f325a9676c | ||
|
|
9a4bdcd07d | ||
|
|
85c903704a | ||
|
|
8cb834d6a1 | ||
|
|
cdd41a6d15 | ||
|
|
de250f72bc | ||
|
|
9a9903cf69 | ||
|
|
cda496ba99 | ||
|
|
a8bb7a25bf | ||
|
|
a5374d4212 | ||
|
|
908b385501 | ||
|
|
f8a23b7ce5 | ||
|
|
614e088079 | ||
|
|
cd1cff0089 | ||
|
|
f03bc38a16 | ||
|
|
a1ef095334 | ||
|
|
5618b4ef1d | ||
|
|
58b161b16d | ||
|
|
8ef3d66b95 | ||
|
|
28dbcb0709 | ||
|
|
6876db1a86 | ||
|
|
f2fac2e8f0 | ||
|
|
4a27dc45d5 | ||
|
|
6ee8e44ce4 | ||
|
|
8131cfe853 | ||
|
|
c39a9b8b17 | ||
|
|
957bdd2d25 | ||
|
|
ddb265b0f8 | ||
|
|
c4b8d44215 | ||
|
|
d949d9618c | ||
|
|
5805999024 | ||
|
|
471ee6763b | ||
|
|
af97df1349 | ||
|
|
ecbb45b8ef | ||
|
|
cd9e20610f | ||
|
|
537728a870 | ||
|
|
f275edafeb | ||
|
|
f0c0fd3bf8 | ||
|
|
c94fe2bf31 | ||
|
|
e3a2fbce13 | ||
|
|
716897f41a | ||
|
|
3d6beb3c35 | ||
|
|
9d55c237ed | ||
|
|
e6c7fa4a43 | ||
|
|
905fb8846c | ||
|
|
f6fc22fef1 | ||
|
|
1c022a1c7c | ||
|
|
91d2268532 | ||
|
|
aeb52b4a64 | ||
|
|
ae30b0cb85 | ||
|
|
900dbea59c | ||
|
|
2ebb9f0b11 | ||
|
|
c08e26ddd1 | ||
|
|
b5a005a8ac | ||
|
|
e9359c4fa7 | ||
|
|
c289ac5767 | ||
|
|
d46e601863 | ||
|
|
e12e0c7d02 | ||
|
|
c1ad4880fe | ||
|
|
b3e492d890 | ||
|
|
2ff909cd36 | ||
|
|
c1cad36449 | ||
|
|
e40bf15a50 | ||
|
|
e0cd625ef4 | ||
|
|
f6257408bf | ||
|
|
5289cde13e | ||
|
|
bab609b2d4 | ||
|
|
3fee4ecba5 | ||
|
|
831e2dc5db | ||
|
|
a9af2bce1c | ||
|
|
885e693118 | ||
|
|
00b3492414 | ||
|
|
8c28132361 | ||
|
|
2ce8b4b17d | ||
|
|
7ea741e4dc | ||
|
|
84f00d36d3 | ||
|
|
2140209506 | ||
|
|
4ca655279a | ||
|
|
4bb26713fc | ||
|
|
8721bd2e83 | ||
|
|
2e8be527b6 | ||
|
|
4cb30e48aa | ||
|
|
91fef6d793 | ||
|
|
ef59df4a67 | ||
|
|
e89ea9becf | ||
|
|
41205b197c | ||
|
|
8b12fc7a65 | ||
|
|
f8031e16f4 | ||
|
|
ba83e0241b | ||
|
|
adc7ab52a9 | ||
|
|
de5d5a9769 | ||
|
|
a28f9e852f | ||
|
|
3a3012f5b6 | ||
|
|
7ac0456c53 | ||
|
|
e9f5c606b6 | ||
|
|
7619fbcf7a | ||
|
|
40fea3236e | ||
|
|
0a11bb870e | ||
|
|
6964e1b5c9 | ||
|
|
341b80e7ab | ||
|
|
4d36a5ff20 | ||
|
|
f3b916c049 | ||
|
|
41a97888b6 | ||
|
|
7038c7d627 | ||
|
|
3258325daa | ||
|
|
0f7c1dfe39 | ||
|
|
ef66877529 | ||
|
|
ec95d71732 | ||
|
|
571e828529 | ||
|
|
bd84d4c0ae | ||
|
|
22add7179c | ||
|
|
27d4a5bbca | ||
|
|
7b754c0ff4 | ||
|
|
576850d0fa | ||
|
|
84725cd362 | ||
|
|
67b271285c | ||
|
|
bfe76e8273 | ||
|
|
d266131f90 | ||
|
|
e8c870af73 | ||
|
|
d0595eb4d7 | ||
|
|
a9c00fd567 | ||
|
|
ebf1a89254 | ||
|
|
d92ef17851 | ||
|
|
8aa6e8d89b | ||
|
|
64ab9a0818 | ||
|
|
d280dd59b2 | ||
|
|
2f7a4b0979 | ||
|
|
a36709ccbc | ||
|
|
250595c925 | ||
|
|
177fca7ae8 | ||
|
|
01f20ae4c3 | ||
|
|
51192321b2 | ||
|
|
e864ecde4b |
+6
-3
@@ -95,6 +95,9 @@
|
||||
|
||||
新增跨边界映射时必须增加单位测试。禁止用“数值小于 1 就乘 100”一类启发式转换,这会掩盖真实数据错误。
|
||||
|
||||
五档盘口的 `bid_volumes` / `ask_volumes` 沿用现有封单量口径,单位为“手”;前端计算
|
||||
封单额时再乘 `100` 换算为股。provider 必须在数据边界完成单位转换。
|
||||
|
||||
### 3.2 价格与复权
|
||||
|
||||
- enriched 的 `open/high/low/close` 为前复权价格。
|
||||
@@ -122,7 +125,7 @@
|
||||
数据源已经插件化。任何通用功能都必须通过 provider 能力和标准化数据集访问数据,不能把 TickFlow SDK 调用硬编码到策略、监控、回测、API 或前端流程中。
|
||||
|
||||
- 使用现有的 `get_provider()`、`provider_has_dataset()` 和 preferences 路由能力。
|
||||
- 支持的数据集包括但不限于 `daily`、`adj_factor`、`minute`、`full_minute`(盘中全市场分钟落盘)、`realtime`、`financial`;新增数据集应先定义清晰的输入输出契约。
|
||||
- 支持的数据集包括但不限于 `daily`、`adj_factor`、`minute`、`full_minute`(盘中全市场分钟落盘)、`realtime`、`depth5`、`financial`;新增数据集应先定义清晰的输入输出契约。
|
||||
- provider 负责把供应商字段、单位、日期和代码格式转换为内部标准格式。
|
||||
- 上层服务依赖标准字段和能力声明,不依赖供应商响应结构。
|
||||
- 只有明确标注为 TickFlow 专属的功能才可以直接依赖 TickFlow,并且不得影响其他 provider。
|
||||
@@ -138,9 +141,9 @@
|
||||
- 各页面能力门控统一以矩阵的 `usable` 为准(生效源当前能否真正提供该能力),不是 TickFlow 套餐视角;缺能力提示统一引导到数据源配置。
|
||||
- 能力层中立:通用界面(侧栏徽章、能力路由卡、各页门控提示)不得出现 TickFlow 档位/订阅词汇;档位信息只在 TickFlow 专属详情卡展示。provider 名称作为路由事实可以出现。
|
||||
- 每个能力独立路由,禁止跟随/派生特殊值(`same_as_daily` 已下线);存量非法偏好值由 preferences getter 回退默认自愈,不做迁移。
|
||||
- 边界注记:分时监控由分钟能力兜底(`intraday_monitor_support`),不单设分时能力;`full_minute`(全量分钟)数据集已开放插件/自定义源声明;`depth5` 已进矩阵但插件数据集白名单暂未开放,当前仅 TickFlow 提供。
|
||||
- 边界注记:分时监控由分钟能力兜底(`intraday_monitor_support`),不单设分时能力;`full_minute`(全量分钟)和 `depth5` 数据集均已开放插件声明。`depth5` 独立路由且失败时不跨源回退。
|
||||
- 实时指数为产品级固定契约,不走路由矩阵:展示层(侧栏指数条、市场总览)固定核心四只(`backend/app/services/index_const.py` 单一权威:上证/深成/创业板/科创综指),后端各消费方与前端 Layout 引用同一份定义不建副本;指数页保留但标的固定为核心四只(无全指数搜索/浏览,`/api/index/list`、`/api/index/search` 已下线);侧栏指数多选配置已下线,相关偏好(`realtime_index_symbols`/`sidebar_index_symbols`/`indices_nav_pinned`/`realtime_pull_index`/`realtime_index_mode`)已删除。监控规则的指数标的不受限——quote_service 把核心四只 + 启用规则的指数并入显式拉取。
|
||||
- 自定义源指数补充协议:A 股快照普遍不含指数(fuyao 实测无指数,指数在其独立端点)。provider 可实现可选方法 `get_realtime_indices(symbols) -> list[dict]`(record 结构与 realtime 一致),quote_service 在自定义源分支鸭子类型调用补拉;未实现的源指数缓存为空,由本地日K兜底接管。fuyao 指数快照有连坐语义——请求混入未知代码整批失败,插件侧必须先行过滤不支持的后缀(如 `.BJ`)。
|
||||
- 自定义源指数补充协议:A 股快照普遍不含指数(fuyao 实测无指数,指数在其独立端点)。provider 可实现可选方法 `get_realtime_indices(symbols) -> list[dict] | None`(record 结构与 realtime 一致),quote_service 在自定义源分支鸭子类型调用补拉;`None` 表示请求失败,保留上轮有效指数缓存,`[]` 表示成功但无数据;未实现的源指数缓存为空,由本地日K兜底接管。fuyao 指数快照有连坐语义——请求混入未知代码整批失败,插件侧必须先行过滤不支持的后缀(如 `.BJ`)。
|
||||
|
||||
## 5. 领域专项要求
|
||||
|
||||
|
||||
+14
-2
@@ -75,9 +75,14 @@ WORKDIR /app
|
||||
# Codex CLI 从官方 npm 包提取原生二进制,不依赖运行时 Node.js。
|
||||
# bookworm 自带 nodejs 18.19, 满足插件 engines>=18; --no-install-recommends 精简,
|
||||
# 自带 libnode/libc-ares 等全部动态依赖, 无需手动补库。
|
||||
# 国内构建走 apt mirror 已在 debian 镜像sources.list 配好, 无需额外换源。
|
||||
# 国内构建换 apt 阿里云镜像 (bookworm deb822 格式)。官方 python slim 镜像的
|
||||
# sources 指向 deb.debian.org, 部分国内网络下 apt 极慢(实测阿里云 ECS 拉官方源
|
||||
# 单阶段可达 8 分钟以上); USE_CN_MIRROR 与 npm/pypi 的换源开关注一脉相承。
|
||||
# tesseract-ocr: 自选截图导入(始终安装); nodejs: 仅 INCLUDE_STOCKSDK=1 时安装
|
||||
RUN apt-get update \
|
||||
RUN if [ "$USE_CN_MIRROR" = "1" ]; then \
|
||||
sed -i 's|deb.debian.org|mirrors.aliyun.com|g' /etc/apt/sources.list.d/debian.sources /etc/apt/sources.list 2>/dev/null || true; \
|
||||
fi \
|
||||
&& apt-get update \
|
||||
&& apt-get install -y --no-install-recommends tesseract-ocr tesseract-ocr-eng \
|
||||
&& if [ "$INCLUDE_STOCKSDK" = "1" ]; then \
|
||||
apt-get install -y --no-install-recommends nodejs \
|
||||
@@ -136,6 +141,13 @@ COPY --from=codex-builder /opt/codex-native /usr/local/bin/codex
|
||||
RUN codex --version
|
||||
|
||||
ENV PYTHONPATH=/app
|
||||
# 运行时 uv 镜像源持久化: CMD 用 `uv run` 启动, 锁与 pyproject 不一致等场景下
|
||||
# uv 会在容器内重新解析/安装 —— 无源配置时默认 pypi.org, 国内网络会卡死启动
|
||||
# (实测阿里云 ECS)。与构建期 RUN 内的 export 同源, 这里让它跨层存活。
|
||||
ARG PYPI_INDEX=https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
ARG PYPI_FALLBACK=https://mirrors.aliyun.com/pypi/simple
|
||||
ENV UV_DEFAULT_INDEX=${PYPI_INDEX} \
|
||||
UV_EXTRA_INDEX_URL=${PYPI_FALLBACK}
|
||||
# 兜底时区: 交易时段判断已在代码里显式用北京时间 (app/market_time.py),
|
||||
# 此处让日志时间戳等其余 naive 时间也对齐北京时间。
|
||||
ENV TZ=Asia/Shanghai
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2026 tickflow-stock-panel contributors
|
||||
Copyright (c) 2026 tick-stock-panel contributors
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
|
||||
@@ -32,7 +32,7 @@
|
||||
>
|
||||
> **明确不做**:不对标同花顺 / 通达信,不内置「AI 荐股 / 涨停预测」。
|
||||
|
||||
有任何项目问题或商务合作 / 广告投放等合作意向,可邮件联系 415333856@qq.com。
|
||||
有任何项目问题可邮件联系 415333856@qq.com。
|
||||
|
||||
觉得有用可以点个 Star
|
||||
|
||||
@@ -40,19 +40,77 @@
|
||||
|
||||
## ✨ 核心功能
|
||||
|
||||
| 模块 | 一句话 | 详见 |
|
||||
| :--------------- | :--------------------------------------------------------------------- | :-------------------------------- |
|
||||
| 🔀 **能力路由** | 多数据集(日K/除权/实时/分钟/盘口/财务,持续扩展)按源能力独立路由,任选组合 | [custom-data-source.md](./docs/custom-data-source.md) |
|
||||
| 🔍 **选股引擎** | 25 个内置策略 + 分钟策略 + 自定义信号 + AI 生成,Polars 毫秒级扫全 A 股 | [strategy.md](./docs/strategy.md) |
|
||||
| 📊 **指标流水线** | MA/EMA/MACD/RSI/KDJ/布林/量比等 68 列指标与信号,一次扫表落盘 enriched Parquet | [features.md](./docs/features.md) |
|
||||
| 🧪 **回测研究** | 因子/策略/分钟回测 + 财务快照因子(点时口径),T+1/费用/滑点约束,结果可导出 | [features.md](./docs/features.md) |
|
||||
| ⛏️ **因子挖掘** | 嵌套样本外搜索多因子排名组合,与自有策略对照,候选库显式发布、永不自动上线 | [mining.md](./docs/mining.md) |
|
||||
| 🌡️ **市场环境** | 情绪周期 6 阶段(连板梯队驱动)+ 概念/行业主线排名,与 5 档环境分并存 | [market-phase.md](./docs/market-phase.md) |
|
||||
| 🚨 **异动监控** | 竞价/盘中/偏移三类异动一页覆盖:同花顺风向标 + 当日信号聚合 + 交易所偏离值口径 | — |
|
||||
| 📡 **监控中心** | 四类监控(策略/个股信号/价格/异动),多条件 AND/OR + 语音播报 + 飞书推送 | [features.md](./docs/features.md) |
|
||||
| 📈 **个股分析** | 9 类关键价位 + AI 四维分析(技术/基本面/财务/消息面) | [features.md](./docs/features.md) |
|
||||
| 🏆 **连板梯队** | 连板层级统计 + 概念涨幅轮动 + 盘后 AI 复盘(龙虎榜/盘前风向标注入) + 炸板/翘板预警 | [features.md](./docs/features.md) |
|
||||
| 🧰 **数据扩展** | 数据源插件化(TickFlow/fuyao/stock-sdk + YAML 自定义源),扩展字段配成一级页面同台分析 | [custom-data-source.md](./docs/custom-data-source.md) |
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th nowrap align="left">模块</th>
|
||||
<th align="left">一句话</th>
|
||||
<th nowrap align="left">详见</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr>
|
||||
<td nowrap>🔀 <b>能力路由</b></td>
|
||||
<td>多数据集(日K/除权/实时/分钟/盘口/财务,持续扩展)按源能力独立路由,任选组合</td>
|
||||
<td nowrap><a href="./docs/custom-data-source.md">custom-data-source.md</a></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>🔍 <b>选股引擎</b></td>
|
||||
<td>25 个内置策略 + 分钟策略 + 自定义信号 + AI 生成,Polars 毫秒级扫全 A 股</td>
|
||||
<td nowrap><a href="./docs/strategy.md">strategy.md</a></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>📊 <b>指标流水线</b></td>
|
||||
<td>MA/EMA/MACD/RSI/KDJ/布林/量比等 68 列指标与信号,一次扫表落盘 enriched Parquet</td>
|
||||
<td nowrap><a href="./docs/features.md">features.md</a></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>🧪 <b>回测研究</b></td>
|
||||
<td>因子/策略/分钟回测 + 财务快照因子(点时口径),T+1/费用/滑点约束,评分策略附带因子归因</td>
|
||||
<td nowrap><a href="./docs/features.md">features.md</a></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>🔬 <b>因子平台</b></td>
|
||||
<td>DSL 自定义因子(编辑器 25 算子点选/试算/版本) + 检验/组合,与策略双向联动(一键生成策略/触发器引用因子/回测归因)</td>
|
||||
<td nowrap><a href="./docs/factor-platform-plan.md">factor-platform-plan.md</a></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>⛏️ <b>因子挖掘</b></td>
|
||||
<td>嵌套样本外搜索多因子排名组合,与自有策略对照,候选库显式发布、永不自动上线</td>
|
||||
<td nowrap><a href="./docs/mining.md">mining.md</a></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>🌡️ <b>市场环境</b></td>
|
||||
<td>情绪周期 6 阶段(连板梯队驱动)+ 概念/行业主线排名,与 5 档环境分并存</td>
|
||||
<td nowrap><a href="./docs/market-phase.md">market-phase.md</a></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>🚨 <b>异动监控</b></td>
|
||||
<td>竞价/盘中/偏移三类异动一页覆盖:同花顺风向标 + 当日信号聚合 + 交易所偏离值口径</td>
|
||||
<td nowrap>—</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>📡 <b>监控中心</b></td>
|
||||
<td>四类监控(策略/个股信号/价格/异动),多条件 AND/OR + 语音播报 + 飞书推送</td>
|
||||
<td nowrap><a href="./docs/features.md">features.md</a></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>📈 <b>个股分析</b></td>
|
||||
<td>9 类关键价位 + AI 四维分析(技术/基本面/财务/消息面)</td>
|
||||
<td nowrap><a href="./docs/features.md">features.md</a></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>🏆 <b>连板梯队</b></td>
|
||||
<td>连板层级统计 + 概念涨幅轮动 + 盘后 AI 复盘(龙虎榜/盘前风向标注入) + 炸板/翘板预警</td>
|
||||
<td nowrap><a href="./docs/features.md">features.md</a></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td nowrap>🧰 <b>数据扩展</b></td>
|
||||
<td>数据源插件化(TickFlow/fuyao/stock-sdk + YAML 自定义源),扩展字段配成一级页面同台分析;时序表(如人气排行)支持按日历史回补,接口配日期参数即可逐日补齐</td>
|
||||
<td nowrap><a href="./docs/custom-data-source.md">custom-data-source.md</a></td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
<details>
|
||||
<summary><b>📦 主要页面与功能</b></summary>
|
||||
@@ -66,10 +124,11 @@
|
||||
- **策略** Screener — Polars 毫秒级扫描全 A 股,日线/分钟策略统一单池,按策略声明周期自动路由执行
|
||||
- **回测** Backtest — 四种研究视图:
|
||||
- **因子回测** — IC/IR、分层收益、多空组合,62+ 因子目录先筛掉无效指标
|
||||
- **策略回测** — 净值曲线、回撤、夏普、胜率,T+1/手续费/滑点/止损,SSE 流式进度
|
||||
- **策略回测** — 净值曲线、回撤、夏普、胜率、盈亏比、蒙卡回撤,T+1/手续费/滑点/止损,SSE 流式进度;评分因子策略附带「因子归因」(胜/败单入场信号日因子对比)
|
||||
- **分钟策略回测** — 逐交易日回放信号、分钟收盘入场,分钟级成交明细
|
||||
- **验证** — 参数敏感性与滚动样本外
|
||||
- 研究闭环:结果导出 CSV(概要/净值/交易明细/分标的统计) → 保存候选 → **一键载入复测**
|
||||
- **因子** Factors — 检验/因子库/编辑器/组合四 tab:IC·分层·Newey-West 检验、自定义 DSL 因子(25 算子点选、双语字段、我的因子模板)、版本与生命周期管理;因子库可**一键生成排名策略**,策略触发器可直接引用因子条件
|
||||
- **挖掘** Mining — 嵌套样本外因子与策略挖掘:训练区间因子方向重估 + 相关性去重 + 多因子排名组合搜索,自有策略作对照轨;候选入库,显式确认后才发布,永不自动上线
|
||||
|
||||
**📈 个股与板块分析**
|
||||
@@ -81,6 +140,8 @@
|
||||
|
||||
**🔔 监控与复盘**
|
||||
- **监控中心** Monitor — 策略/个股信号/价格/异动四类规则,支持自选分组作用域,盘中实时弹窗 + 语音播报(播报个股名称与信号) + 触发记录持久化
|
||||
- **持仓提醒** Lots — 记录个股/ETF 买入批次,自动生成止盈止损/到期监控规则
|
||||
- **信号库** Signals — 内置预计算信号 + 自定义条件信号(含因子条件与 AI 生成),供策略触发器/回测/监控统一取用
|
||||
- **异动监控** Abnormal Moves — 按交易时间线三 tab:
|
||||
- **竞价异动** — 同花顺盘前风向标(含当日/次日真实收益对照、追高风险标记)+ 全市场竞价扫描(待采集任务)
|
||||
- **盘中异动** — 涨停/炸板/翘板/跌停/新高/新低/放量当日信号聚合,零新增采集
|
||||
@@ -90,7 +151,7 @@
|
||||
**🗄️ 数据与扩展**
|
||||
- **数据** Data — 本地数据画像与同步状态(维表/日K/除权/Enriched/指数/ETF/分钟K/财务),盘后管道与历史扩展
|
||||
- **扩展分析** (动态菜单) — 把任意第三方/扩展数据字段配成一级菜单,与内置数据同台分析
|
||||
- **设置** Settings — 数据源与能力检测(能力路由矩阵、档位徽章)、AI 接口、实时监控、扩展页面、信号库、菜单与系统设置
|
||||
- **设置** Settings — 数据源与能力检测(能力路由矩阵、档位徽章)、AI 接口、实时监控、扩展页面、菜单与系统设置
|
||||
|
||||
</details>
|
||||
|
||||
@@ -252,18 +313,25 @@ flowchart TB
|
||||
|
||||
## 🚀 快速开始
|
||||
|
||||
> 前置依赖:Python ≥ 3.11 · Node ≥ 20 · [`uv`](https://docs.astral.sh/uv/) · `pnpm`(`npm i -g pnpm`)
|
||||
> 前置依赖(仅方式 D 需要):Python ≥ 3.11 · Node ≥ 20 · [`uv`](https://docs.astral.sh/uv/) · `pnpm`(`npm i -g pnpm`)
|
||||
>
|
||||
> 有 Docker 直接看 **方式 A**,一条命令拉现成镜像;完全不想碰命令行:看 **方式 C**,让本机 AI 帮你部署
|
||||
|
||||
### 方式 A:Dev 模式(二次开发推荐)
|
||||
### 方式 A:GHCR 现成镜像(免本地构建,多数用户推荐)
|
||||
|
||||
本项目每次推送都由 GitHub Actions 自动构建多架构镜像(linux/amd64 · arm64)并发布到 GHCR,拿来即用,本地无需装 Python / Node,也不用现场 build:
|
||||
|
||||
```bash
|
||||
cp .env.example .env # 按需填 TICKFLOW_API_KEY(留空 = None 模式)
|
||||
./dev.sh # Windows: .\dev.ps1
|
||||
docker run -d --name tsp -p 3018:3018 -v ${PWD}/data:/app/data ghcr.io/shy3130/tick-stock-panel:latest
|
||||
# 打开 http://localhost:3018
|
||||
```
|
||||
|
||||
自动检查 / 下载依赖、释放端口、同时起前后端。后端 → <http://localhost:3018> · 前端 → <http://localhost:3011>。
|
||||
- 需要配置时:从 [.env.example](./.env.example) 复制出 `.env`,命令里加 `--env-file .env`。
|
||||
- 跑自己改过的代码:fork 后到仓库 **Actions** 页启用 workflow(fork 默认禁用),构建出的 `ghcr.io/<你的用户名>/tick-stock-panel` 用法相同。
|
||||
- 想用 compose 编排(挂载 `.env` / `tiers.yaml`):参考 [docker-compose.yml](./docker-compose.yml),把 `build:` 段换成 `image: ghcr.io/shy3130/tick-stock-panel:latest`。
|
||||
- 现成镜像默认不含 stock-sdk 插件与老 CPU 兼容内核(合规与体积考虑),有此需求请用方式 B 自构建,详见 [docs/deployment.md](./docs/deployment.md)。
|
||||
|
||||
### 方式 B:Docker(部署最省心)
|
||||
### 方式 B:Docker Compose(本地构建,全套挂载)
|
||||
|
||||
```bash
|
||||
cp .env.example .env
|
||||
@@ -285,9 +353,31 @@ CODEX_CLI_VERSION=0.144.3 docker compose up --build
|
||||
|
||||
> Codex CLI 模式允许 TickFlow 容器读取本机 Codex 登录凭据,仅应在受信任的本机环境启用。凭据目录以只读方式挂载,不会写入镜像。
|
||||
|
||||
镜像已内置 **stock-sdk** 数据源插件(Node 运行时 + 依赖),开箱即用。
|
||||
镜像默认**不含** stock-sdk 插件(合规考虑);确需启用执行 `docker compose build --build-arg INCLUDE_STOCKSDK=1` 后再 `docker compose up -d`,详见 [docs/deployment.md](./docs/deployment.md)。
|
||||
|
||||
> 📖 Docker 进阶、GitHub Actions 自构建、老 CPU 兼容、访问密码设置等见 [docs/deployment.md](./docs/deployment.md)。
|
||||
> 📖 Docker 进阶、老 CPU 兼容、访问密码设置等见 [docs/deployment.md](./docs/deployment.md)。
|
||||
|
||||
### 方式 C:本机 AI 代部署(AI玩家首选)
|
||||
|
||||
装一个本机 AI 编程助手(Trae / Codex / OpenCode / ZCode / WorkBuddy 等,任选其一),新建一个空文件夹用助手打开,把下面这段话原样发给它:
|
||||
|
||||
```text
|
||||
帮我部署开源项目 https://github.com/shy3130/tick-stock-panel 到本机:
|
||||
克隆到当前文件夹;有 Docker 优先拉 ghcr.io/shy3130/tick-stock-panel:latest 现成镜像,没有就走 Dev 模式;
|
||||
缺少的依赖(Docker / Python / Node)帮我一起装好;
|
||||
最后告诉我浏览器打开哪个地址、需要填哪些 Key。
|
||||
```
|
||||
|
||||
AI 会自动完成克隆、装依赖、启动服务,完成后浏览器打开 <http://localhost:3018> 即可;`TICKFLOW_API_KEY` 等配置按 AI 提示填,详见 [配置](#️-配置)。
|
||||
|
||||
### 方式 D:Dev 模式(二次开发推荐)
|
||||
|
||||
```bash
|
||||
cp .env.example .env # 按需填 TICKFLOW_API_KEY(留空 = None 模式)
|
||||
./dev.sh # Windows: .\dev.ps1
|
||||
```
|
||||
|
||||
自动检查 / 下载依赖、释放端口、同时起前后端。后端 → <http://localhost:3018> · 前端 → <http://localhost:3011>。
|
||||
|
||||
### 跑起来后的第一次使用
|
||||
|
||||
@@ -336,6 +426,7 @@ PORT=3018 # 服务端口
|
||||
| [docs/features.md](./docs/features.md) | 各功能模块详细说明(选股/指标/回测/监控/个股分析/数据扩展) |
|
||||
| [docs/custom-data-source.md](./docs/custom-data-source.md) | 自定义数据源接入、能力路由契约、YAML 配置与 mock 联调示例 |
|
||||
| [docs/strategy.md](./docs/strategy.md) | 策略体系(25 内置策略 + 三种扩展方式 + 文件结构) |
|
||||
| [docs/strategy-iteration.md](./docs/strategy-iteration.md) | AI 策略迭代协议:台账 / 证据包 / 门槛判定 / 提示词卡片 |
|
||||
| [docs/mining.md](./docs/mining.md) | 因子与策略挖掘口径、防泄漏、任务隔离和发布边界 |
|
||||
| [docs/market-phase.md](./docs/market-phase.md) | 市场情绪周期 6 阶段与概念/行业主线识别的口径与设计 |
|
||||
| [docs/plugin-development.md](./docs/plugin-development.md) | 数据源插件开发规范(以 stock-sdk / fuyao 为参考实现) |
|
||||
@@ -346,37 +437,32 @@ fork同时请点个star哦,欢迎 Issue 和 PR。
|
||||
|
||||
---
|
||||
|
||||
## 💬 交流群
|
||||
|
||||
欢迎加入交流群,一起讨论交流。作者个人维护的部分个性化接口,统一公布在群公告中,供大家免费使用。
|
||||
|
||||
<img src="./community-qr-code.jpg" alt="交流群二维码" width="240" />
|
||||
|
||||
---
|
||||
|
||||
## ❤️ 支持项目
|
||||
## ❤️ 支持项目 / 💬 交流群
|
||||
|
||||
<div align="center">
|
||||
|
||||
如果这个项目对你有帮助,欢迎请作者喝杯咖啡 ☕
|
||||
|
||||
<table>
|
||||
<tr>
|
||||
<td width="50%" align="center"><b>微信赞赏</b></td>
|
||||
<td width="50%" align="center"><b>支付宝</b></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td width="50%" align="center"><img src="./assets/support/wechat-appreciation.jpg" alt="微信赞赏码 · 感谢道友支持 愿一路长红" height="280" /></td>
|
||||
<td width="50%" align="center"><img src="./assets/support/alipay.jpg" alt="支付宝收款码 · 打开支付宝扫一扫" height="280" /></td>
|
||||
<td width="50%" align="center">
|
||||
<b>❤️ 支持项目</b><br/>
|
||||
<sub>如果这个项目对你有帮助,欢迎请作者喝杯咖啡 ☕</sub>
|
||||
<table>
|
||||
<tr><td align="center"><img src="./assets/support/wechat-appreciation.jpg" alt="微信赞赏码 · 感谢道友支持 愿一路长红" height="280" /></td></tr>
|
||||
<tr><td align="center"><sub>愿道友一路长红 📈</sub></td></tr>
|
||||
</table>
|
||||
</td>
|
||||
<td width="50%" align="center">
|
||||
<b>💬 交流群</b><br/>
|
||||
<sub>欢迎加入交流群,一起讨论交流<br/>个性化接口统一公布在群公告,免费使用</sub>
|
||||
<table>
|
||||
<tr><td align="center"><img src="./community-qr-code.jpg" alt="交流群二维码 · 个人维护的个性化接口见群公告" height="280" /></td></tr>
|
||||
</table>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
愿道友一路长红 📈
|
||||
|
||||
</div>
|
||||
|
||||
> 打赏完全自愿,金额不限;不用于购买任何功能、数据权限、投资建议
|
||||
>
|
||||
> 作者精力有限,优先响应赞助回馈,希望理解
|
||||
|
||||
---
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
import sys
|
||||
|
||||
__version__ = "0.2.2"
|
||||
__version__ = "0.2.3"
|
||||
|
||||
# Windows 默认 stdout/stderr 编码为 GBK(cp936),TickFlow SDK 内部输出含 emoji 的
|
||||
# 指数/标的名称(如 \U0001f193)时会抛 UnicodeEncodeError,导致请求失败。
|
||||
|
||||
@@ -5,7 +5,7 @@ import random
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from fastapi import APIRouter, HTTPException, Query, Request
|
||||
|
||||
from app.services import alert_store
|
||||
|
||||
@@ -19,8 +19,10 @@ def _data_dir(request: Request) -> Path:
|
||||
@router.get("")
|
||||
def list_alerts(
|
||||
request: Request,
|
||||
days: int = 7,
|
||||
limit: int = 5000,
|
||||
# 上限取存储侧保留策略 (alert_store.MAX_DAYS / MAX_RECORDS): 超出也没有可返回的记录。
|
||||
# 无下限时 days=-1 会把 cutoff 推到未来直接返回空, limit=-1 会走负数切片静默丢掉最旧一条。
|
||||
days: int = Query(alert_store.MAX_DAYS, ge=1, le=alert_store.MAX_DAYS),
|
||||
limit: int = Query(alert_store.MAX_RECORDS, ge=1, le=alert_store.MAX_RECORDS),
|
||||
source: str | None = None,
|
||||
type: str | None = None,
|
||||
ext_columns: str | None = None,
|
||||
|
||||
+29
-22
@@ -128,9 +128,9 @@ class FactorColumnsResponse(BaseModel):
|
||||
|
||||
@router.get("/factor/columns")
|
||||
def factor_columns():
|
||||
"""返回可用的因子列列表。"""
|
||||
from app.backtest.factor import FACTOR_COLUMNS
|
||||
return {"columns": FACTOR_COLUMNS}
|
||||
"""返回可用的因子列列表 (含运行期注册的自定义/复合因子)。"""
|
||||
from app.factors.registry import factor_columns_view
|
||||
return {"columns": factor_columns_view()}
|
||||
|
||||
|
||||
class FactorBacktestRequest(BaseModel):
|
||||
@@ -149,16 +149,17 @@ class FactorBacktestRequest(BaseModel):
|
||||
@router.post("/factor/run")
|
||||
def factor_run(req: FactorBacktestRequest, request: Request):
|
||||
"""因子回测 — IC/IR 分析 + 分层回测。"""
|
||||
from app.backtest.factor import FACTOR_COLUMNS, FactorBacktestService, FactorConfig
|
||||
from app.backtest.factor import FactorBacktestService, FactorConfig
|
||||
from app.factors.registry import factor_columns_view
|
||||
|
||||
if req.factor_name not in {item["id"] for item in FACTOR_COLUMNS}:
|
||||
if req.factor_name not in {item["id"] for item in factor_columns_view()}:
|
||||
raise HTTPException(status_code=400, detail=f"不支持的因子: {req.factor_name}")
|
||||
|
||||
engine = _get_engine(request)
|
||||
svc = FactorBacktestService(engine)
|
||||
|
||||
end = req.end or date.today()
|
||||
start = _resolve_start(req, end, STRATEGY_DEFAULT_DAYS)
|
||||
start = _resolve_start(req, end, FACTOR_DEFAULT_DAYS)
|
||||
_guard_server_backtest_range(start, end)
|
||||
symbols = req.symbols if req.symbols else None
|
||||
if symbols is not None and len(symbols) > FACTOR_MAX_SYMBOLS:
|
||||
@@ -184,7 +185,7 @@ def factor_run(req: FactorBacktestRequest, request: Request):
|
||||
|
||||
|
||||
class FactorBatchRequest(BaseModel):
|
||||
factor_names: list[str] = Field(..., min_length=1, max_length=64)
|
||||
factor_names: list[str] = Field(..., min_length=1, max_length=96) # 目录 77 + 自定义余量
|
||||
symbols: list[str] | None = None
|
||||
start: date | None = None
|
||||
end: date | None = None
|
||||
@@ -200,19 +201,19 @@ class FactorBatchRequest(BaseModel):
|
||||
def factor_batch(req: FactorBatchRequest, request: Request):
|
||||
"""批量筛选因子, 同一批次只加载并计算一次数据面板。"""
|
||||
from app.backtest.factor import (
|
||||
FACTOR_COLUMNS,
|
||||
FactorBacktestService,
|
||||
FactorBatchConfig,
|
||||
)
|
||||
from app.factors.registry import factor_columns_view
|
||||
|
||||
factor_names = list(dict.fromkeys(req.factor_names))
|
||||
allowed = {item["id"] for item in FACTOR_COLUMNS}
|
||||
allowed = {item["id"] for item in factor_columns_view()}
|
||||
invalid = [name for name in factor_names if name not in allowed]
|
||||
if invalid:
|
||||
raise HTTPException(status_code=400, detail=f"不支持的因子: {', '.join(invalid)}")
|
||||
|
||||
end = req.end or date.today()
|
||||
start = _resolve_start(req, end, STRATEGY_DEFAULT_DAYS)
|
||||
start = _resolve_start(req, end, FACTOR_DEFAULT_DAYS)
|
||||
_guard_server_backtest_range(start, end)
|
||||
symbols = req.symbols if req.symbols else None
|
||||
if symbols is not None and len(symbols) > FACTOR_MAX_SYMBOLS:
|
||||
@@ -526,10 +527,12 @@ async def strategy_stream(
|
||||
from app.backtest.strategy import StrategyBacktestConfig
|
||||
from app.backtest.worker import make_worker_task, run_worker_task
|
||||
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
if start:
|
||||
start_date = date.fromisoformat(start)
|
||||
else:
|
||||
try:
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
start_date = date.fromisoformat(start) if start else None
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=f"日期格式错误: {e}") from e
|
||||
if start_date is None:
|
||||
# 空 start = 全部历史: 用本地最早日K日期, 查不到再回退到默认窗口
|
||||
earliest = request.app.state.repo.earliest_daily_date()
|
||||
start_date = earliest or (end_date - timedelta(days=FACTOR_DEFAULT_DAYS))
|
||||
@@ -829,10 +832,12 @@ async def optimize_stream(
|
||||
from app.backtest.optimizer import OptimizeConfig
|
||||
from app.backtest.worker import make_worker_task, run_worker_task
|
||||
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
if start:
|
||||
start_date = date.fromisoformat(start)
|
||||
else:
|
||||
try:
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
start_date = date.fromisoformat(start) if start else None
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=f"日期格式错误: {e}") from e
|
||||
if start_date is None:
|
||||
earliest = request.app.state.repo.earliest_daily_date()
|
||||
start_date = earliest or (end_date - timedelta(days=FACTOR_DEFAULT_DAYS))
|
||||
|
||||
@@ -1051,10 +1056,12 @@ async def walkforward_stream(
|
||||
|
||||
direction = direction or None
|
||||
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
if start:
|
||||
start_date = date.fromisoformat(start)
|
||||
else:
|
||||
try:
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
start_date = date.fromisoformat(start) if start else None
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=f"日期格式错误: {e}") from e
|
||||
if start_date is None:
|
||||
earliest = request.app.state.repo.earliest_daily_date()
|
||||
start_date = earliest or (end_date - timedelta(days=STRATEGY_DEFAULT_DAYS))
|
||||
|
||||
|
||||
+158
-29
@@ -16,6 +16,7 @@ import polars as pl
|
||||
from fastapi import APIRouter, File, HTTPException, Query, Request, UploadFile
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.market_time import CN_TZ
|
||||
from app.services.ext_data import (
|
||||
ExtConfig,
|
||||
ExtConfigStore,
|
||||
@@ -24,13 +25,15 @@ from app.services.ext_data import (
|
||||
apply_config_mapping,
|
||||
detect_symbol_candidates,
|
||||
ensure_utf8_csv,
|
||||
ext_api_key_field,
|
||||
fix_symbol_format,
|
||||
get_ext_api_key,
|
||||
infer_fields_from_df,
|
||||
parse_upload_file,
|
||||
write_ext_parquet,
|
||||
rows_to_parquet,
|
||||
)
|
||||
from app.services.ext_pull import fetch_and_ingest, pull_scheduler
|
||||
from app.services.ext_pull import _request_json, fetch_and_ingest, pull_scheduler
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/api/ext-data", tags=["ext-data"])
|
||||
@@ -70,6 +73,16 @@ class IngestReq(BaseModel):
|
||||
rows: list[dict] = Field(..., min_length=1)
|
||||
|
||||
|
||||
class PullAuthReq(BaseModel):
|
||||
"""拉取接口鉴权方式 (与自定义行情源 AuthConfig 同口径)。
|
||||
|
||||
Key 本体存 secrets_store (secrets.json), 不写入 config.json。
|
||||
"""
|
||||
type: Literal["none", "bearer", "header", "query"] = "none"
|
||||
header: str = Field("Authorization", min_length=1, max_length=64) # bearer/header 用
|
||||
param: str = Field("token", min_length=1, max_length=64) # query 用
|
||||
|
||||
|
||||
class PullConfigReq(BaseModel):
|
||||
"""定时拉取配置请求。"""
|
||||
url: str = Field(..., min_length=1)
|
||||
@@ -82,6 +95,15 @@ class PullConfigReq(BaseModel):
|
||||
enabled: bool = False
|
||||
time_window_start: str | None = None # "HH:MM", None=不限
|
||||
time_window_end: str | None = None # "HH:MM", None=不限
|
||||
# 接口按日查询的参数名 (如 "date"): 配置后支持历史回补, 且当日拉取也带日期参数
|
||||
date_param: str | None = Field(None, min_length=1, max_length=16, pattern=r"^[A-Za-z_][A-Za-z0-9_]*$")
|
||||
# 鉴权方式; 请求中缺省 (None) = 保留现有配置, {"type":"none"} = 关闭鉴权
|
||||
auth: PullAuthReq | None = None
|
||||
|
||||
|
||||
class ApiKeyReq(BaseModel):
|
||||
"""设置拉取接口 API Key; 空串 = 清除。"""
|
||||
key: str = Field(..., max_length=4096)
|
||||
|
||||
|
||||
class DetectUrlReq(BaseModel):
|
||||
@@ -182,6 +204,20 @@ def _safe_json_value(value):
|
||||
return value
|
||||
|
||||
|
||||
def _partition_date(raw: str) -> str:
|
||||
"""把 `date` 入参规范成 `YYYY-MM-DD` 分区名。
|
||||
|
||||
这个值直接拼进分区目录名 (`timeseries/date=<value>`), 所以非法值不只是格式问题:
|
||||
`date=x/../../../../kline_daily` 会让读取路径离开 `ext_data/<id>/timeseries/`。
|
||||
同一文件的 `/sync`、`/ingest`、`/backfill` 都先 `date.fromisoformat` 再用, 只有
|
||||
`/rows` 和 `/dimension-members` 走的这条路把原始字符串直接拼进了路径。
|
||||
"""
|
||||
try:
|
||||
return date.fromisoformat(raw).isoformat()
|
||||
except ValueError as e:
|
||||
raise HTTPException(400, f"日期格式错误: {raw}") from e
|
||||
|
||||
|
||||
def _read_ext_dataframe(
|
||||
config: ExtConfig,
|
||||
data_dir: Path,
|
||||
@@ -200,10 +236,11 @@ def _read_ext_dataframe(
|
||||
return pl.DataFrame(), None
|
||||
|
||||
if snapshot_date:
|
||||
path = base / f"date={snapshot_date}" / "part.parquet"
|
||||
day = _partition_date(snapshot_date)
|
||||
path = base / f"date={day}" / "part.parquet"
|
||||
if not path.exists():
|
||||
return pl.DataFrame(), snapshot_date
|
||||
return pl.read_parquet(path), snapshot_date
|
||||
return pl.DataFrame(), day
|
||||
return pl.read_parquet(path), day
|
||||
|
||||
partitions = sorted(
|
||||
d for d in base.iterdir()
|
||||
@@ -234,10 +271,14 @@ def _with_instrument_name(df: pl.DataFrame, data_dir: Path) -> pl.DataFrame:
|
||||
|
||||
|
||||
def _latest_sync_date(config: ExtConfig, data_dir: Path) -> str | None:
|
||||
"""扫描数据文件,返回该扩展配置的最新同步时间(含时分秒)。
|
||||
"""扫描数据文件,返回该扩展配置的最新同步时间(北京墙钟, 含时分秒)。
|
||||
|
||||
- snapshot: 直接取 ext_data/{id}/part.parquet 的 mtime
|
||||
- timeseries: 扫描 ext_data/{id}/timeseries/date=xxx 分区目录
|
||||
|
||||
用北京时间而非宿主机时钟: 前端 ExtDataStatCard 原样展示这串裸时间,
|
||||
容器默认 UTC 时会比同一页拉取面板里的 pull.last_run(带时区 ISO,
|
||||
浏览器按本地时区渲染)整整差一个时区。
|
||||
"""
|
||||
from datetime import datetime
|
||||
|
||||
@@ -245,7 +286,7 @@ def _latest_sync_date(config: ExtConfig, data_dir: Path) -> str | None:
|
||||
# 快照: part.parquet 与 config.json 同级
|
||||
p = data_dir / "ext_data" / config.id / "part.parquet"
|
||||
if p.exists():
|
||||
ts = datetime.fromtimestamp(p.stat().st_mtime).strftime("%Y-%m-%d %H:%M:%S")
|
||||
ts = datetime.fromtimestamp(p.stat().st_mtime, tz=CN_TZ).strftime("%Y-%m-%d %H:%M:%S")
|
||||
return ts
|
||||
# 兼容旧路径
|
||||
old = data_dir / "instruments_ext"
|
||||
@@ -264,7 +305,7 @@ def _latest_sync_date(config: ExtConfig, data_dir: Path) -> str | None:
|
||||
|
||||
|
||||
def _latest_sync_from_partitions(base: Path) -> str | None:
|
||||
"""从 date=xxx 分区目录中找到最新分区的修改时间。"""
|
||||
"""从 date=xxx 分区目录中找到最新分区的修改时间 (北京墙钟)。"""
|
||||
from datetime import datetime
|
||||
latest_ts: float = 0
|
||||
latest_date: str | None = None
|
||||
@@ -276,7 +317,7 @@ def _latest_sync_from_partitions(base: Path) -> str | None:
|
||||
latest_ts = mtime
|
||||
latest_date = d.name[5:]
|
||||
if latest_date and latest_ts > 0:
|
||||
ts = datetime.fromtimestamp(latest_ts).strftime("%H:%M:%S")
|
||||
ts = datetime.fromtimestamp(latest_ts, tz=CN_TZ).strftime("%H:%M:%S")
|
||||
return f"{latest_date} {ts}"
|
||||
return latest_date
|
||||
|
||||
@@ -352,6 +393,7 @@ def create_config(request: Request, body: CreateExtReq):
|
||||
code_map=body.code_map,
|
||||
)
|
||||
store.upsert(config)
|
||||
_refresh_views(request)
|
||||
return config.to_dict()
|
||||
|
||||
|
||||
@@ -373,6 +415,7 @@ def update_config(request: Request, config_id: str, body: UpdateExtReq):
|
||||
if body.code_map is not None:
|
||||
config.code_map = body.code_map
|
||||
store.upsert(config)
|
||||
_refresh_views(request)
|
||||
return config.to_dict()
|
||||
|
||||
|
||||
@@ -382,6 +425,11 @@ def delete_config(request: Request, config_id: str):
|
||||
store = _store(request)
|
||||
if not store.delete(config_id):
|
||||
raise HTTPException(404, f"配置 '{config_id}' 不存在")
|
||||
# 同步清掉 secrets.json 里残留的拉取 API Key, 避免同名重建配置时误用旧 Key
|
||||
from app import secrets_store
|
||||
|
||||
secrets_store.clear(ext_api_key_field(config_id))
|
||||
_refresh_views(request)
|
||||
return {"status": "deleted"}
|
||||
|
||||
|
||||
@@ -682,6 +730,30 @@ def dimension_intraday(
|
||||
# 文件上传
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# 扩展数据 CSV/Excel 上传上限(与自选截图 OCR 的 12MB 上限属同类保护, 见 watchlist.py)。
|
||||
# 通过分块写入临时文件, 超限即拒绝, 避免 `await file.read()` 把整个文件读入内存。
|
||||
_MAX_UPLOAD_BYTES = 50 * 1024 * 1024
|
||||
_UPLOAD_CHUNK_BYTES = 1024 * 1024
|
||||
|
||||
|
||||
async def _write_upload_capped(file: UploadFile, dest: Path, max_bytes: int) -> None:
|
||||
"""分块把上传文件写入 dest, 累计超过 max_bytes 立即拒绝(413)。
|
||||
|
||||
避免一次性 `await file.read()` 把整个文件读入内存(大文件可能触发高内存占用、
|
||||
进程 OOM 或服务不可用); 超限时停止继续读取与落盘。
|
||||
"""
|
||||
total = 0
|
||||
with dest.open("wb") as f:
|
||||
while True:
|
||||
chunk = await file.read(_UPLOAD_CHUNK_BYTES)
|
||||
if not chunk:
|
||||
break
|
||||
total += len(chunk)
|
||||
if total > max_bytes:
|
||||
raise HTTPException(413, f"文件过大(上限 {max_bytes // (1024 * 1024)}MB)")
|
||||
f.write(chunk)
|
||||
|
||||
|
||||
@router.post("/{config_id}/upload")
|
||||
async def upload_data(
|
||||
request: Request,
|
||||
@@ -704,9 +776,7 @@ async def upload_data(
|
||||
tmp_dir = Path(tempfile.mkdtemp())
|
||||
tmp_path = tmp_dir / f"upload{suffix}"
|
||||
try:
|
||||
with tmp_path.open("wb") as f:
|
||||
content = await file.read()
|
||||
f.write(content)
|
||||
await _write_upload_capped(file, tmp_path, _MAX_UPLOAD_BYTES)
|
||||
|
||||
# 直接读取文件,不做列重命名
|
||||
if suffix == ".csv":
|
||||
@@ -792,7 +862,7 @@ def configure_pull(request: Request, config_id: str, body: PullConfigReq):
|
||||
if not config:
|
||||
raise HTTPException(404, f"配置 '{config_id}' 不存在")
|
||||
|
||||
# 保留历史状态字段
|
||||
# 保留历史状态字段; auth 缺省时沿用现有配置 (关闭鉴权需显式传 {"type":"none"})
|
||||
old_pull = config.pull
|
||||
config.pull = PullConfig(
|
||||
url=body.url,
|
||||
@@ -805,6 +875,8 @@ def configure_pull(request: Request, config_id: str, body: PullConfigReq):
|
||||
enabled=body.enabled,
|
||||
time_window_start=body.time_window_start,
|
||||
time_window_end=body.time_window_end,
|
||||
date_param=body.date_param,
|
||||
auth=body.auth.model_dump() if body.auth else (old_pull.auth if old_pull else None),
|
||||
last_run=old_pull.last_run if old_pull else None,
|
||||
last_status=old_pull.last_status if old_pull else None,
|
||||
last_message=old_pull.last_message if old_pull else None,
|
||||
@@ -825,6 +897,38 @@ def configure_pull(request: Request, config_id: str, body: PullConfigReq):
|
||||
return {"status": "ok", "pull": config.pull.to_dict()}
|
||||
|
||||
|
||||
@router.get("/{config_id}/api-key")
|
||||
def get_pull_api_key(request: Request, config_id: str):
|
||||
"""查询拉取接口 API Key 状态。只返回脱敏值, 不返回明文。"""
|
||||
store = _store(request)
|
||||
config = store.get(config_id)
|
||||
if not config:
|
||||
raise HTTPException(404, f"配置 '{config_id}' 不存在")
|
||||
|
||||
from app import secrets_store
|
||||
|
||||
key = get_ext_api_key(config_id)
|
||||
return {"key_set": bool(key), "masked_key": secrets_store.mask(key) if key else ""}
|
||||
|
||||
|
||||
@router.put("/{config_id}/api-key")
|
||||
def set_pull_api_key(request: Request, config_id: str, body: ApiKeyReq):
|
||||
"""设置 (或空串清除) 拉取接口的 API Key, 存 secrets.json (权限 0600)。"""
|
||||
store = _store(request)
|
||||
config = store.get(config_id)
|
||||
if not config:
|
||||
raise HTTPException(404, f"配置 '{config_id}' 不存在")
|
||||
|
||||
from app import secrets_store
|
||||
|
||||
value = body.key.strip()
|
||||
if value:
|
||||
secrets_store.save({ext_api_key_field(config_id): value})
|
||||
else:
|
||||
secrets_store.clear(ext_api_key_field(config_id))
|
||||
return {"status": "ok", "key_set": bool(value), "masked_key": secrets_store.mask(value) if value else ""}
|
||||
|
||||
|
||||
@router.post("/{config_id}/pull/test")
|
||||
async def test_pull(request: Request, config_id: str):
|
||||
"""测试拉取:请求外部 API 并返回预览数据,不写入。"""
|
||||
@@ -835,23 +939,12 @@ async def test_pull(request: Request, config_id: str):
|
||||
if not config.pull or not config.pull.url:
|
||||
raise HTTPException(400, "拉取未配置或 URL 为空")
|
||||
|
||||
# 临时构建一个带新配置的 config 用于测试
|
||||
from app.services.ext_pull import _extract_rows, _apply_field_map
|
||||
import httpx
|
||||
# 复用正式拉取的请求实现 (UA 标识头 + 鉴权注入同一套口径), 不带日期参数
|
||||
from app.services.ext_pull import _apply_field_map, _extract_rows
|
||||
|
||||
pull = config.pull
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30) as client:
|
||||
headers = pull.headers or {}
|
||||
kwargs: dict = {"headers": headers}
|
||||
if pull.method.upper() == "POST" and pull.body:
|
||||
kwargs["content"] = pull.body
|
||||
if "content-type" not in {k.lower() for k in headers}:
|
||||
kwargs["headers"]["Content-Type"] = "application/json"
|
||||
resp = await client.request(pull.method.upper(), pull.url, **kwargs)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
data = await _request_json(pull, config.id)
|
||||
rows = _extract_rows(data, pull.response_path)
|
||||
preview = _apply_field_map(rows[:5], pull.field_map)
|
||||
return {
|
||||
@@ -899,6 +992,38 @@ async def run_pull(request: Request, config_id: str):
|
||||
raise HTTPException(400, f"拉取失败: {e}") from e
|
||||
|
||||
|
||||
@router.post("/{config_id}/backfill")
|
||||
async def backfill_history_ep(
|
||||
request: Request,
|
||||
config_id: str,
|
||||
start: str = Query(..., description="开始日期 YYYY-MM-DD"),
|
||||
end: str = Query(..., description="结束日期 YYYY-MM-DD (含)"),
|
||||
):
|
||||
"""历史回补: 按本地交易日逐日拉取并写入 timeseries 分区。
|
||||
|
||||
前提: 配置为 timeseries 模式且拉取配置了 date_param (接口支持按日期
|
||||
查询)。幂等 —— 已存在的分区跳过, 失败单日不中断, 结果逐项返回。
|
||||
"""
|
||||
store = _store(request)
|
||||
config = store.get(config_id)
|
||||
if not config:
|
||||
raise HTTPException(404, f"配置 '{config_id}' 不存在")
|
||||
try:
|
||||
start_d = date.fromisoformat(start)
|
||||
end_d = date.fromisoformat(end)
|
||||
except ValueError as e:
|
||||
raise HTTPException(422, f"日期格式错误 (应为 YYYY-MM-DD): {e}") from e
|
||||
|
||||
from app.services.ext_pull import backfill_history
|
||||
|
||||
try:
|
||||
result = await backfill_history(config, _data_dir(request), start_d, end_d)
|
||||
except ValueError as e:
|
||||
raise HTTPException(400, str(e)) from e
|
||||
_refresh_views(request)
|
||||
return {"status": "ok", **result}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ---------------------------------------------------------------------------
|
||||
# Symbol 格式修复
|
||||
@@ -939,9 +1064,7 @@ async def detect_fields(
|
||||
tmp_dir = Path(tempfile.mkdtemp())
|
||||
tmp_path = tmp_dir / f"upload{suffix}"
|
||||
try:
|
||||
with tmp_path.open("wb") as f:
|
||||
content = await file.read()
|
||||
f.write(content)
|
||||
await _write_upload_capped(file, tmp_path, _MAX_UPLOAD_BYTES)
|
||||
|
||||
# 直接读取,不要求 symbol 列
|
||||
if suffix == ".csv":
|
||||
@@ -1160,3 +1283,9 @@ def _refresh_views(request: Request) -> None:
|
||||
db.execute(sql)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 扩展列已接入 enriched 帧 (compute_signals/compute_enriched_today 注入):
|
||||
# repo 内存 enriched 缓存 (_enriched_cache/_etf_/_index_) 持有含旧扩展列的
|
||||
# 帧, 必须一并清理, 否则写入后监控/列表仍用旧值 (服务层已清扩展帧与策略缓存)。
|
||||
if hasattr(repo, "clear_cache"):
|
||||
repo.clear_cache()
|
||||
|
||||
@@ -0,0 +1,456 @@
|
||||
"""因子注册表 API — 因子库 (P1) + 公式校验/试算 (P2) + 自定义/复合因子 CRUD (P3)。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, timedelta
|
||||
|
||||
import polars as pl
|
||||
from fastapi import APIRouter, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.factors import store
|
||||
from app.factors.dsl import FACTOR_COLUMN, compile_formula
|
||||
from app.factors.registry import all_factors, unregister_factor
|
||||
|
||||
router = APIRouter(prefix="/api/factors", tags=["factors"])
|
||||
|
||||
|
||||
@router.get("")
|
||||
def list_factors(asset_type: str | None = Query(default=None, pattern="^(stock|etf)$")) -> dict:
|
||||
"""注册表因子列表; asset_type 过滤适用资产 (财务因子仅股票)。"""
|
||||
specs = all_factors(asset_type=asset_type)
|
||||
return {
|
||||
"factors": [
|
||||
{
|
||||
"id": spec.id,
|
||||
"label": spec.label,
|
||||
"group": spec.group,
|
||||
"kind": spec.kind,
|
||||
"version": spec.version,
|
||||
"formula": spec.formula_text,
|
||||
"direction": spec.direction,
|
||||
"unit": spec.unit,
|
||||
"warmup_bars": spec.warmup_bars,
|
||||
"pit": spec.pit,
|
||||
"asset_types": sorted(spec.asset_types),
|
||||
"stability": spec.stability,
|
||||
"scale_free": spec.scale_free,
|
||||
"dependencies": sorted(spec.dependencies),
|
||||
}
|
||||
for spec in specs
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
class FormulaValidateRequest(BaseModel):
|
||||
formula: str = Field(..., min_length=1, max_length=2000)
|
||||
|
||||
|
||||
def _compiled_payload(compiled) -> dict:
|
||||
return {
|
||||
"ok": compiled.ok,
|
||||
"errors": [error.to_dict() for error in compiled.errors],
|
||||
"dependencies": sorted(compiled.dependencies),
|
||||
"referenced_factors": sorted(compiled.referenced_factors),
|
||||
"warmup_bars": compiled.warmup_bars,
|
||||
"cross_sectional": compiled.cross_sectional,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/validate")
|
||||
def validate_formula(req: FormulaValidateRequest) -> dict:
|
||||
"""公式校验: 语法/语义/窗口纪律/依赖推导, 编译期 fail-closed。"""
|
||||
return _compiled_payload(compile_formula(req.formula))
|
||||
|
||||
|
||||
class FormulaTrialRequest(FormulaValidateRequest):
|
||||
asset_type: str = Field(default="stock", pattern="^(stock|etf)$")
|
||||
days: int = Field(default=40, ge=20, le=120)
|
||||
|
||||
|
||||
@router.post("/trial")
|
||||
def trial_formula(req: FormulaTrialRequest, request: Request) -> dict:
|
||||
"""公式试算: 最近 N 个交易日截面 Rank IC 快照 (复用回测面板与虚拟因子物化路径)。"""
|
||||
compiled = compile_formula(req.formula)
|
||||
if not compiled.ok:
|
||||
raise HTTPException(status_code=400, detail={"errors": [error.to_dict() for error in compiled.errors]})
|
||||
|
||||
from app.api.backtest import _get_engine
|
||||
|
||||
# 交易日 → 自然日换算 (A股年均 243 交易日 ≈ 1.48 自然日/交易日), 留 buffer
|
||||
calendar_days = int((compiled.warmup_bars + req.days) * 1.6) + 15
|
||||
start = date.today() - timedelta(days=calendar_days)
|
||||
# 面板基础物理列 (load_panel 只返回 parquet 物理列, 因子列由补算路径生成)
|
||||
base_columns = ["symbol", "date", "open", "high", "low", "close", "volume", "amount", "turnover_rate"]
|
||||
if "consecutive_limit_ups" in compiled.dependencies:
|
||||
base_columns.append("consecutive_limit_ups")
|
||||
engine = _get_engine(request)
|
||||
panel = engine.load_panel(None, start, date.today(), columns=base_columns, asset_type=req.asset_type)
|
||||
if panel.is_empty():
|
||||
raise HTTPException(status_code=400, detail="当前数据目录无可用历史数据, 无法试算")
|
||||
|
||||
# 复用检验引擎同一条补算路径 (compute_indicators + 虚拟因子物化), 禁止第二套计算逻辑
|
||||
from app.backtest.factor import FactorBacktestService
|
||||
|
||||
physical = set(panel.columns)
|
||||
to_compute = set(compiled.referenced_factors) | {
|
||||
dep for dep in compiled.dependencies if dep not in physical
|
||||
}
|
||||
if to_compute:
|
||||
panel = FactorBacktestService._compute_missing_factors(panel, to_compute)
|
||||
|
||||
if compiled.frame_transform is None:
|
||||
raise HTTPException(status_code=500, detail="编译产物缺少帧变换")
|
||||
prepared = compiled.frame_transform(panel)
|
||||
if prepared is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"errors": [{
|
||||
"code": "E013", "message": "依赖列不可用: 面板缺少公式所需列",
|
||||
"position": {"offset": 0, "line": 1},
|
||||
"detail": {"missing": sorted((compiled.dependencies | compiled.referenced_factors) - set(panel.columns))},
|
||||
}]},
|
||||
)
|
||||
|
||||
total_rows = panel.height
|
||||
frame = (
|
||||
prepared
|
||||
.with_columns(
|
||||
(pl.col("close").shift(-1).over("symbol") / pl.col("close") - 1.0).alias("_next_return")
|
||||
)
|
||||
.filter(pl.col(FACTOR_COLUMN).is_not_null())
|
||||
.unique(subset=["symbol", "date"], keep="last")
|
||||
.sort(["symbol", "date"])
|
||||
)
|
||||
non_null_rows = frame.height
|
||||
if non_null_rows == 0:
|
||||
return {
|
||||
"ok": True, "n_dates": 0, "null_ratio": 1.0,
|
||||
"ic_mean": None, "ic_std": None, "ir": None, "ic_win_rate": None,
|
||||
"ic_series": [], "message": "试算区间内公式输出全为空 (检查预热窗口与数据范围)",
|
||||
}
|
||||
|
||||
ic_frame = (
|
||||
frame.filter(pl.col("_next_return").is_not_null())
|
||||
.group_by("date")
|
||||
.agg(
|
||||
pl.corr(pl.col(FACTOR_COLUMN).rank(method="average"), pl.col("_next_return").rank(method="average")).alias("ic"),
|
||||
pl.len().alias("n_symbols"),
|
||||
)
|
||||
.filter(pl.col("ic").is_not_null())
|
||||
.sort("date")
|
||||
.tail(req.days)
|
||||
)
|
||||
ic_series = [
|
||||
{"date": str(row["date"]), "ic": round(row["ic"], 4), "n_symbols": row["n_symbols"]}
|
||||
for row in ic_frame.to_dicts()
|
||||
]
|
||||
if ic_frame.is_empty():
|
||||
return {
|
||||
"ok": True, "n_dates": 0, "null_ratio": round(1.0 - non_null_rows / max(total_rows, 1), 4),
|
||||
"ic_mean": None, "ic_std": None, "ir": None, "ic_win_rate": None,
|
||||
"ic_series": [], "message": "无有效 IC 截面 (需每期 ≥2 只标的)",
|
||||
}
|
||||
stats = ic_frame.select(
|
||||
pl.col("ic").mean().alias("mean"),
|
||||
pl.col("ic").std(ddof=0).alias("std"),
|
||||
(pl.col("ic") > 0).mean().alias("win"),
|
||||
).row(0, named=True)
|
||||
ic_std = stats["std"]
|
||||
# Newey-West t (lag=1): 与检验页同源口径, 样本过少时不给 (fail-closed)
|
||||
t_newey_west = None
|
||||
if ic_frame.height >= 5:
|
||||
from app.backtest.stats_v2 import newey_west_t
|
||||
|
||||
values = ic_frame["ic"].to_numpy()
|
||||
nw = newey_west_t(values, lag=1)
|
||||
if nw is not None:
|
||||
t_newey_west = round(float(nw[0]), 3)
|
||||
return {
|
||||
"ok": True,
|
||||
"n_dates": ic_frame.height,
|
||||
"null_ratio": round(1.0 - non_null_rows / max(total_rows, 1), 4),
|
||||
"ic_mean": round(stats["mean"], 4),
|
||||
"ic_std": None if ic_std is None else round(ic_std, 4),
|
||||
"ir": None if not ic_std or ic_std == 0 else round(stats["mean"] / ic_std, 3),
|
||||
"ic_win_rate": round(stats["win"], 4),
|
||||
"t_newey_west": t_newey_west,
|
||||
"ic_series": ic_series,
|
||||
}
|
||||
|
||||
|
||||
# ── 自定义/复合因子 CRUD (P3) ──────────────────────────────
|
||||
|
||||
|
||||
class CustomFactorCreateRequest(BaseModel):
|
||||
id: str | None = Field(default=None, max_length=48)
|
||||
label: str = Field(..., min_length=1, max_length=32)
|
||||
group: str = Field(default="自定义", max_length=16)
|
||||
formula: str = Field(..., min_length=1, max_length=2000)
|
||||
description: str = Field(default="", max_length=500)
|
||||
direction: str = Field(default="none", pattern="^(high|low|none)$")
|
||||
|
||||
|
||||
class CompositeFactorCreateRequest(BaseModel):
|
||||
id: str | None = Field(default=None, max_length=48)
|
||||
label: str = Field(..., min_length=1, max_length=32)
|
||||
group: str = Field(default="组合", max_length=16)
|
||||
members: dict[str, float] = Field(..., min_length=2, max_length=8)
|
||||
description: str = Field(default="", max_length=500)
|
||||
direction: str = Field(default="none", pattern="^(high|low|none)$")
|
||||
|
||||
|
||||
def _data_dir(request: Request):
|
||||
from pathlib import Path
|
||||
|
||||
data_dir = getattr(getattr(request.app.state, "repo", None), "store", None)
|
||||
root = getattr(data_dir, "data_dir", None) if data_dir is not None else None
|
||||
if root is None:
|
||||
raise HTTPException(status_code=500, detail="数据目录不可用")
|
||||
return Path(root)
|
||||
|
||||
|
||||
def _slugify_id(label: str, prefix: str) -> str:
|
||||
base = "".join(ch if ch.isascii() and (ch.isalnum() or ch == "_") else "_" for ch in label.lower())
|
||||
candidate = f"{prefix}_{base}".strip("_")[:44]
|
||||
import re
|
||||
|
||||
candidate = re.sub(r"_+", "_", candidate)
|
||||
return candidate or f"{prefix}_f"
|
||||
|
||||
|
||||
def _resolve_id(requested: str | None, label: str, prefix: str) -> str:
|
||||
return requested.strip() if requested and requested.strip() else _slugify_id(label, prefix)
|
||||
|
||||
|
||||
def _next_version(data_dir, factor_id: str) -> int:
|
||||
for definition in store.load_all(data_dir):
|
||||
if str(definition.get("id")) == factor_id:
|
||||
return int(definition.get("version", 1)) + 1
|
||||
return 1
|
||||
|
||||
|
||||
def _trial_nonempty(request: Request, formula: str, asset_type: str = "stock") -> None:
|
||||
"""保存前置校验: 公式在最近 40 个交易日有非空输出 (设计 §3.5, fail-closed)。"""
|
||||
compiled = compile_formula(formula)
|
||||
if not compiled.ok:
|
||||
raise HTTPException(status_code=400, detail={"errors": [e.to_dict() for e in compiled.errors]})
|
||||
from app.api.backtest import _get_engine
|
||||
from app.backtest.factor import FactorBacktestService
|
||||
|
||||
calendar_days = int((compiled.warmup_bars + 40) * 1.6) + 15
|
||||
base_columns = ["symbol", "date", "open", "high", "low", "close", "volume", "amount", "turnover_rate"]
|
||||
engine = _get_engine(request)
|
||||
panel = engine.load_panel(None, date.today() - timedelta(days=calendar_days), date.today(), columns=base_columns, asset_type=asset_type)
|
||||
if panel.is_empty():
|
||||
raise HTTPException(status_code=400, detail="当前无历史数据, 无法完成保存前试算 (fail-closed)")
|
||||
physical = set(panel.columns)
|
||||
to_compute = set(compiled.referenced_factors) | {d for d in compiled.dependencies if d not in physical}
|
||||
if to_compute:
|
||||
panel = FactorBacktestService._compute_missing_factors(panel, to_compute)
|
||||
prepared = compiled.frame_transform(panel) if compiled.frame_transform else None
|
||||
if prepared is None or prepared[FACTOR_COLUMN].is_not_null().sum() == 0:
|
||||
raise HTTPException(status_code=400, detail="公式在最近 40 个交易日输出全为空, 拒绝保存")
|
||||
|
||||
|
||||
@router.post("/custom")
|
||||
def create_custom_factor(req: CustomFactorCreateRequest, request: Request) -> dict:
|
||||
"""保存自定义公式因子: 编译通过 + 服务端试算非空 (fail-closed)。"""
|
||||
data_dir = _data_dir(request)
|
||||
factor_id = _resolve_id(req.id, req.label, "uf")
|
||||
definition = {
|
||||
"id": factor_id,
|
||||
"kind": "custom",
|
||||
"version": _next_version(data_dir, factor_id),
|
||||
"label": req.label,
|
||||
"group": req.group,
|
||||
"formula": req.formula,
|
||||
"description": req.description,
|
||||
"direction": req.direction,
|
||||
"status": "draft",
|
||||
"created_at": store._now(),
|
||||
"updated_at": store._now(),
|
||||
}
|
||||
try:
|
||||
store.to_spec(definition) # 先做 schema/id/编译校验
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
_trial_nonempty(request, req.formula)
|
||||
try:
|
||||
store.register_definition(definition)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
store.save_one(data_dir, definition)
|
||||
return {"ok": True, "id": factor_id, "version": definition["version"]}
|
||||
|
||||
|
||||
@router.post("/composite")
|
||||
def create_composite_factor(req: CompositeFactorCreateRequest, request: Request) -> dict:
|
||||
"""保存复合因子: 成员校验 + 循环引用检查 (无需试算, 值由成员物化路径计算)。"""
|
||||
data_dir = _data_dir(request)
|
||||
factor_id = _resolve_id(req.id, req.label, "cf")
|
||||
definition = {
|
||||
"id": factor_id,
|
||||
"kind": "composite",
|
||||
"version": _next_version(data_dir, factor_id),
|
||||
"label": req.label,
|
||||
"group": req.group,
|
||||
"members": req.members,
|
||||
"description": req.description,
|
||||
"direction": req.direction,
|
||||
"status": "draft",
|
||||
"created_at": store._now(),
|
||||
"updated_at": store._now(),
|
||||
}
|
||||
try:
|
||||
store.register_definition(definition)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
store.save_one(data_dir, definition)
|
||||
return {"ok": True, "id": factor_id, "version": definition["version"]}
|
||||
|
||||
|
||||
class CustomFactorUpdateRequest(BaseModel):
|
||||
label: str = Field(..., min_length=1, max_length=32)
|
||||
group: str = Field(default="自定义", max_length=16)
|
||||
formula: str = Field(..., min_length=1, max_length=2000)
|
||||
description: str = Field(default="", max_length=500)
|
||||
direction: str = Field(default="none", pattern="^(high|low|none)$")
|
||||
|
||||
|
||||
@router.post("/custom/{factor_id}/update")
|
||||
def update_custom_factor(factor_id: str, req: CustomFactorUpdateRequest, request: Request) -> dict:
|
||||
"""编辑已有自定义因子: 编译校验 + 试算非空 (与创建同一门禁) → 版本提升注册。
|
||||
|
||||
公式变化时状态回 draft (生命周期语义: 编辑后需重新检验激活); 仅改名称/分组保留状态。
|
||||
"""
|
||||
data_dir = _data_dir(request)
|
||||
target = None
|
||||
for definition in store.load_all(data_dir):
|
||||
if str(definition.get("id")) == factor_id:
|
||||
target = definition
|
||||
break
|
||||
if target is None:
|
||||
raise HTTPException(status_code=404, detail=f"自定义因子不存在: {factor_id}")
|
||||
if str(target.get("kind", "custom")) != "custom":
|
||||
raise HTTPException(status_code=400, detail=f"仅自定义因子支持公式编辑 (kind={target.get('kind')})")
|
||||
formula_changed = str(target.get("formula")) != req.formula
|
||||
if formula_changed:
|
||||
_trial_nonempty(request, req.formula)
|
||||
target.update({
|
||||
"label": req.label,
|
||||
"group": req.group,
|
||||
"formula": req.formula,
|
||||
"description": req.description,
|
||||
"direction": req.direction,
|
||||
"version": int(target.get("version", 1)) + 1, # 版本提升 → 注册表允许覆盖
|
||||
"status": "draft" if formula_changed else str(target.get("status", "draft")),
|
||||
"updated_at": store._now(),
|
||||
})
|
||||
try:
|
||||
store.register_definition(target)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
store.save_one(data_dir, target)
|
||||
return {"ok": True, "id": factor_id, "version": target["version"], "status": target["status"]}
|
||||
|
||||
|
||||
def _find_references(data_dir, factor_id: str) -> list[str]:
|
||||
"""扫描策略与复合因子定义中的引用 (删除前 fail-closed 检查)。"""
|
||||
references: list[str] = []
|
||||
strategies_dir = data_dir / "strategies"
|
||||
if strategies_dir.is_dir():
|
||||
for file in strategies_dir.glob("*.json"):
|
||||
try:
|
||||
text = file.read_text(encoding="utf-8")
|
||||
if factor_id in text:
|
||||
references.append(f"strategies/{file.name}")
|
||||
except OSError:
|
||||
continue
|
||||
for definition in store.load_all(data_dir):
|
||||
if str(definition.get("id")) == factor_id:
|
||||
continue
|
||||
members = definition.get("members")
|
||||
if isinstance(members, dict) and factor_id in members:
|
||||
references.append(f"custom_factors/{definition.get('id')}.json")
|
||||
return references
|
||||
|
||||
|
||||
@router.delete("/custom/{factor_id}")
|
||||
def delete_custom_factor(factor_id: str, request: Request, force: bool = Query(default=False)) -> dict:
|
||||
"""删除自定义/复合因子; 有引用时列出引用方并拒绝 (需 force)。"""
|
||||
data_dir = _data_dir(request)
|
||||
from app.factors.registry import get_factor
|
||||
|
||||
# 引用检查必须排在存在性判定之前: 下面用来探测「盘上是否有定义」的
|
||||
# store.delete_one 本身就会删文件, 反过来会出现「拒绝删除」但定义已被删掉。
|
||||
references = _find_references(data_dir, factor_id)
|
||||
if references and not force:
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail={"message": "该因子仍有引用, 拒绝删除 (可带 force=true 强制)", "references": references},
|
||||
)
|
||||
if get_factor(factor_id) is None and not store.delete_one(data_dir, factor_id):
|
||||
raise HTTPException(status_code=404, detail=f"因子不存在: {factor_id}")
|
||||
try:
|
||||
unregister_factor(factor_id)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
store.delete_one(data_dir, factor_id)
|
||||
return {"ok": True, "id": factor_id, "removed_references": references}
|
||||
|
||||
|
||||
class FactorStatusRequest(BaseModel):
|
||||
status: str = Field(..., pattern="^(draft|active|watch|retired)$")
|
||||
|
||||
|
||||
@router.post("/custom/{factor_id}/status")
|
||||
def update_factor_status(factor_id: str, req: FactorStatusRequest, request: Request) -> dict:
|
||||
"""生命周期状态迁移 (P4): draft->active->watch->retired, 编辑后回 draft。"""
|
||||
data_dir = _data_dir(request)
|
||||
target = None
|
||||
for definition in store.load_all(data_dir):
|
||||
if str(definition.get("id")) == factor_id:
|
||||
target = definition
|
||||
break
|
||||
if target is None:
|
||||
raise HTTPException(status_code=404, detail=f"自定义因子不存在: {factor_id}")
|
||||
target["status"] = req.status
|
||||
target["updated_at"] = store._now()
|
||||
try:
|
||||
# 动态因子先注销再注册: 元数据变更 (status/group) 不提升版本,
|
||||
# 直接 register 会因"版本未提升"被拒 (启动加载后的真实路径)
|
||||
unregister_factor(factor_id)
|
||||
store.register_definition(target) # 状态与 stability 联动
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
store.save_one(data_dir, target)
|
||||
return {"ok": True, "id": factor_id, "status": req.status}
|
||||
|
||||
|
||||
class FactorGroupRequest(BaseModel):
|
||||
group: str = Field(..., min_length=1, max_length=24)
|
||||
|
||||
|
||||
@router.post("/custom/{factor_id}/group")
|
||||
def update_factor_group(factor_id: str, req: FactorGroupRequest, request: Request) -> dict:
|
||||
"""修改单个自定义/复合因子的分组 (内置因子分组与快照/预设绑定, 不可改)。"""
|
||||
data_dir = _data_dir(request)
|
||||
group = req.group.strip()
|
||||
if not group:
|
||||
raise HTTPException(status_code=400, detail="分组名不能为空")
|
||||
target = None
|
||||
for definition in store.load_all(data_dir):
|
||||
if str(definition.get("id")) == factor_id:
|
||||
target = definition
|
||||
break
|
||||
if target is None:
|
||||
raise HTTPException(status_code=404, detail=f"自定义因子不存在: {factor_id}")
|
||||
target["group"] = group
|
||||
target["updated_at"] = store._now()
|
||||
try:
|
||||
unregister_factor(factor_id)
|
||||
store.register_definition(target)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
store.save_one(data_dir, target)
|
||||
return {"ok": True, "id": factor_id, "group": group}
|
||||
@@ -70,9 +70,12 @@ def get_index_minute(
|
||||
):
|
||||
"""实时读取指数分钟 K。不写入股票分钟 parquet。"""
|
||||
repo = request.app.state.repo
|
||||
capset = request.app.state.capabilities
|
||||
info = _index_info(repo, symbol)
|
||||
day = trade_date or date.today()
|
||||
df = kline_sync.fetch_minute_single(symbol, day, asset_type="index")
|
||||
df = kline_sync.fetch_minute_single(
|
||||
symbol, day, asset_type="index", capset=capset,
|
||||
)
|
||||
return {
|
||||
"symbol": symbol,
|
||||
"name": info.get("name"),
|
||||
|
||||
+163
-84
@@ -27,7 +27,7 @@ router = APIRouter(prefix="/api/kline", tags=["kline"])
|
||||
def _gzip_payload(request: Request, payload: dict, *, pref_key: str) -> dict | Response:
|
||||
"""大 JSON 响应的传输压缩: 偏好开启 + 客户端接受 gzip + 响应超阈值才压。
|
||||
|
||||
分时/日K批量各自独立偏好键 (网络设置里大开关批量、子开关单独控制)。
|
||||
分时/日K各自使用独立偏好键 (沿用已有 *_batch_compress 存储键保证兼容)。
|
||||
level 6 实测 13MB ≈ 290ms CPU 压掉 87%; level 9 要 2.5s 不可用。
|
||||
datetime → isoformat, 与 FastAPI jsonable_encoder 输出一致
|
||||
(前端 since 增量按字符串字典序比较, 格式必须与非压缩路径相同)。
|
||||
@@ -393,7 +393,11 @@ def get_daily(
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=502, detail=f"TickFlow fetch failed: {e}") from e
|
||||
if raw.is_empty():
|
||||
return {"symbol": symbol, "name": stock_name, "stock_info": stock_info, "rows": []}
|
||||
return _gzip_payload(
|
||||
request,
|
||||
{"symbol": symbol, "name": stock_name, "stock_info": stock_info, "rows": []},
|
||||
pref_key="daily_batch_compress",
|
||||
)
|
||||
# 拉除权因子做前复权 (Starter+ 有权限), 否则空 df → compute_enriched 退回未复权
|
||||
factors = pl.DataFrame()
|
||||
capset = getattr(request.app.state, "capabilities", None)
|
||||
@@ -408,7 +412,11 @@ def get_daily(
|
||||
# 即使 live 模式也尝试追加实时蜡烛
|
||||
rows = _maybe_inject_live_candle(request, symbol, rows, asset_type)
|
||||
resp = {"symbol": symbol, "name": stock_name, "stock_info": stock_info, "rows": rows, "source": "live"}
|
||||
return _attach_ext(resp, repo, symbol, ext_columns)
|
||||
return _gzip_payload(
|
||||
request,
|
||||
_attach_ext(resp, repo, symbol, ext_columns),
|
||||
pref_key="daily_batch_compress",
|
||||
)
|
||||
|
||||
rows = df.to_dicts()
|
||||
|
||||
@@ -416,7 +424,11 @@ def get_daily(
|
||||
rows = _maybe_inject_live_candle(request, symbol, rows, asset_type)
|
||||
|
||||
resp = {"symbol": symbol, "name": stock_name, "stock_info": stock_info, "rows": rows, "source": "enriched"}
|
||||
return _attach_ext(resp, repo, symbol, ext_columns)
|
||||
return _gzip_payload(
|
||||
request,
|
||||
_attach_ext(resp, repo, symbol, ext_columns),
|
||||
pref_key="daily_batch_compress",
|
||||
)
|
||||
|
||||
|
||||
def _attach_ext(resp: dict, repo, symbol: str, ext_columns: Optional[str]) -> dict:
|
||||
@@ -459,51 +471,54 @@ def _attach_ext(resp: dict, repo, symbol: str, ext_columns: Optional[str]) -> di
|
||||
return resp
|
||||
|
||||
|
||||
def _maybe_inject_live_candle(request: Request, symbol: str, rows: list[dict], asset_type: str = "stock") -> list[dict]:
|
||||
"""如果有当日实时 enriched 数据, 用实时数据生成今日蜡烛并追加/覆盖。
|
||||
def _latest_live_candle(
|
||||
request: Request,
|
||||
symbol: str,
|
||||
asset_type: str = "stock",
|
||||
*,
|
||||
refresh_asset: bool = True,
|
||||
) -> dict | None:
|
||||
"""从内存缓存读取单只标的的当日实时 enriched 行。"""
|
||||
|
||||
stock 走 QuoteService 的股票实时缓存; etf 走 ETF enriched 缓存 (开启实时 ETF
|
||||
拉取时为盘中数据, 否则为磁盘最新日, 由下方"非今日不注入"守卫自然跳过)。
|
||||
"""
|
||||
if asset_type == "stock":
|
||||
qs = getattr(request.app.state, "quote_service", None)
|
||||
if not qs:
|
||||
return rows
|
||||
return None
|
||||
df_today, enriched_date = qs.get_enriched_today()
|
||||
elif asset_type == "etf":
|
||||
df_today, enriched_date = request.app.state.repo.get_enriched_latest_asset("etf")
|
||||
df_today, enriched_date = request.app.state.repo.get_enriched_latest_asset(
|
||||
"etf", refresh=refresh_asset,
|
||||
)
|
||||
else:
|
||||
return rows
|
||||
return None
|
||||
if df_today.is_empty():
|
||||
return rows
|
||||
return None
|
||||
|
||||
# 非交易日(周末/假日)缓存的行情日期 != 今天,跳过注入避免产生重复蜡烛
|
||||
# 非交易日(周末/假日)缓存日期 != 今天, 跳过注入避免产生重复蜡烛
|
||||
if not enriched_date or enriched_date != date.today():
|
||||
return rows
|
||||
return None
|
||||
|
||||
# 查找该 symbol 的实时 enriched 行
|
||||
import polars as pl
|
||||
try:
|
||||
q = df_today.filter(pl.col("symbol") == symbol).to_dicts()
|
||||
if not q:
|
||||
return rows
|
||||
return None
|
||||
q = q[0]
|
||||
except Exception: # noqa: BLE001
|
||||
return rows
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
close_price = q.get("close")
|
||||
if not close_price or close_price <= 0:
|
||||
return rows
|
||||
return None
|
||||
|
||||
today_str = str(enriched_date)
|
||||
|
||||
# enriched 行已包含 OHLCV + 全套指标, 直接用它
|
||||
# 修复: API 在非交易时段可能返回 open/high/low=0, 用 close 填充避免异常蜡烛
|
||||
# 沿用完整日K接口原有的实时行投影, 避免增量接口形成第二套字段契约。
|
||||
# API 在非交易时段可能返回 open/high/low=0, 用 close 填充避免异常蜡烛。
|
||||
raw_open = q.get("open")
|
||||
raw_high = q.get("high")
|
||||
raw_low = q.get("low")
|
||||
live_row: dict = {
|
||||
"date": today_str,
|
||||
live_row = {
|
||||
"date": str(enriched_date),
|
||||
"symbol": symbol,
|
||||
"open": raw_open if raw_open and raw_open > 0 else close_price,
|
||||
"high": raw_high if raw_high and raw_high > 0 else close_price,
|
||||
@@ -514,7 +529,6 @@ def _maybe_inject_live_candle(request: Request, symbol: str, rows: list[dict], a
|
||||
"change_pct": q.get("change_pct"),
|
||||
"is_live": True,
|
||||
}
|
||||
# 补上 enriched 的技术指标字段
|
||||
for key in ("ma5", "ma10", "ma20", "ma30", "ma60",
|
||||
"macd_dif", "macd_dea", "macd_hist",
|
||||
"kdj_k", "kdj_d", "kdj_j",
|
||||
@@ -523,11 +537,19 @@ def _maybe_inject_live_candle(request: Request, symbol: str, rows: list[dict], a
|
||||
"atr_14", "vol_ratio_5d"):
|
||||
if key in q and q[key] is not None:
|
||||
live_row[key] = q[key]
|
||||
return live_row
|
||||
|
||||
|
||||
def _maybe_inject_live_candle(request: Request, symbol: str, rows: list[dict], asset_type: str = "stock") -> list[dict]:
|
||||
"""如果有当日实时 enriched 数据, 用实时数据生成今日蜡烛并追加/覆盖。"""
|
||||
live_row = _latest_live_candle(request, symbol, asset_type)
|
||||
if live_row is None:
|
||||
return rows
|
||||
|
||||
# 如果已有今天的 enriched 行, 覆盖; 否则追加
|
||||
found = False
|
||||
for i, r in enumerate(rows):
|
||||
if str(r.get("date")) == today_str:
|
||||
for r in rows:
|
||||
if str(r.get("date")) == live_row["date"]:
|
||||
r.update(live_row)
|
||||
found = True
|
||||
break
|
||||
@@ -538,6 +560,22 @@ def _maybe_inject_live_candle(request: Request, symbol: str, rows: list[dict], a
|
||||
return rows
|
||||
|
||||
|
||||
@router.get("/daily/latest")
|
||||
def get_daily_latest(
|
||||
request: Request,
|
||||
symbol: str = Query(..., description="标的代码,如 000001.SZ"),
|
||||
):
|
||||
"""返回内存中的当日单行 K 线, 供详情页实时增量更新。"""
|
||||
repo = request.app.state.repo
|
||||
asset_type = repo.resolve_asset_type(symbol)
|
||||
row = _latest_live_candle(request, symbol, asset_type, refresh_asset=False)
|
||||
return {
|
||||
"symbol": symbol,
|
||||
"row": row,
|
||||
"source": "live" if row is not None else "none",
|
||||
}
|
||||
|
||||
|
||||
class DailyBatchRequest:
|
||||
"""批量日K请求。"""
|
||||
symbols: list[str]
|
||||
@@ -703,12 +741,19 @@ def get_minute_batch(request: Request, body: dict):
|
||||
# 本地状态分类 (补拉已改为取到即落盘, 完整性判定随之收紧):
|
||||
# - fresh: 根数 >= 期望-2 (时间边界容差), 直接用本地。原 0.9 比例阈值会让
|
||||
# 持久化数据在 90% 处冻结尾巴, 必须按根数差判。
|
||||
# - holes: 中间缺K (相邻间距非 1 分钟 / 非午休 91 分钟) → 全天重拉回填,
|
||||
# 否则"最后一根+1min"的增量窗口永远不会回看中间的洞。
|
||||
# - holes: 缺K → 全天重拉回填, 否则"最后一根+1min"的增量窗口永远不会
|
||||
# 回看洞。含两种: 中间的洞 (相邻间距非 1 分钟 / 非午休 91 分钟)
|
||||
# 与前部的洞 (首根显著晚于开盘 — 盘中重启/停机跨开盘的残留,
|
||||
# 连续的尾部K会被增量锚定锁死, 同样必须全天重拉)。
|
||||
# - stale: 仅尾部落后 → 增量拉, 请求量从"每轮全天"降为"每轮一根"量级。
|
||||
_LUNCH_GAP_MIN = 91 # 11:30 → 13:01
|
||||
# 前部洞基准: 开盘后 6 分钟 (容许无集合竞价K的数据源)。晚开/停牌复牌的票
|
||||
# 也会命中 → 全天拉幂等, 至多多一次批量请求, 与中间洞同一代价模型。
|
||||
day_open_floor = datetime(trade_date.year, trade_date.month, trade_date.day, 9, 36, 0)
|
||||
|
||||
def _has_holes(sub: pl.DataFrame) -> bool:
|
||||
if not sub.is_empty() and sub["datetime"][0] > day_open_floor:
|
||||
return True
|
||||
gaps = sub["datetime"].diff().dt.total_minutes().drop_nulls()
|
||||
return gaps.filter((gaps != 1) & (gaps != _LUNCH_GAP_MIN)).len() > 0
|
||||
|
||||
@@ -737,14 +782,15 @@ def get_minute_batch(request: Request, body: dict):
|
||||
svc = getattr(request.app.state, "minute_refresh", None)
|
||||
full_minute_healthy = bool(svc is not None and svc.is_healthy())
|
||||
if full_minute_healthy:
|
||||
# 股票缺口不补拉, 本地有多少给多少 (服务下一轮写入补全);
|
||||
# ETF 不在 universe 内, 维持补拉
|
||||
for sym in [*full_pull, *stale_last]:
|
||||
# 纯尾部落后 (stale_last): 服务的增量轮下一轮就会补上, 股票不补拉省请求。
|
||||
# 空洞 (full_pull: 空分区 / 中间洞 / 前部洞): 服务增量锚定本地最新时间,
|
||||
# 永远不会回看洞 → 不压制, 由端点全天拉取并落盘修复。
|
||||
# ETF 不在服务 universe 内, 两类均维持补拉。
|
||||
for sym in stale_last:
|
||||
if sym not in etf_set:
|
||||
sub = local_parts.get(sym)
|
||||
if sub is not None and not sub.is_empty():
|
||||
result[sym] = sub.to_dicts()
|
||||
full_pull = [s for s in full_pull if s in etf_set]
|
||||
stale_last = {s: t for s, t in stale_last.items() if s in etf_set}
|
||||
|
||||
# Step 2: 补拉并落盘 (取到即写, upsert 语义; 下一轮命中本地, 请求量骤降)。
|
||||
@@ -853,13 +899,21 @@ def get_minute_range(
|
||||
|
||||
# 指数分钟 K 不落本地仓库, 最新分时仍由 /api/index/minute 实时读取。
|
||||
if asset_type == "index":
|
||||
return {**base_response, "sessions": [], "source": "none"}
|
||||
return _gzip_payload(
|
||||
request,
|
||||
{**base_response, "sessions": [], "source": "none"},
|
||||
pref_key="minute_batch_compress",
|
||||
)
|
||||
|
||||
end = cn_today()
|
||||
start = end - timedelta(days=days * 3 + 20)
|
||||
minute = repo.get_minute_range([symbol], start, end, asset_type=asset_type)
|
||||
if minute.is_empty() or "datetime" not in minute.columns:
|
||||
return {**base_response, "sessions": [], "source": "none"}
|
||||
return _gzip_payload(
|
||||
request,
|
||||
{**base_response, "sessions": [], "source": "none"},
|
||||
pref_key="minute_batch_compress",
|
||||
)
|
||||
|
||||
minute = minute.with_columns(
|
||||
pl.col("datetime").dt.date().alias("_trade_date"),
|
||||
@@ -888,11 +942,15 @@ def get_minute_range(
|
||||
"rows": rows,
|
||||
})
|
||||
|
||||
return {
|
||||
**base_response,
|
||||
"sessions": sessions,
|
||||
"source": "local" if sessions else "none",
|
||||
}
|
||||
return _gzip_payload(
|
||||
request,
|
||||
{
|
||||
**base_response,
|
||||
"sessions": sessions,
|
||||
"source": "local" if sessions else "none",
|
||||
},
|
||||
pref_key="minute_batch_compress",
|
||||
)
|
||||
|
||||
|
||||
@router.get("/minute")
|
||||
@@ -905,12 +963,14 @@ def get_minute(
|
||||
"""读取某只股票某天的分钟 K 线。
|
||||
|
||||
- 本地有完整数据(240条) → 直接返回
|
||||
- 本地无数据或不完整 → 从 TickFlow 实时拉取返回(不写入)
|
||||
- 本地无数据或不完整 → 从有效分钟数据源实时拉取返回(不写入)
|
||||
- 自定义源失败时, 仅具备 TickFlow 单股分钟能力才回退 TickFlow
|
||||
- live=true 且当日连续竞价时段 → 跳过本地优先直接实时拉取:
|
||||
盘中分钟增量落盘的本地分区按 ≥60s 轮次更新, 90% 完整度启发式会让
|
||||
详情分时图停在上一增量轮, 与行情列表的节奏脱节
|
||||
"""
|
||||
repo = request.app.state.repo
|
||||
capset = request.app.state.capabilities
|
||||
asset_type = repo.resolve_asset_type(symbol)
|
||||
stock_info = _get_stock_info(repo, symbol) if asset_type == "stock" else _get_asset_info(repo, symbol, asset_type)
|
||||
stock_name = stock_info.get("name")
|
||||
@@ -935,22 +995,29 @@ def get_minute(
|
||||
else:
|
||||
trade_date = today
|
||||
if trade_date is None:
|
||||
# 本地无任何分钟K,尝试从 TickFlow 拉取当天
|
||||
# 本地无任何分钟K, 尝试从当前有效分钟源拉取当天
|
||||
trade_date = cn_today()
|
||||
df = kline_sync.fetch_minute_single(symbol, trade_date, asset_type=asset_type)
|
||||
df = kline_sync.fetch_minute_single(
|
||||
symbol, trade_date, asset_type=asset_type, capset=capset,
|
||||
)
|
||||
price_limit = _get_price_limit_info(
|
||||
repo, symbol, trade_date, asset_type, stock_name,
|
||||
)
|
||||
prev_close = _get_previous_closes(
|
||||
repo, symbol, [trade_date], asset_type,
|
||||
).get(trade_date)
|
||||
return {
|
||||
"symbol": symbol, "name": stock_name, "stock_info": stock_info,
|
||||
"date": str(trade_date), "rows": df.to_dicts(), "source": "live",
|
||||
"asset_type": asset_type,
|
||||
"price_limit": price_limit,
|
||||
"prev_close": prev_close,
|
||||
}
|
||||
return _gzip_payload(
|
||||
request,
|
||||
{
|
||||
"symbol": symbol, "name": stock_name, "stock_info": stock_info,
|
||||
"date": str(trade_date), "rows": df.to_dicts(),
|
||||
"source": "live" if not df.is_empty() else "none",
|
||||
"asset_type": asset_type,
|
||||
"price_limit": price_limit,
|
||||
"prev_close": prev_close,
|
||||
},
|
||||
pref_key="minute_batch_compress",
|
||||
)
|
||||
|
||||
prev_close = _get_previous_closes(
|
||||
repo, symbol, [trade_date], asset_type,
|
||||
@@ -962,14 +1029,20 @@ def get_minute(
|
||||
if live and trade_date == cn_today() and in_continuous_session():
|
||||
# 详情分时轮询: 当日盘中实时拉取最新一根K, 不落盘; 拉空(源侧延迟/
|
||||
# 时段边界)则落回下方本地优先路径。
|
||||
live_df = kline_sync.fetch_minute_single(symbol, trade_date, asset_type=asset_type)
|
||||
live_df = kline_sync.fetch_minute_single(
|
||||
symbol, trade_date, asset_type=asset_type, capset=capset,
|
||||
)
|
||||
if not live_df.is_empty():
|
||||
return {
|
||||
"symbol": symbol, "name": stock_name, "stock_info": stock_info,
|
||||
"date": str(trade_date), "rows": live_df.to_dicts(),
|
||||
"source": "live", "asset_type": asset_type,
|
||||
"price_limit": price_limit, "prev_close": prev_close,
|
||||
}
|
||||
return _gzip_payload(
|
||||
request,
|
||||
{
|
||||
"symbol": symbol, "name": stock_name, "stock_info": stock_info,
|
||||
"date": str(trade_date), "rows": live_df.to_dicts(),
|
||||
"source": "live", "asset_type": asset_type,
|
||||
"price_limit": price_limit, "prev_close": prev_close,
|
||||
},
|
||||
pref_key="minute_batch_compress",
|
||||
)
|
||||
|
||||
df = repo.get_minute(symbol, trade_date, asset_type=asset_type)
|
||||
|
||||
@@ -993,24 +1066,34 @@ def get_minute(
|
||||
is_complete = not df.is_empty() and len(df) >= expected * 0.9 # 允许 10% 容差
|
||||
|
||||
if is_complete:
|
||||
return {
|
||||
return _gzip_payload(
|
||||
request,
|
||||
{
|
||||
"symbol": symbol, "name": stock_name, "stock_info": stock_info,
|
||||
"date": str(trade_date), "rows": df.to_dicts(), "source": "local",
|
||||
"asset_type": asset_type,
|
||||
"price_limit": price_limit,
|
||||
"prev_close": prev_close,
|
||||
},
|
||||
pref_key="minute_batch_compress",
|
||||
)
|
||||
|
||||
# 本地不完整或无数据 → 从当前有效分钟源实时拉取
|
||||
live_df = kline_sync.fetch_minute_single(
|
||||
symbol, trade_date, asset_type=asset_type, capset=capset,
|
||||
)
|
||||
return _gzip_payload(
|
||||
request,
|
||||
{
|
||||
"symbol": symbol, "name": stock_name, "stock_info": stock_info,
|
||||
"date": str(trade_date), "rows": df.to_dicts(), "source": "local",
|
||||
"date": str(trade_date), "rows": live_df.to_dicts(),
|
||||
"source": "live" if not live_df.is_empty() else "none",
|
||||
"asset_type": asset_type,
|
||||
"price_limit": price_limit,
|
||||
"prev_close": prev_close,
|
||||
}
|
||||
|
||||
# 本地不完整或无数据 → 从 TickFlow 实时拉取
|
||||
live_df = kline_sync.fetch_minute_single(symbol, trade_date, asset_type=asset_type)
|
||||
return {
|
||||
"symbol": symbol, "name": stock_name, "stock_info": stock_info,
|
||||
"date": str(trade_date), "rows": live_df.to_dicts(),
|
||||
"source": "live" if not live_df.is_empty() else "none",
|
||||
"asset_type": asset_type,
|
||||
"price_limit": price_limit,
|
||||
"prev_close": prev_close,
|
||||
}
|
||||
},
|
||||
pref_key="minute_batch_compress",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/sync")
|
||||
@@ -1055,7 +1138,7 @@ async def sync_minute(request: Request):
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
from app.services.pipeline_jobs import JobCancelledError, job_store, release_run_slot, try_acquire_run_slot
|
||||
from app.services.pipeline_jobs import JobCancelledError, job_store, release_run_slot, run_with_capacity, try_acquire_run_slot
|
||||
from app.api.data import invalidate_storage_cache
|
||||
from app.services.preferences import get_minute_sync_days
|
||||
from app.tickflow.capabilities import Cap
|
||||
@@ -1092,7 +1175,6 @@ async def sync_minute(request: Request):
|
||||
job_store.progress(job_id, stage, pct, msg)
|
||||
|
||||
try:
|
||||
job_store.start(job_id)
|
||||
progress("sync_minute", 5, "解析标的池…")
|
||||
universe = sorted(set(get_pool("watchlist")) | set(get_pool("CN_Equity_A")))
|
||||
# 补充 instruments 全量标的,覆盖北交所、新股等
|
||||
@@ -1125,7 +1207,7 @@ async def sync_minute(request: Request):
|
||||
on_chunk_done=_on_chunk,
|
||||
)
|
||||
|
||||
written = await loop.run_in_executor(_long_task_executor, _run)
|
||||
written = await loop.run_in_executor(_long_task_executor, run_with_capacity, job_id, _run)
|
||||
|
||||
# 刷新视图
|
||||
from app.jobs.daily_pipeline import _refresh_single_view
|
||||
@@ -1261,7 +1343,7 @@ async def extend_history(request: Request):
|
||||
raise HTTPException(status_code=403, detail="需要 Pro+ 权限 (batch K-line)")
|
||||
|
||||
from app.services.extend_history import run_extend_history
|
||||
from app.services.pipeline_jobs import JobCancelledError, job_store, release_run_slot, try_acquire_run_slot
|
||||
from app.services.pipeline_jobs import JobCancelledError, job_store, release_run_slot, run_with_capacity, try_acquire_run_slot
|
||||
from app.api.data import invalidate_storage_cache
|
||||
|
||||
job_id, is_new = job_store.create()
|
||||
@@ -1280,9 +1362,8 @@ async def extend_history(request: Request):
|
||||
stage_pct=stage_pct, skip_log=skip_log)
|
||||
|
||||
try:
|
||||
job_store.start(job_id)
|
||||
result = await loop.run_in_executor(
|
||||
_long_task_executor,
|
||||
_long_task_executor, run_with_capacity, job_id,
|
||||
lambda: run_extend_history(repo, capset, value, unit, on_progress=progress),
|
||||
)
|
||||
if "error" in result:
|
||||
@@ -1343,7 +1424,7 @@ async def repair_daily(request: Request):
|
||||
raise HTTPException(status_code=403, detail="需要 Pro+ 权限 (batch K-line)")
|
||||
|
||||
from app.services.repair_daily import run_repair_daily
|
||||
from app.services.pipeline_jobs import JobCancelledError, job_store, release_run_slot, try_acquire_run_slot
|
||||
from app.services.pipeline_jobs import JobCancelledError, job_store, release_run_slot, run_with_capacity, try_acquire_run_slot
|
||||
from app.api.data import invalidate_storage_cache
|
||||
|
||||
job_id, is_new = job_store.create()
|
||||
@@ -1370,8 +1451,7 @@ async def repair_daily(request: Request):
|
||||
return run_repair_daily(repo, capset, start_date, on_progress=progress)
|
||||
|
||||
try:
|
||||
job_store.start(job_id)
|
||||
result = await loop.run_in_executor(_long_task_executor, _run)
|
||||
result = await loop.run_in_executor(_long_task_executor, run_with_capacity, job_id, _run)
|
||||
if "error" in result:
|
||||
job_store.fail(job_id, result["error"])
|
||||
else:
|
||||
@@ -1406,7 +1486,7 @@ async def rebuild_enriched(request: Request):
|
||||
try:
|
||||
repo = request.app.state.repo
|
||||
|
||||
from app.services.pipeline_jobs import JobCancelledError, job_store, release_run_slot, try_acquire_run_slot
|
||||
from app.services.pipeline_jobs import JobCancelledError, job_store, release_run_slot, run_with_capacity, try_acquire_run_slot
|
||||
from app.api.data import invalidate_storage_cache
|
||||
|
||||
job_id, is_new = job_store.create()
|
||||
@@ -1425,7 +1505,6 @@ async def rebuild_enriched(request: Request):
|
||||
stage_pct=stage_pct, skip_log=skip_log)
|
||||
|
||||
try:
|
||||
job_store.start(job_id)
|
||||
progress("rebuild_enriched", 10, "全量计算 enriched…")
|
||||
from app.indicators.pipeline import run_pipeline
|
||||
|
||||
@@ -1436,7 +1515,7 @@ async def rebuild_enriched(request: Request):
|
||||
stage_pct=int(100 * cur / tot), skip_log=True)
|
||||
|
||||
written = await loop.run_in_executor(
|
||||
_long_task_executor,
|
||||
_long_task_executor, run_with_capacity, job_id,
|
||||
lambda: run_pipeline(on_batch_done=_batch_progress),
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,131 @@
|
||||
"""批次登记 API — 薄"批次"页 (持仓提醒), 只做胶水, 不含会计语义。
|
||||
|
||||
映射/校验/持久化在 strategy.lots 域; 写完派生规则后复用 monitor_rules 的 _sync_engine 同步引擎。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import secrets
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.strategy import lots as lots_domain
|
||||
from app.strategy import monitor_rules
|
||||
|
||||
router = APIRouter(prefix="/api/lots", tags=["lots"])
|
||||
|
||||
# 批次 + 派生规则 + 引擎重载的跨请求互斥; 规则全部校验通过才落盘, 避免半成品 (镜像 watchlist 服务层)。
|
||||
_write_lock = threading.Lock()
|
||||
|
||||
|
||||
def _data_dir(request: Request) -> Path:
|
||||
return request.app.state.repo.store.data_dir
|
||||
|
||||
|
||||
def _resolve_asset_type(request: Request, symbol: str) -> str:
|
||||
"""按 symbol 解析资产类型 (stock/etf); 解析失败默认 stock (fail-safe)。"""
|
||||
repo = getattr(request.app.state, "repo", None)
|
||||
try:
|
||||
return repo.resolve_asset_type(symbol) if repo is not None else "stock"
|
||||
except Exception:
|
||||
# 回退为 stock 会让 etf 批次的止盈止损规则落入错误的监控轮, 必须留痕排查
|
||||
logging.getLogger(__name__).warning(
|
||||
"resolve_asset_type failed for %s, falling back to stock", symbol, exc_info=True
|
||||
)
|
||||
return "stock"
|
||||
|
||||
|
||||
class LotModel(BaseModel):
|
||||
id: str | None = None
|
||||
symbol: str
|
||||
qty: float = 0
|
||||
cost_price: float = 0
|
||||
buy_date: str | None = None
|
||||
target_pct: float = 0
|
||||
stop_pct: float = 0
|
||||
remind_date: str | None = None
|
||||
lead_days: int = 1
|
||||
|
||||
|
||||
def _reload_engine(request: Request) -> None:
|
||||
"""批次规则保存/删除后重载引擎 — 复用监控规则 API 的共享重载 (含指数纠正)。"""
|
||||
from app.api.monitor_rules import _sync_engine
|
||||
|
||||
_sync_engine(request)
|
||||
|
||||
|
||||
def sync_lot(request: Request, lot: dict) -> None:
|
||||
"""写批次文件 + 同步其两条派生监控规则 + 重载引擎。
|
||||
|
||||
派生规则继承用户默认推送渠道 (webhook_default_channels), 否则批次告警会静默只走应用内。
|
||||
"""
|
||||
from app.services import preferences
|
||||
|
||||
data_dir = _data_dir(request)
|
||||
with _write_lock:
|
||||
default_channels = preferences.get_webhook_default_channels()
|
||||
# ETF/指数等资产类型解析 (止盈止损价格规则须走对应资产监控轮才会触发)
|
||||
asset_type = _resolve_asset_type(request, lot["symbol"])
|
||||
price_rule, date_rule = lots_domain.lot_to_rules(lot)
|
||||
rules_to_write: list[dict] = []
|
||||
rules_to_delete: list[str] = []
|
||||
for rid, rule in ((f"{lot['id']}_p", price_rule), (f"{lot['id']}_d", date_rule)):
|
||||
if rule is None:
|
||||
rules_to_delete.append(rid)
|
||||
continue
|
||||
rule["asset_type"] = asset_type
|
||||
rule.setdefault("webhook_channels", list(default_channels))
|
||||
# 保留旧 created_at, 避免编辑批次后派生规则在监控中心列表跳位
|
||||
existing = monitor_rules.load_one(data_dir, rid)
|
||||
if existing and existing.get("created_at"):
|
||||
rule["created_at"] = existing["created_at"]
|
||||
try:
|
||||
monitor_rules.validate(rule)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e)) from e
|
||||
rules_to_write.append(monitor_rules.normalize(rule))
|
||||
lots_domain.save_one(data_dir, lot)
|
||||
for rid in rules_to_delete:
|
||||
monitor_rules.delete_one(data_dir, rid)
|
||||
for rule in rules_to_write:
|
||||
monitor_rules.save_one(data_dir, rule)
|
||||
_reload_engine(request)
|
||||
|
||||
|
||||
@router.get("")
|
||||
def list_lots(request: Request):
|
||||
return {"lots": lots_domain.load_all(_data_dir(request))}
|
||||
|
||||
|
||||
@router.post("")
|
||||
def upsert_lot(lot_in: LotModel, request: Request):
|
||||
"""新建/更新一个批次。id 缺省时服务端生成 (紧凑, 保证 {id}_p/_d 规则 id ≤ 40 字符)。"""
|
||||
lot = lot_in.model_dump()
|
||||
if not lot.get("id"):
|
||||
lot["id"] = f"lot_{int(time.time() * 1000):x}_{secrets.token_hex(2)}"
|
||||
try:
|
||||
lots_domain.validate_lot(lot)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e)) from e
|
||||
lot = lots_domain.normalize_lot(lot)
|
||||
sync_lot(request, lot)
|
||||
return {"ok": True, "lot": lot}
|
||||
|
||||
|
||||
@router.delete("/{lot_id}")
|
||||
def delete_lot(lot_id: str, request: Request):
|
||||
if not monitor_rules.ID_RE.match(lot_id):
|
||||
raise HTTPException(status_code=400, detail="批次 id 非法")
|
||||
data_dir = _data_dir(request)
|
||||
with _write_lock:
|
||||
deleted = lots_domain.delete_one(data_dir, lot_id)
|
||||
# 两条派生规则都要删 (用 or 会短路跳过第二条)
|
||||
deleted_p = monitor_rules.delete_one(data_dir, f"{lot_id}_p")
|
||||
deleted_d = monitor_rules.delete_one(data_dir, f"{lot_id}_d")
|
||||
if deleted or deleted_p or deleted_d:
|
||||
_reload_engine(request)
|
||||
return {"ok": True}
|
||||
@@ -19,7 +19,7 @@ from fastapi import APIRouter, HTTPException, Query, Request
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.services import auction_benchmark, dragon_tiger, market_recap_reports
|
||||
from app.services import auction_benchmark, dragon_tiger, market_recap_reports, preferences
|
||||
from app.services.market_recap import recap_market_stream
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -122,6 +122,7 @@ class SaveReportRequest(BaseModel):
|
||||
summary: str = ""
|
||||
emotion_score: int | None = None
|
||||
emotion_label: str = ""
|
||||
push: bool = False # 是否显式外发推送(manual 模式下需显式传 true)
|
||||
|
||||
|
||||
@router.get("/reports")
|
||||
@@ -132,7 +133,7 @@ def list_reports(request: Request):
|
||||
|
||||
@router.post("/reports")
|
||||
def save_report(request: Request, req: SaveReportRequest):
|
||||
"""保存一条复盘报告。"""
|
||||
"""保存一条复盘报告。req.push=True 或 review_push_mode=auto 时才推送到外部渠道。"""
|
||||
report = market_recap_reports.save_report({
|
||||
"as_of": req.as_of,
|
||||
"focus": req.focus,
|
||||
@@ -141,13 +142,14 @@ def save_report(request: Request, req: SaveReportRequest):
|
||||
"emotion_score": req.emotion_score,
|
||||
"emotion_label": req.emotion_label,
|
||||
})
|
||||
# 推送到飞书(可选): 与定时复盘共用同一开关 review_push_enabled 与 _maybe_push_review。
|
||||
# 推送门控: manual 模式需显式 push=True; auto 模式保持归档即推。
|
||||
# 内部 try/except 静默降级, 不影响归档返回值。
|
||||
from app.jobs.daily_pipeline import _maybe_push_review
|
||||
_maybe_push_review(req.content, {
|
||||
"as_of": req.as_of,
|
||||
"emotion_label": req.emotion_label,
|
||||
})
|
||||
if req.push or preferences.get_review_push_mode() == "auto":
|
||||
from app.jobs.daily_pipeline import _maybe_push_review
|
||||
_maybe_push_review(req.content, {
|
||||
"as_of": req.as_of,
|
||||
"emotion_label": req.emotion_label,
|
||||
})
|
||||
return {"ok": True, "report": report}
|
||||
|
||||
|
||||
|
||||
+125
-3
@@ -13,7 +13,6 @@ from fastapi import APIRouter, Header, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
from sse_starlette.sse import EventSourceResponse
|
||||
|
||||
from app.backtest.factor import FACTOR_COLUMNS
|
||||
from app.backtest.mining import (
|
||||
MAX_BEAM_WIDTH,
|
||||
MAX_COMBINATION_SIZE,
|
||||
@@ -21,6 +20,7 @@ from app.backtest.mining import (
|
||||
evaluate_candidate_gate,
|
||||
)
|
||||
from app.enriched_generation import EnrichedGenerationUnavailableError
|
||||
from app.factors.registry import factor_columns_view
|
||||
from app.services import preferences
|
||||
from app.services.mining_jobs import (
|
||||
RUN_STATUSES,
|
||||
@@ -31,6 +31,7 @@ from app.services.mining_jobs import (
|
||||
MiningRunValidationError,
|
||||
)
|
||||
from app.services.mining_preflight import (
|
||||
enriched_partition_dates,
|
||||
mining_availability,
|
||||
require_mining_availability,
|
||||
)
|
||||
@@ -40,7 +41,9 @@ from app.services.mining_schedule import (
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/api/backtest/mining", tags=["backtest"])
|
||||
_FACTOR_IDS = frozenset(str(item["id"]) for item in FACTOR_COLUMNS)
|
||||
# 校验时动态读取 (含运行期注册的自定义/复合因子)
|
||||
def _known_factor_ids() -> frozenset[str]:
|
||||
return frozenset(str(item["id"]) for item in factor_columns_view())
|
||||
_MAX_ARTIFACT_BYTES = 64 * 1024 * 1024
|
||||
_SSE_POLL_SECONDS = 0.5
|
||||
_SSE_HEARTBEAT_SECONDS = 15.0
|
||||
@@ -87,7 +90,7 @@ class MiningStartRequest(BaseModel):
|
||||
@field_validator("factor_names")
|
||||
@classmethod
|
||||
def _known_factors(cls, values: list[str]) -> list[str]:
|
||||
unknown = sorted(set(values) - _FACTOR_IDS)
|
||||
unknown = sorted(set(values) - _known_factor_ids())
|
||||
if unknown:
|
||||
raise ValueError(f"unknown mining factors: {unknown}")
|
||||
return values
|
||||
@@ -119,6 +122,38 @@ class MiningSchedulePatch(BaseModel):
|
||||
mining_budget_profile: Literal["balanced", "strict"] | None = None
|
||||
|
||||
|
||||
class MiningAutoStartRequest(BaseModel):
|
||||
"""自动挖掘: 因子池由 L1 统计筛选自动生成, 不接受手动指定。"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid", strict=True)
|
||||
|
||||
asset_type: Literal["stock", "etf"] = "stock"
|
||||
start: date | None = None
|
||||
end: date | None = None
|
||||
budget_profile: Literal["exploratory", "balanced", "strict"] = "balanced"
|
||||
commission_pct: float = Field(0.0002, ge=0.0, le=0.05, allow_inf_nan=False)
|
||||
stamp_tax_pct: float = Field(0.0005, ge=0.0, le=0.05, allow_inf_nan=False)
|
||||
slippage_bps: float = Field(5.0, ge=0.0, le=1000.0, allow_inf_nan=False)
|
||||
correlation_threshold: float = Field(0.75, gt=0.0, le=1.0, allow_inf_nan=False)
|
||||
force: bool = False
|
||||
|
||||
@field_validator("start", "end", mode="before")
|
||||
@classmethod
|
||||
def _iso_dates(cls, value: Any) -> Any:
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return date.fromisoformat(value)
|
||||
except ValueError as exc:
|
||||
raise ValueError("dates must use ISO YYYY-MM-DD format") from exc
|
||||
return value
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _date_range(self) -> MiningAutoStartRequest:
|
||||
if self.start is not None and self.end is not None and self.start > self.end:
|
||||
raise ValueError("start must not be after end")
|
||||
return self
|
||||
|
||||
|
||||
@router.get("/availability")
|
||||
def get_availability(
|
||||
request: Request,
|
||||
@@ -237,6 +272,93 @@ def cancel_run(run_id: str, request: Request) -> dict[str, Any]:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/auto")
|
||||
def start_auto_run(payload: MiningAutoStartRequest, request: Request) -> dict[str, Any]:
|
||||
"""自动挖掘: L1 统计筛选全量因子 → 达标池 → 复用挖掘任务管理启动嵌套样本外验证。
|
||||
|
||||
筛选结果随请求持久化 (request.auto_screening), 供结果页展示达标因子清单与
|
||||
失败原因分布; 无达标因子时返回 started=false 而不是报错。
|
||||
"""
|
||||
from app.services.auto_mining import screen_all_factors
|
||||
|
||||
manager = _manager(request)
|
||||
data_dir = request.app.state.repo.store.data_dir
|
||||
try:
|
||||
require_mining_availability(
|
||||
data_dir,
|
||||
asset_type=payload.asset_type,
|
||||
budget_profile=payload.budget_profile,
|
||||
start=payload.start,
|
||||
end=payload.end,
|
||||
)
|
||||
engine = getattr(request.app.state, "backtest_engine", None)
|
||||
if engine is None:
|
||||
from app.backtest.engine import BacktestEngine
|
||||
|
||||
engine = BacktestEngine(request.app.state.repo)
|
||||
request.app.state.backtest_engine = engine
|
||||
all_dates = enriched_partition_dates(data_dir, payload.asset_type)
|
||||
screen_end = payload.end or (all_dates[-1] if all_dates else date.today())
|
||||
screening = screen_all_factors(
|
||||
engine,
|
||||
asset_type=payload.asset_type,
|
||||
start=payload.start,
|
||||
end=screen_end,
|
||||
profile=payload.budget_profile,
|
||||
)
|
||||
except (MiningRunValidationError, ValueError) as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
except EnrichedGenerationUnavailableError as exc:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="行情数据正在更新(enriched 发布中), 请等数据更新完成后再开始挖掘",
|
||||
) from exc
|
||||
|
||||
if not screening["pool"]:
|
||||
return {"started": False, "reason": "no_qualified_factors", "screening": screening}
|
||||
|
||||
worker_request = {
|
||||
"factor_names": screening["pool"],
|
||||
"strategy_ids": [],
|
||||
"symbols": None,
|
||||
"asset_type": payload.asset_type,
|
||||
"start": payload.start.isoformat() if payload.start else None,
|
||||
"end": payload.end.isoformat() if payload.end else None,
|
||||
"budget_profile": payload.budget_profile,
|
||||
"commission_pct": payload.commission_pct,
|
||||
"stamp_tax_pct": payload.stamp_tax_pct,
|
||||
"slippage_bps": payload.slippage_bps,
|
||||
"correlation_threshold": payload.correlation_threshold,
|
||||
"max_combination_factors": 4,
|
||||
"beam_width": 12,
|
||||
"max_finalists": MAX_FINALISTS,
|
||||
"auto": True,
|
||||
"auto_screening": screening,
|
||||
}
|
||||
try:
|
||||
fingerprint = build_data_fingerprint(
|
||||
request.app.state.repo,
|
||||
request.app.state,
|
||||
worker_request,
|
||||
)
|
||||
manifest = manager.start(
|
||||
worker_request,
|
||||
fingerprint,
|
||||
force=payload.force,
|
||||
source="auto",
|
||||
)
|
||||
except (MiningRunValidationError, ValueError) as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
except EnrichedGenerationUnavailableError as exc:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="行情数据正在更新(enriched 发布中), 请等数据更新完成后再开始挖掘",
|
||||
) from exc
|
||||
except MiningRunStoreError as exc:
|
||||
raise HTTPException(status_code=500, detail="failed to persist mining run") from exc
|
||||
return {"started": True, "run": _project_run(manager.store, manifest), "screening": screening}
|
||||
|
||||
|
||||
@router.get("/runs/{run_id}/result")
|
||||
def get_result(run_id: str, request: Request) -> dict[str, Any]:
|
||||
store = _manager(request).store
|
||||
|
||||
@@ -97,10 +97,13 @@ class RuleModel(BaseModel):
|
||||
conditions: list[ConditionModel] = []
|
||||
logic: str = "and" # and | or
|
||||
cooldown_seconds: int = 3600
|
||||
# date 类型 (日期提醒): 纯日历窗口, 无 conditions
|
||||
remind_date: str | None = None # YYYY-MM-DD
|
||||
lead_days: int = 0 # 提前 N 天进入提醒窗口
|
||||
severity: str = "info" # info | warn | critical
|
||||
webhook_url: str = "" # Webhook 推送地址 (推送到 QMT 等外部软件, 待定)
|
||||
webhook_enabled: bool = False # 兼容老规则 (已由 webhook_channels 取代, 仅做向后兼容读)
|
||||
webhook_channels: list[str] = [] # 命中时推送的外部渠道 (合法值 'feishu' | 'wecom')
|
||||
webhook_channels: list[str] = [] # 合法值: feishu | wecom | custom | email
|
||||
message: str = ""
|
||||
# abnormal 专属 (异动边缘监控): any | 3d | 10d | 30d
|
||||
abnormal_window: str = "any"
|
||||
@@ -166,6 +169,7 @@ def get_options(request: Request):
|
||||
{"key": "abnormal", "label": "异动监控"},
|
||||
{"key": "sector", "label": "板块监控"},
|
||||
{"key": "volume_delta", "label": "轮询放量"},
|
||||
{"key": "date", "label": "日期提醒"},
|
||||
],
|
||||
"scopes": [
|
||||
{"key": "symbols", "label": "指定标的"},
|
||||
@@ -288,6 +292,9 @@ def save_rule(req: RuleModel, request: Request):
|
||||
raise HTTPException(status_code=400, detail=str(e)) from e
|
||||
# 编辑现有规则时, 保留原 created_at (避免按时间排序时位置跳动)
|
||||
existing = monitor_rules.load_one(_data_dir(request), rule["id"])
|
||||
# 批次派生规则由「持仓提醒」页托管, 监控中心只读 (启停/改/删均回持仓页)
|
||||
if existing and existing.get("lot_id"):
|
||||
raise HTTPException(status_code=409, detail="该规则由「持仓提醒」页托管, 请在持仓提醒页修改")
|
||||
if existing and existing.get("created_at"):
|
||||
rule["created_at"] = existing["created_at"]
|
||||
try:
|
||||
@@ -344,6 +351,10 @@ def save_rule(req: RuleModel, request: Request):
|
||||
def delete_rule(rule_id: str, request: Request):
|
||||
if not monitor_rules.ID_RE.match(rule_id):
|
||||
raise HTTPException(status_code=400, detail="规则 id 非法")
|
||||
# 批次派生规则由「持仓提醒」页托管, 删除需在持仓页操作 (级联清理派生规则)
|
||||
existing = monitor_rules.load_one(_data_dir(request), rule_id)
|
||||
if existing and existing.get("lot_id"):
|
||||
raise HTTPException(status_code=409, detail="该规则由「持仓提醒」页托管, 请在持仓提醒页删除批次")
|
||||
deleted = monitor_rules.delete_one(_data_dir(request), rule_id)
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=404, detail="规则不存在")
|
||||
|
||||
@@ -12,6 +12,7 @@ from app.services.pipeline_jobs import (
|
||||
JobCancelledError,
|
||||
job_store,
|
||||
release_run_slot,
|
||||
run_with_capacity,
|
||||
try_acquire_run_slot,
|
||||
)
|
||||
from app.api.data import invalidate_storage_cache
|
||||
@@ -52,7 +53,6 @@ async def run_now(request: Request) -> dict:
|
||||
# 管道运行期间暂停实时行情取数, 防止覆写同一批 parquet 竞态
|
||||
qs = getattr(request.app.state, "quote_service", None)
|
||||
try:
|
||||
job_store.start(job_id)
|
||||
loop = asyncio.get_event_loop()
|
||||
|
||||
def progress(stage: str, pct: int, msg: str, stage_pct: int | None = None,
|
||||
@@ -60,15 +60,17 @@ async def run_now(request: Request) -> dict:
|
||||
job_store.progress(job_id, stage, pct, msg, stage_pct=stage_pct, skip_log=skip_log)
|
||||
|
||||
def _run() -> dict:
|
||||
if qs:
|
||||
with qs.paused():
|
||||
return daily_pipeline.run_now(repo, capset, on_progress=progress)
|
||||
return daily_pipeline.run_now(repo, capset, on_progress=progress)
|
||||
try:
|
||||
if qs:
|
||||
with qs.paused():
|
||||
return daily_pipeline.run_now(repo, capset, on_progress=progress)
|
||||
return daily_pipeline.run_now(repo, capset, on_progress=progress)
|
||||
finally:
|
||||
repo.refresh_cache()
|
||||
|
||||
result = await loop.run_in_executor(_long_task_executor, _run)
|
||||
result = await loop.run_in_executor(_long_task_executor, run_with_capacity, job_id, _run)
|
||||
job_store.succeed(job_id, result)
|
||||
invalidate_storage_cache()
|
||||
repo.refresh_cache() # 刷新 Polars 缓存
|
||||
except JobCancelledError:
|
||||
# 已被 reap/手动取消终止: job 状态已由 terminate() 写为 failed,
|
||||
# 拉取线程在分块回调处自行退出, 这里无需(也无法)再写状态。
|
||||
|
||||
+178
-15
@@ -1,21 +1,23 @@
|
||||
"""Screener API。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import glob as _glob
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from dataclasses import asdict
|
||||
from dataclasses import asdict, replace
|
||||
from datetime import date, datetime
|
||||
from typing import Any, Optional
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Query, Request
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.config import settings
|
||||
from app.db_safe import is_valid_ext_ident, quote_ident
|
||||
from app.services import strategy_cache
|
||||
from app.services import strategy_cache, strategy_run_queue
|
||||
from app.services.screener import ScreenerService
|
||||
from app.strategy import config as strategy_config
|
||||
|
||||
@@ -379,6 +381,9 @@ def get_cached_summary(request: Request):
|
||||
sid: {
|
||||
"total": int(result.get("total") or 0),
|
||||
"as_of": result.get("as_of"),
|
||||
# 渐进式 run_all 写入的计算时间戳; 监控实时叠加/旧缓存无此字段 → None,
|
||||
# 前端视为新鲜 (有值即为最新一轮实时结果)
|
||||
"computed_at": result.get("computed_at"),
|
||||
}
|
||||
for sid, result in results.items()
|
||||
if isinstance(result, dict)
|
||||
@@ -497,6 +502,120 @@ def market_snapshot(request: Request):
|
||||
return {"as_of": str(as_of), "rows": rows}
|
||||
|
||||
|
||||
def _run_all_progressive(
|
||||
*,
|
||||
repo,
|
||||
engine,
|
||||
svc: ScreenerService,
|
||||
as_of,
|
||||
asset_type: str,
|
||||
timeframe: str,
|
||||
all_ids: list[str],
|
||||
params_map: dict,
|
||||
overrides_map: dict,
|
||||
first_return_s: float,
|
||||
t_total: float,
|
||||
) -> dict:
|
||||
"""run_all 渐进式执行: 快策略随响应先返回, 慢策略后台算完逐个落缓存。
|
||||
|
||||
执行全程在单飞执行器里 (见 services/strategy_run_queue.py): 相同请求
|
||||
搭车现有执行, 不同请求排队; HTTP 侧只轮询状态快照到首返时限。
|
||||
"""
|
||||
data_dir = repo.store.data_dir
|
||||
key = (asset_type, timeframe, str(as_of), tuple(sorted(all_ids)))
|
||||
ordered_ids = strategy_run_queue.order_strategy_ids(
|
||||
all_ids, strategy_run_queue.load_run_timings(data_dir)
|
||||
)
|
||||
|
||||
def job(handle: strategy_run_queue.StrategyRunHandle) -> None:
|
||||
context = svc.build_strategy_context(
|
||||
engine,
|
||||
as_of,
|
||||
ordered_ids,
|
||||
timeframe=timeframe,
|
||||
params_map=params_map,
|
||||
overrides_map=overrides_map,
|
||||
)
|
||||
# 逐策略 run_all 不会把矩阵回写 context.market → 每个矩阵策略都会重建
|
||||
# 全市场矩阵 (小服务器上单次数秒到十余秒)。这里按字段并集一次建好复用;
|
||||
# FakeEngine 等无该方法的实现跳过 (保持旧行为)。
|
||||
if getattr(context, "market", None) is None:
|
||||
build_matrix = getattr(engine, "build_shared_matrix", None)
|
||||
if callable(build_matrix):
|
||||
matrix = build_matrix(
|
||||
context,
|
||||
[(sid, engine.get(sid)) for sid in ordered_ids],
|
||||
params_map,
|
||||
overrides_map,
|
||||
)
|
||||
if matrix is not None:
|
||||
context = replace(context, market=matrix)
|
||||
all_results: dict[str, dict] = {}
|
||||
elapsed_map: dict[str, float] = {}
|
||||
for sid in ordered_ids:
|
||||
t0 = time.perf_counter()
|
||||
# 逐策略隔离: 单个策略崩溃 (如自定义代码的数据类型错误) 只记
|
||||
# 错误跳过, 不让整批剩余策略陪葬 — 其余策略照常算完落缓存。
|
||||
try:
|
||||
single = engine.run_all(
|
||||
context,
|
||||
params_map=params_map,
|
||||
overrides_map=overrides_map,
|
||||
strategy_ids=[sid],
|
||||
parallel=False,
|
||||
)
|
||||
result = single[sid]
|
||||
except Exception as e:
|
||||
logger.warning("run_all: 策略 %s 执行失败, 跳过: %s", sid, e, exc_info=True)
|
||||
handle.fail_one(sid, str(e))
|
||||
continue
|
||||
payload = {
|
||||
"total": result.total,
|
||||
"as_of": str(as_of),
|
||||
"rows": _safe(asdict(result)).get("rows", []),
|
||||
"computed_at": int(time.time() * 1000),
|
||||
}
|
||||
all_results[sid] = payload
|
||||
elapsed_map[sid] = (time.perf_counter() - t0) * 1000
|
||||
# 逐策略增量落盘 (write_cache 同日按 sid 合并), 前端轮询即可逐个看到
|
||||
try:
|
||||
strategy_cache.write_cache(data_dir, str(as_of), {sid: payload})
|
||||
except Exception:
|
||||
logger.warning("run_all 渐进写入缓存失败: %s", sid, exc_info=True)
|
||||
handle.complete(sid, {k: v for k, v in payload.items() if k != "rows"})
|
||||
# 收尾: 与旧版口径一致的整体重写 + 耗时落盘供下次排序
|
||||
if all_results:
|
||||
with contextlib.suppress(Exception):
|
||||
strategy_cache.write_cache(data_dir, str(as_of), all_results)
|
||||
strategy_run_queue.record_run_timings(data_dir, elapsed_map)
|
||||
|
||||
handle = strategy_run_queue.MANAGER.get_or_submit(key, ordered_ids, job)
|
||||
deadline = time.perf_counter() + first_return_s
|
||||
snap = handle.snapshot()
|
||||
while not snap["done"] and time.perf_counter() < deadline:
|
||||
time.sleep(0.2)
|
||||
snap = handle.snapshot()
|
||||
|
||||
done_results = snap["results"]
|
||||
if snap["error"] and not done_results:
|
||||
raise HTTPException(status_code=500, detail=snap["error"])
|
||||
logger.info(
|
||||
"run_all: first return %.1fms (%d done, %d pending)",
|
||||
(time.perf_counter() - t_total) * 1000,
|
||||
len(done_results),
|
||||
len(snap["pending"]),
|
||||
)
|
||||
return {
|
||||
"as_of": str(as_of),
|
||||
"results": done_results,
|
||||
"pending": snap["pending"],
|
||||
"errors": snap["errors"],
|
||||
"complete": snap["done"] and not snap["error"],
|
||||
"error": snap["error"],
|
||||
"started_at": snap["started_at_ms"],
|
||||
}
|
||||
|
||||
|
||||
@router.post("/run_all")
|
||||
def run_all(request: Request, body: Optional[dict] = None):
|
||||
"""批量运行指定策略;注册、路由和执行均由 StrategyEngine 负责。"""
|
||||
@@ -516,7 +635,15 @@ def run_all(request: Request, body: Optional[dict] = None):
|
||||
# 解析日期
|
||||
raw_date = body.get("as_of")
|
||||
if raw_date:
|
||||
as_of = date_type.fromisoformat(str(raw_date)) if isinstance(raw_date, str) else raw_date
|
||||
# 与 /custom、/preset 的 `as_of: date` 同口径: 只收 ISO 日期字符串。
|
||||
# 非字符串原样透传会让 str(as_of) 把 "20260904" 之类写进 strategy_cache.json,
|
||||
# 与其它入口写的 "2026-09-04" 不是同一格式, 后续按 as_of 比对缓存永远失配。
|
||||
if not isinstance(raw_date, str):
|
||||
raise HTTPException(status_code=400, detail="as_of 必须是 YYYY-MM-DD 日期字符串")
|
||||
try:
|
||||
as_of = date_type.fromisoformat(raw_date)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=f"日期格式错误: {e}") from e
|
||||
else:
|
||||
as_of = svc.latest_date()
|
||||
if not as_of:
|
||||
@@ -556,6 +683,26 @@ def run_all(request: Request, body: Optional[dict] = None):
|
||||
for sid in all_ids
|
||||
}
|
||||
overrides_map = {sid: all_overrides.get(sid, {}) for sid in all_ids}
|
||||
|
||||
# 渐进式返回 (页面首屏路径): 按历史耗时升序执行, 首返时限内算完的随响应
|
||||
# 返回, 慢策略转后台继续算并逐个写入策略缓存, 前端轮询 cached-summary 点亮。
|
||||
# 仅日线 + summary_only (策略页卡片) 启用; 分钟/明细请求保持整段阻塞。
|
||||
first_return_s = settings.strategy_run_all_first_return_s
|
||||
if body.get("summary_only") and timeframe == "1d" and first_return_s > 0:
|
||||
return _run_all_progressive(
|
||||
repo=repo,
|
||||
engine=engine,
|
||||
svc=svc,
|
||||
as_of=as_of,
|
||||
asset_type=asset_type,
|
||||
timeframe=timeframe,
|
||||
all_ids=all_ids,
|
||||
params_map=params_map,
|
||||
overrides_map=overrides_map,
|
||||
first_return_s=first_return_s,
|
||||
t_total=t_total,
|
||||
)
|
||||
|
||||
try:
|
||||
context = svc.build_strategy_context(
|
||||
engine,
|
||||
@@ -786,32 +933,48 @@ def limit_ladder(
|
||||
if ext_specs:
|
||||
db = repo.store.db
|
||||
data_dir = repo.store.data_dir
|
||||
from app.api.ext_data import _read_ext_dataframe
|
||||
from app.services.ext_data import ExtConfigStore
|
||||
|
||||
ext_store = ExtConfigStore(data_dir)
|
||||
configs = {c.id: c for c in ext_store.load_all()}
|
||||
|
||||
def _dedup_ext(frame: pl.DataFrame, field: str, out_col: str) -> pl.DataFrame | None:
|
||||
"""(symbol, 字段) 两列并按 symbol 去重; 缺列时返回 None。"""
|
||||
if frame.is_empty() or "symbol" not in frame.columns or field not in frame.columns:
|
||||
return None
|
||||
return (
|
||||
frame
|
||||
.select(["symbol", field])
|
||||
.unique(subset=["symbol"], keep="last")
|
||||
.rename({field: out_col})
|
||||
)
|
||||
|
||||
for config_id, field_name in ext_specs:
|
||||
view_name = f"ext_{config_id}"
|
||||
ext_col_name = f"{config_id}__{field_name}"
|
||||
try:
|
||||
ext_df = pl.from_arrow(db.query(
|
||||
f"SELECT symbol, {quote_ident(field_name)} FROM {view_name}"
|
||||
).arrow())
|
||||
if not ext_df.is_empty() and "symbol" in ext_df.columns:
|
||||
ext_df = ext_df.rename({field_name: ext_col_name})
|
||||
df = df.join(ext_df.select(["symbol", ext_col_name]), on="symbol", how="left")
|
||||
# 扩展时序数据必须只取最新分区; 否则一个 symbol 会按历史分区数被 JOIN 放大
|
||||
# (ext_{id} 视图覆盖 timeseries/**), 与自选股列表同口径。
|
||||
cfg = configs.get(config_id)
|
||||
if cfg:
|
||||
ext_df, _ = _read_ext_dataframe(cfg, data_dir)
|
||||
else:
|
||||
ext_df = pl.from_arrow(db.query(
|
||||
f"SELECT symbol, {quote_ident(field_name)} FROM {view_name}"
|
||||
).arrow())
|
||||
joined = _dedup_ext(ext_df, field_name, ext_col_name)
|
||||
if joined is not None:
|
||||
df = df.join(joined, on="symbol", how="left")
|
||||
ext_col_names.append(ext_col_name)
|
||||
except Exception:
|
||||
cfg = configs.get(config_id)
|
||||
if cfg:
|
||||
try:
|
||||
from app.api.ext_data import _parquet_glob
|
||||
glob = _parquet_glob(cfg, data_dir)
|
||||
ext_df = pl.read_parquet(glob)
|
||||
if not ext_df.is_empty() and "symbol" in ext_df.columns and field_name in ext_df.columns:
|
||||
ext_df = ext_df.select(["symbol", field_name]).rename({field_name: ext_col_name})
|
||||
df = df.join(ext_df, on="symbol", how="left")
|
||||
ext_df, _ = _read_ext_dataframe(cfg, data_dir)
|
||||
joined = _dedup_ext(ext_df, field_name, ext_col_name)
|
||||
if joined is not None:
|
||||
df = df.join(joined, on="symbol", how="left")
|
||||
ext_col_names.append(ext_col_name)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
+171
-9
@@ -537,6 +537,10 @@ def get_preferences() -> dict:
|
||||
"feishu_webhook_url": preferences.get_feishu_webhook_url(),
|
||||
"feishu_webhook_secret": preferences.get_feishu_webhook_secret(),
|
||||
"wecom_webhook_url": preferences.get_wecom_webhook_url(),
|
||||
"custom_webhook_url": preferences.get_custom_webhook_url(),
|
||||
"custom_webhook_secret_set": bool(secrets_store.get_custom_webhook_secret()),
|
||||
"email_smtp_config": preferences.get_email_smtp_config(),
|
||||
"email_smtp_password_set": bool(secrets_store.get_email_smtp_password()),
|
||||
"wecom_bot_id": preferences.get_wecom_bot_id(),
|
||||
"wecom_bot_secret": preferences.get_wecom_bot_secret(),
|
||||
"wecom_bot_enabled": preferences.get_wecom_bot_enabled(),
|
||||
@@ -553,6 +557,7 @@ def get_preferences() -> dict:
|
||||
"depth_finalize_time": preferences.get_depth_finalize_time(),
|
||||
"review_schedule": preferences.get_review_schedule(),
|
||||
"review_push_channels": preferences.get_review_push_channels(),
|
||||
"review_push_mode": preferences.get_review_push_mode(),
|
||||
**preferences.get_mining_schedule(),
|
||||
}
|
||||
|
||||
@@ -587,6 +592,7 @@ def get_capability_matrix() -> dict:
|
||||
"realtime_data_provider": preferences.get_realtime_data_provider(),
|
||||
"daily_data_provider": preferences.get_daily_data_provider(),
|
||||
"minute_data_provider": preferences.get_minute_data_provider(),
|
||||
"full_minute_data_provider": preferences.get_full_minute_data_provider(),
|
||||
"depth5_data_provider": preferences.get_depth5_data_provider(),
|
||||
"adj_factor_provider": preferences.get_adj_factor_provider(),
|
||||
"financial_data_provider": preferences.get_financial_provider(),
|
||||
@@ -805,7 +811,7 @@ def update_data_source_job_timeouts(req: DataSourceJobTimeoutPrefs) -> dict:
|
||||
|
||||
@router.put("/preferences/minute-batch-compress")
|
||||
def update_minute_batch_compress(req: MinuteBatchCompressPrefs) -> dict:
|
||||
"""保存分时批量响应的 gzip 传输压缩开关。逐请求即时读取, 保存后立即生效。"""
|
||||
"""保存分时详情与批量响应的 gzip 传输压缩开关。逐请求即时读取, 保存后立即生效。"""
|
||||
from app.services import preferences
|
||||
preferences.save({"minute_batch_compress": req.minute_batch_compress})
|
||||
return {"minute_batch_compress": preferences.get_minute_batch_compress()}
|
||||
@@ -813,7 +819,7 @@ def update_minute_batch_compress(req: MinuteBatchCompressPrefs) -> dict:
|
||||
|
||||
@router.put("/preferences/daily-batch-compress")
|
||||
def update_daily_batch_compress(req: DailyBatchCompressPrefs) -> dict:
|
||||
"""保存日K批量响应的 gzip 传输压缩开关 (与分时独立)。逐请求即时读取。"""
|
||||
"""保存日K详情与批量响应的 gzip 传输压缩开关 (与分时独立)。逐请求即时读取。"""
|
||||
from app.services import preferences
|
||||
preferences.save({"daily_batch_compress": req.daily_batch_compress})
|
||||
return {"daily_batch_compress": preferences.get_daily_batch_compress()}
|
||||
@@ -1232,6 +1238,153 @@ def update_wecom_webhook(req: WecomWebhookPrefsIn) -> dict:
|
||||
return {"wecom_webhook_url": saved_url}
|
||||
|
||||
|
||||
class CustomWebhookPrefsIn(BaseModel):
|
||||
url: str
|
||||
# None preserves the stored secret; an explicit empty string clears it.
|
||||
secret: str | None = None
|
||||
|
||||
|
||||
@router.put("/preferences/custom-webhook")
|
||||
def update_custom_webhook(req: CustomWebhookPrefsIn) -> dict:
|
||||
"""Configure the generic third-party JSON webhook and optional HMAC secret."""
|
||||
from app.services import preferences, webhook_adapter
|
||||
|
||||
url = (req.url or "").strip()
|
||||
if url and not webhook_adapter.is_valid_custom_url(url):
|
||||
raise HTTPException(status_code=400, detail="Webhook 地址必须是完整的 HTTP(S) URL")
|
||||
saved_url = preferences.set_custom_webhook_url(url)
|
||||
if not saved_url:
|
||||
secrets_store.set_custom_webhook_secret("")
|
||||
elif req.secret is not None:
|
||||
secrets_store.set_custom_webhook_secret(req.secret)
|
||||
return {
|
||||
"custom_webhook_url": saved_url,
|
||||
"custom_webhook_secret_set": bool(secrets_store.get_custom_webhook_secret()),
|
||||
}
|
||||
|
||||
|
||||
class EmailSmtpPrefsIn(BaseModel):
|
||||
host: str
|
||||
port: int = Field(default=465, ge=1, le=65535)
|
||||
security: Literal["ssl", "starttls", "none"] = "ssl"
|
||||
username: str = ""
|
||||
# None preserves the stored password; an explicit empty string clears it.
|
||||
password: str | None = None
|
||||
from_address: str = ""
|
||||
to_addresses: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
@router.put("/preferences/email-smtp")
|
||||
def update_email_smtp(req: EmailSmtpPrefsIn) -> dict:
|
||||
"""Configure the SMTP transport shared by monitor alerts and review reports."""
|
||||
from app.services import email_adapter, preferences
|
||||
|
||||
host = (req.host or "").strip()
|
||||
username = (req.username or "").strip()
|
||||
from_address = (req.from_address or username).strip()
|
||||
recipients = list(dict.fromkeys(item.strip() for item in req.to_addresses if item.strip()))
|
||||
if host:
|
||||
if not from_address or not email_adapter.is_valid_email(from_address):
|
||||
raise HTTPException(status_code=400, detail="请填写有效的发件人邮箱")
|
||||
if not recipients or any(not email_adapter.is_valid_email(item) for item in recipients):
|
||||
raise HTTPException(status_code=400, detail="请至少填写一个有效的收件人邮箱")
|
||||
effective_password = (
|
||||
secrets_store.get_email_smtp_password()
|
||||
if req.password is None
|
||||
else req.password
|
||||
)
|
||||
if username and not effective_password:
|
||||
raise HTTPException(status_code=400, detail="已填写 SMTP 登录用户名, 请同时填写密码或授权码")
|
||||
else:
|
||||
username = ""
|
||||
from_address = ""
|
||||
recipients = []
|
||||
|
||||
config = preferences.set_email_smtp_config({
|
||||
"host": host,
|
||||
"port": req.port,
|
||||
"security": req.security,
|
||||
"username": username,
|
||||
"from_address": from_address,
|
||||
"to_addresses": recipients,
|
||||
})
|
||||
if not host or not username:
|
||||
secrets_store.set_email_smtp_password("")
|
||||
elif req.password is not None:
|
||||
secrets_store.set_email_smtp_password(req.password)
|
||||
return {
|
||||
"email_smtp_config": config,
|
||||
"email_smtp_password_set": bool(secrets_store.get_email_smtp_password()),
|
||||
}
|
||||
|
||||
|
||||
class WebhookTestIn(BaseModel):
|
||||
channel: Literal["feishu", "wecom", "custom", "email"]
|
||||
|
||||
|
||||
@router.post("/preferences/webhook-test")
|
||||
def test_webhook(req: WebhookTestIn) -> dict:
|
||||
"""向已保存的 Webhook 地址发送一条测试消息,验证配置是否正确。
|
||||
|
||||
只测试已保存的配置(与生产推送同源),不测试未保存草稿。
|
||||
未配置 / 地址非法 / 发送失败均返回 HTTP 200 + {ok: False},
|
||||
前端统一读 detail 渲染绿/红,不抛 400。
|
||||
"""
|
||||
from app.services import preferences
|
||||
from app.services import webhook_adapter
|
||||
|
||||
title = "TickFlow Stock Panel 推送测试"
|
||||
body = "如果你看到这条消息,说明推送配置正确 🎉"
|
||||
|
||||
if req.channel == "feishu":
|
||||
url = preferences.get_feishu_webhook_url()
|
||||
if not url:
|
||||
return {"ok": False, "detail": "尚未配置飞书 Webhook,请先保存"}
|
||||
if not webhook_adapter.is_valid_feishu_url(url):
|
||||
return {"ok": False, "detail": "已保存的飞书 Webhook 地址非法,请重新保存"}
|
||||
secret = preferences.get_feishu_webhook_secret()
|
||||
# 诊断用途单次尝试: 失败即返回, 不等生产退避重试 (~17s)
|
||||
ok = webhook_adapter.send_feishu(url, title, body, secret, max_attempts=1)
|
||||
elif req.channel == "wecom":
|
||||
url = preferences.get_wecom_webhook_url()
|
||||
if not url:
|
||||
return {"ok": False, "detail": "尚未配置企业微信 Webhook,请先保存"}
|
||||
if not webhook_adapter.is_valid_wecom_url(url):
|
||||
return {"ok": False, "detail": "已保存的企业微信 Webhook 地址非法,请重新保存"}
|
||||
ok = webhook_adapter.send_wecom(url, title, body)
|
||||
elif req.channel == "custom":
|
||||
url = preferences.get_custom_webhook_url()
|
||||
if not url:
|
||||
return {"ok": False, "detail": "尚未配置第三方 Webhook, 请先保存"}
|
||||
if not webhook_adapter.is_valid_custom_url(url):
|
||||
return {"ok": False, "detail": "已保存的第三方 Webhook 地址非法, 请重新保存"}
|
||||
ok = webhook_adapter.send_custom(
|
||||
url,
|
||||
title,
|
||||
body,
|
||||
event_type="test",
|
||||
secret=secrets_store.get_custom_webhook_secret(),
|
||||
max_attempts=1,
|
||||
)
|
||||
else: # email
|
||||
from app.services import email_adapter
|
||||
|
||||
config = preferences.get_email_smtp_config()
|
||||
if not email_adapter.is_configured(config):
|
||||
return {"ok": False, "detail": "尚未完整配置邮件 SMTP, 请先保存"}
|
||||
ok = email_adapter.send_email(
|
||||
config,
|
||||
secrets_store.get_email_smtp_password(),
|
||||
title,
|
||||
body,
|
||||
max_attempts=1,
|
||||
)
|
||||
|
||||
if ok:
|
||||
return {"ok": True, "detail": "测试消息已发送, 请检查对应接收端"}
|
||||
return {"ok": False, "detail": "推送失败:网络不可达或地址/密钥不正确,详情见后端日志"}
|
||||
|
||||
|
||||
class WecomBotPrefsIn(BaseModel):
|
||||
bot_id: str
|
||||
secret: str
|
||||
@@ -1313,7 +1466,7 @@ def update_webhook_enabled_default(req: WebhookEnabledDefaultIn) -> dict:
|
||||
|
||||
|
||||
class WebhookDefaultChannelsIn(BaseModel):
|
||||
channels: list[str] # 多选: ['feishu','wecom'] 等; 空数组=默认不推送
|
||||
channels: list[str] # 多选: feishu / wecom / custom / email; 空数组=不推送
|
||||
|
||||
|
||||
@router.put("/preferences/webhook-default-channels")
|
||||
@@ -1334,7 +1487,7 @@ def update_quote_interval(req: QuoteIntervalIn, request: Request) -> dict:
|
||||
"""更新行情轮询间隔。按档位自动 clamp。"""
|
||||
qs = getattr(request.app.state, "quote_service", None)
|
||||
if not qs:
|
||||
return {"interval": req.interval, "min_interval": qs.get_min_interval(), "max_interval": 60.0}
|
||||
return {"interval": req.interval, "min_interval": 6.0, "max_interval": 60.0}
|
||||
clamped = qs.set_interval(req.interval)
|
||||
return {
|
||||
"interval": clamped,
|
||||
@@ -1514,7 +1667,9 @@ async def test_endpoint(req: TestEndpointIn) -> dict:
|
||||
import statistics
|
||||
|
||||
base = req.url.rstrip("/")
|
||||
rounds = max(1, min(10, req.rounds or _endpoints_cache.get("data", {}).get("testRounds", 5)))
|
||||
# 缓存初值的 "data" 是 None(键存在, get 的默认值不生效), 端点清单未预热时要兜底
|
||||
manifest = _endpoints_cache.get("data") or {}
|
||||
rounds = max(1, min(10, req.rounds or manifest.get("testRounds", 5)))
|
||||
health_url = base + "/health"
|
||||
|
||||
latencies: list[float] = []
|
||||
@@ -1756,17 +1911,24 @@ def update_review_schedule(req: ReviewScheduleIn, request: Request) -> dict:
|
||||
|
||||
|
||||
class ReviewPushIn(BaseModel):
|
||||
channels: list[str] # 多选: ['feishu'] 等; 空数组=不推送。微信等开发中
|
||||
channels: list[str] # 多选: feishu / wecom / custom / email; 空数组=不推送
|
||||
mode: str | None = None # 可选: auto=归档即推 / manual=仅显式 push; 不传则不变
|
||||
|
||||
|
||||
@router.put("/preferences/review-push")
|
||||
def update_review_push(req: ReviewPushIn) -> dict:
|
||||
"""复盘推送渠道(多选) — 选定把复盘报告(手动生成 / 定时生成归档后)推送到哪些外部工具。
|
||||
"""复盘推送设置(渠道多选 + 触发方式)。
|
||||
|
||||
纯偏好, 与定时复盘 / 实时行情完全独立, 常驻可单独设置。空数组=不推送。
|
||||
实际推送由归档端点(POST /api/market-recap/reports)与定时任务(_run_scheduled_review)
|
||||
在归档后读取本列表逐个推送。白名单外的渠道会被过滤掉。
|
||||
在归档后读取渠道列表, 并按 review_push_mode 决定是否外发:
|
||||
- manual: 定时复盘只归档不推送, 手动保存需显式 push=true
|
||||
- auto: 归档即推(行为与旧逻辑一致)
|
||||
白名单外的渠道会被过滤掉, 白名单外的 mode 值回退 manual。
|
||||
"""
|
||||
from app.services import preferences
|
||||
saved = preferences.set_review_push_channels(req.channels)
|
||||
return {"review_push_channels": saved}
|
||||
mode = preferences.get_review_push_mode()
|
||||
if req.mode is not None:
|
||||
mode = preferences.set_review_push_mode(req.mode)
|
||||
return {"review_push_channels": saved, "review_push_mode": mode}
|
||||
|
||||
+160
-4
@@ -10,6 +10,7 @@ from fastapi import APIRouter, HTTPException, Request
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.strategy import custom_signals
|
||||
from app.strategy.intraday_features import INTRADAY_FEATURES
|
||||
|
||||
router = APIRouter(prefix="/api/custom-signals", tags=["custom-signals"])
|
||||
|
||||
@@ -24,7 +25,9 @@ def _invalidate(request: Request) -> None:
|
||||
信号增删会改变注入列集合: 只清表达式缓存不够, repo 内存缓存 /
|
||||
strategy 磁盘缓存里算好的历史窗口仍不含新 csg_ 列 (或仍含已删列),
|
||||
需要一并清除, 否则创建信号后立即运行策略仍会报缺列。
|
||||
盘中信号定义缓存(intraday)一并失效, 下一分钟 bucket 即生效。
|
||||
"""
|
||||
custom_signals.invalidate_intraday_cache()
|
||||
from app.indicators.pipeline import invalidate_custom_signals
|
||||
invalidate_custom_signals()
|
||||
from app.services import strategy_cache
|
||||
@@ -35,11 +38,11 @@ def _invalidate(request: Request) -> None:
|
||||
|
||||
|
||||
class ConditionModel(BaseModel):
|
||||
left: str # 字段名(须在白名单)
|
||||
op: str # > >= < <= == !=
|
||||
left: str # 字段名(日线在白名单 / 盘中在特征白名单)
|
||||
op: str # > >= < <= == != ; 盘中额外: cross_up cross_down
|
||||
right: str # "field:xxx" 或数字字符串
|
||||
leftDays: int = 0 # 左字段取几日前 (0=当日, 默认)
|
||||
rightDays: int = 0 # 右字段取几日前 (仅 right 为字段时有意义)
|
||||
leftDays: int = 0 # 左字段取几日前 (0=当日, 默认; 盘中信号必须为 0)
|
||||
rightDays: int = 0 # 右字段取几日前 (仅 right 为字段时有意义; 盘中信号必须为 0)
|
||||
|
||||
|
||||
class SignalModel(BaseModel):
|
||||
@@ -48,6 +51,17 @@ class SignalModel(BaseModel):
|
||||
kind: str # entry | exit | both
|
||||
conditions: list[ConditionModel]
|
||||
enabled: bool = True
|
||||
timeframe: str = "daily" # daily | intraday(分钟K特征, 输出当日条件上升沿)
|
||||
min_bars: int = 0 # 仅 intraday: 当日最少已完成 bar 数, 不足不触发
|
||||
|
||||
|
||||
class IntradayReplayRequest(BaseModel):
|
||||
"""盘中信号历史回放 — 用本地分钟K重放触发时点, 不消耗盘中数据能力。"""
|
||||
signal_id: str
|
||||
start_date: str # YYYY-MM-DD
|
||||
end_date: str # YYYY-MM-DD
|
||||
symbols: list[str]
|
||||
asset_type: str = "stock"
|
||||
|
||||
|
||||
class AIGenerateRequest(BaseModel):
|
||||
@@ -87,16 +101,59 @@ def get_options():
|
||||
groups.append({"key": cat, "label": label,
|
||||
"fields": [{"key": f, "label": ENRICHED_COLUMNS.get(f, f)} for f in cat_fields]})
|
||||
|
||||
# 注册表因子 (虚拟/自定义/复合): 历史路径由 compute_signals 复用评分物化
|
||||
# 管线补算; 已是物化列的基础因子 (rsi_14 等) 上面已分组, 此处跳过。
|
||||
from app.factors.registry import all_factors
|
||||
|
||||
factor_groups: dict[str, list[dict[str, str]]] = {}
|
||||
for spec in all_factors():
|
||||
if spec.id in allowed:
|
||||
continue
|
||||
label = spec.label
|
||||
if spec.warmup_bars > 1:
|
||||
label = f"{label} · 预热{spec.warmup_bars}日"
|
||||
if list(spec.asset_types) == ["stock"]:
|
||||
label = f"{label} · 仅股票"
|
||||
factor_groups.setdefault(spec.group or "因子", []).append({"key": spec.id, "label": label})
|
||||
for group_label, group_fields in factor_groups.items():
|
||||
groups.append({"key": f"factor:{group_label}", "label": f"因子 · {group_label}", "fields": group_fields})
|
||||
fields.extend(group_fields)
|
||||
|
||||
# string 扩展字段 (概念/行业归属等): 只进信号条件, 不注册为因子。
|
||||
# stringFields 标记 + 独立分组, 前端据此切换运算符 (包含/等于/不等于)
|
||||
# 与右值输入 (字符串文本, 不支持字段引用)。
|
||||
from app.factors.ext_factors import ext_string_field_entries
|
||||
|
||||
str_entries = ext_string_field_entries()
|
||||
if str_entries:
|
||||
str_group = {"key": "ext_string", "label": "扩展 · 字符串", "fields": str_entries}
|
||||
groups.append(str_group)
|
||||
fields.extend(str_entries)
|
||||
|
||||
return {
|
||||
"fields": fields,
|
||||
"groups": groups,
|
||||
"maxDays": custom_signals.MAX_DAYS,
|
||||
"operators": [">", ">=", "<", "<=", "==", "!="],
|
||||
"stringFields": [e["key"] for e in str_entries],
|
||||
"stringOperators": ["contains", "==", "!="],
|
||||
"kinds": [
|
||||
{"key": "entry", "label": "入场"},
|
||||
{"key": "exit", "label": "出场"},
|
||||
{"key": "both", "label": "出入通用"},
|
||||
],
|
||||
# 盘中信号(timeframe=intraday): 分钟K特征白名单 + 额外穿越算子
|
||||
"intraday": {
|
||||
"fields": [
|
||||
{"key": f, "label": label}
|
||||
for f, label in sorted(INTRADAY_FEATURES.items())
|
||||
],
|
||||
"operators": [">", ">=", "<", "<=", "==", "!=", "cross_up", "cross_down"],
|
||||
},
|
||||
"timeframes": [
|
||||
{"key": "daily", "label": "日线"},
|
||||
{"key": "intraday", "label": "盘中(分钟K)"},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@@ -171,3 +228,102 @@ def delete_signal(signal_id: str, request: Request):
|
||||
raise HTTPException(status_code=404, detail="信号不存在")
|
||||
_invalidate(request)
|
||||
return {"ok": True}
|
||||
|
||||
|
||||
# ── 盘中信号历史回放 ────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/intraday/replay")
|
||||
def intraday_replay(req: IntradayReplayRequest, request: Request):
|
||||
"""用本地历史分钟K回放盘中信号的触发时点。
|
||||
|
||||
只读本地分钟分区, 不消耗盘中数据能力 — 用户可先在历史区间验证信号,
|
||||
再决定是否配置到监控/分钟策略。昨收取自本地日K(无昨日数据的日子该特征降级)。
|
||||
"""
|
||||
from datetime import date, timedelta
|
||||
|
||||
import polars as pl
|
||||
|
||||
from app.strategy.intraday_features import build_feature_frame
|
||||
|
||||
try:
|
||||
start = date.fromisoformat(req.start_date)
|
||||
end = date.fromisoformat(req.end_date)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=f"日期格式错误: {e}") from e
|
||||
if start > end:
|
||||
raise HTTPException(status_code=400, detail="start_date 不能晚于 end_date")
|
||||
if (end - start).days > 60:
|
||||
raise HTTPException(status_code=400, detail="回放区间最长 60 天")
|
||||
symbols = [s for s in dict.fromkeys(req.symbols) if s]
|
||||
if not symbols:
|
||||
raise HTTPException(status_code=400, detail="symbols 不能为空")
|
||||
if len(symbols) > 200:
|
||||
raise HTTPException(status_code=400, detail="单次回放最多 200 只标的")
|
||||
|
||||
# 信号定义必须存在且为盘中类型
|
||||
sig = next(
|
||||
(s for s in custom_signals.load_all(_data_dir(request)) if s.get("id") == req.signal_id),
|
||||
None,
|
||||
)
|
||||
if sig is None:
|
||||
raise HTTPException(status_code=404, detail="信号不存在")
|
||||
if sig.get("timeframe") != custom_signals.TIMEFRAME_INTRADAY:
|
||||
raise HTTPException(status_code=400, detail="该信号不是盘中(timeframe=intraday)信号")
|
||||
exprs = custom_signals.build_intraday_expressions([sig])
|
||||
col = custom_signals.intraday_column_name(sig["id"])
|
||||
if col not in exprs:
|
||||
raise HTTPException(status_code=400, detail="信号编译失败, 请检查条件字段")
|
||||
min_bars = int(sig.get("min_bars", 0) or 0)
|
||||
|
||||
repo = request.app.state.repo
|
||||
# 昨收映射: 一次性取区间(含前置 15 天)日K, 按「严格早于当日」取最近收盘
|
||||
daily = repo.get_daily_batch(symbols, start - timedelta(days=15), end, columns=["symbol", "date", "close"])
|
||||
close_by_sym_date: dict[str, dict[date, float]] = {}
|
||||
if not daily.is_empty():
|
||||
for row in daily.sort(["symbol", "date"]).iter_rows(named=True):
|
||||
close_by_sym_date.setdefault(str(row["symbol"]), {})[row["date"]] = float(row["close"])
|
||||
|
||||
triggers: list[dict] = []
|
||||
days_scanned = 0
|
||||
bars_scanned = 0
|
||||
day = start
|
||||
while day <= end:
|
||||
minute_df = repo.get_minute_batch(symbols, day, asset_type=req.asset_type)
|
||||
if minute_df is not None and not minute_df.is_empty():
|
||||
days_scanned += 1
|
||||
bars_scanned += minute_df.height
|
||||
prev_close = {
|
||||
sym: closes_map[max(d for d in closes_map if d < day)]
|
||||
for sym, closes_map in close_by_sym_date.items()
|
||||
if any(d < day for d in closes_map)
|
||||
}
|
||||
frame = build_feature_frame(minute_df, prev_close=prev_close)
|
||||
if not frame.is_empty():
|
||||
evaluated = custom_signals.apply_intraday_edges(frame, {col: exprs[col]}).with_columns(
|
||||
pl.int_range(pl.len()).over(["symbol", "date"]).alias("_bar_idx")
|
||||
)
|
||||
if min_bars > 0:
|
||||
evaluated = evaluated.with_columns(
|
||||
pl.when(pl.col("_bar_idx") + 1 >= min_bars)
|
||||
.then(pl.col(col))
|
||||
.otherwise(False)
|
||||
.alias(col)
|
||||
)
|
||||
for row in evaluated.filter(pl.col(col)).sort(["datetime", "symbol"]).iter_rows(named=True):
|
||||
triggers.append({
|
||||
"date": day.isoformat(),
|
||||
"time": str(row["datetime"].time()),
|
||||
"symbol": row["symbol"],
|
||||
})
|
||||
day += timedelta(days=1)
|
||||
|
||||
return {
|
||||
"signal_id": req.signal_id,
|
||||
"start_date": req.start_date,
|
||||
"end_date": req.end_date,
|
||||
"symbols": symbols,
|
||||
"days_scanned": days_scanned,
|
||||
"bars_scanned": bars_scanned,
|
||||
"triggers": triggers,
|
||||
}
|
||||
|
||||
@@ -188,6 +188,7 @@ def _strategy_detail(
|
||||
"description": description or s.meta.get("description", ""),
|
||||
"tags": s.meta.get("tags", []),
|
||||
"source": s.source,
|
||||
"research_only": s.meta.get("research_only", False),
|
||||
"execution_backend": s.execution_backend,
|
||||
"asset_types": s.meta.get("asset_types", ["stock"]),
|
||||
"timeframes": s.meta.get("timeframes", ["1d"]),
|
||||
@@ -307,14 +308,17 @@ def list_strategies(
|
||||
request: Request,
|
||||
asset_type: str | None = None,
|
||||
timeframe: str | None = None,
|
||||
include_research: bool = False,
|
||||
):
|
||||
engine = _get_engine(request)
|
||||
data_dir = _data_dir(request)
|
||||
all_overrides = strategy_config.list_overrides(data_dir)
|
||||
|
||||
result = []
|
||||
for meta in engine.list_strategies():
|
||||
if meta.get("research_only"):
|
||||
# include_research=True 时返回 research_only 草稿(供前端「草稿」分区展示/发布)。
|
||||
# 默认 False 保持既有行为: 草稿不进公开列表。
|
||||
for meta in engine.list_strategies(include_research=include_research):
|
||||
if meta.get("research_only") and not include_research:
|
||||
continue
|
||||
if asset_type and asset_type not in meta.get("asset_types", ["stock"]):
|
||||
continue
|
||||
@@ -560,7 +564,11 @@ def _set_meta_string_field(block: str, field: str, value: str) -> str:
|
||||
)
|
||||
if count:
|
||||
return next_block
|
||||
return _insert_meta_field(block, field, _py_string(value))
|
||||
|
||||
|
||||
def _insert_meta_field(block: str, field: str, value_repr: str) -> str:
|
||||
"""在 META 字典末尾(闭合 `}` 之前)插入一个字段。value_repr 已是 Python 源码。"""
|
||||
lines = block.splitlines(keepends=True)
|
||||
key_indent = None
|
||||
for line in lines:
|
||||
@@ -585,7 +593,33 @@ def _set_meta_string_field(block: str, field: str, value: str) -> str:
|
||||
newline = lines[i][len(body):]
|
||||
lines[i] = body.rstrip() + "," + newline
|
||||
break
|
||||
lines.insert(insert_at, f'{key_indent}"{field}": {_py_string(value)},\n')
|
||||
lines.insert(insert_at, f'{key_indent}"{field}": {value_repr},\n')
|
||||
return "".join(lines)
|
||||
|
||||
|
||||
def _set_meta_bool_field(code: str, field: str, value: bool) -> str:
|
||||
"""设置 META 里的布尔字段(纯文本改写, 不执行代码): 存在则替换, 不存在则追加。"""
|
||||
found = find_meta_assignment(code)
|
||||
if found is None:
|
||||
raise ValueError("找不到 META 字典")
|
||||
meta_node = found[1]
|
||||
lines = code.splitlines(keepends=True)
|
||||
start = meta_node.lineno - 1
|
||||
end = meta_node.end_lineno or meta_node.lineno
|
||||
block = "".join(lines[start:end])
|
||||
|
||||
value_repr = "True" if value else "False"
|
||||
key_pattern = re.compile(
|
||||
rf"(?m)^(\s*[\"']{re.escape(field)}[\"']\s*:\s*)(?:True|False|[\"'][^\"'\n]*[\"'])"
|
||||
)
|
||||
next_block, count = key_pattern.subn(
|
||||
lambda m: f"{m.group(1)}{value_repr}",
|
||||
block,
|
||||
count=1,
|
||||
)
|
||||
if not count:
|
||||
next_block = _insert_meta_field(block, field, value_repr)
|
||||
lines[start:end] = next_block.splitlines(keepends=True)
|
||||
return "".join(lines)
|
||||
|
||||
|
||||
@@ -732,6 +766,13 @@ def _save_strategy_code(req: StrategyCodeSaveRequest, request: Request, *, legac
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
prepared = _prepare_strategy_code(req)
|
||||
|
||||
# AI 新建策略默认草稿态(research_only=True): 不进公开列表、不可运行, 需显式 publish。
|
||||
# 仅 create 注入; update 保留既有 research_only, 避免静默取消已发布状态。
|
||||
if expected_source == "ai" and (legacy_ai_path or req.mode == "create"):
|
||||
prepared["code"] = _set_meta_bool_field(prepared["code"], "research_only", True)
|
||||
prepared["meta"] = AIStrategyGenerator._extract_meta(prepared["code"])
|
||||
|
||||
previous_code = path.read_text(encoding="utf-8") if path.exists() else None
|
||||
path.write_text(prepared["code"], encoding="utf-8")
|
||||
|
||||
@@ -763,6 +804,7 @@ def _save_strategy_code(req: StrategyCodeSaveRequest, request: Request, *, legac
|
||||
"source": expected_source,
|
||||
"path": str(path),
|
||||
"meta": prepared["meta"],
|
||||
"research_only": prepared["meta"].get("research_only", False),
|
||||
}
|
||||
|
||||
|
||||
@@ -1062,6 +1104,45 @@ async def ai_save(req: AISaveRequest, request: Request):
|
||||
raise HTTPException(status_code=400, detail=str(e)) from e
|
||||
|
||||
|
||||
@router.post("/{strategy_id}/publish")
|
||||
def publish_ai_strategy(strategy_id: str, request: Request):
|
||||
"""把 research_only 的 AI 草稿策略翻转为公开(research_only=False)。
|
||||
|
||||
门 = 人的显式动作: 只有 AI 来源且仍处于草稿态的策略才能被发布。
|
||||
发布后即进入公开列表、可 run、可监控。
|
||||
"""
|
||||
sid = _validate_strategy_id(strategy_id)
|
||||
engine = _get_engine(request)
|
||||
try:
|
||||
s = engine.get(sid)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=f"策略 {sid} 不存在") from e
|
||||
|
||||
if s.source != "ai":
|
||||
raise HTTPException(status_code=400, detail="仅 AI 策略可经发布端点上线")
|
||||
if not s.meta.get("research_only"):
|
||||
raise HTTPException(status_code=400, detail="该策略已是公开状态")
|
||||
|
||||
path = s.file_path
|
||||
if path is None:
|
||||
raise HTTPException(status_code=400, detail="策略源文件路径无效, 无法发布")
|
||||
previous_code = path.read_text(encoding="utf-8")
|
||||
path.write_text(_set_meta_bool_field(previous_code, "research_only", False), encoding="utf-8")
|
||||
|
||||
try:
|
||||
engine.reload()
|
||||
loaded = engine.get(sid)
|
||||
if loaded.meta.get("research_only"):
|
||||
raise ValueError("发布后策略仍为草稿态")
|
||||
except Exception as e:
|
||||
_restore_strategy_file(path, previous_code)
|
||||
engine.reload()
|
||||
raise HTTPException(status_code=500, detail=f"策略发布失败: {e}") from e
|
||||
|
||||
_invalidate_strategy_runtime(request)
|
||||
return {"ok": True, "strategy_id": sid}
|
||||
|
||||
|
||||
@router.delete("/{strategy_id}")
|
||||
def delete_strategy(strategy_id: str, request: Request):
|
||||
"""删除自定义策略 — 清除源文件、运行时注册和关联状态。内置策略不可删除。"""
|
||||
|
||||
@@ -5,6 +5,7 @@ import logging
|
||||
import math
|
||||
import time
|
||||
from datetime import date
|
||||
from typing import Callable
|
||||
|
||||
import anyio
|
||||
import polars as pl
|
||||
@@ -13,6 +14,7 @@ from pydantic import BaseModel
|
||||
|
||||
from app.db_safe import is_valid_ext_ident, quote_ident
|
||||
from app.services import watchlist
|
||||
from app.services.watchlist_csv import import_watchlist_codes, import_watchlist_csv
|
||||
from app.services.watchlist_ocr import import_watchlist_image
|
||||
from app.services.watchlist_ocr.provider import get_ocr_provider
|
||||
|
||||
@@ -31,6 +33,35 @@ _IMPORT_IMAGE_TYPES = {
|
||||
}
|
||||
# OCR 独立并发上限:避免多张大图同时解码 + 多 Tesseract 子进程
|
||||
_OCR_LIMITER = anyio.CapacityLimiter(2)
|
||||
# CSV/TXT 导入:文本远小于截图,上限 5MB 足够
|
||||
_MAX_IMPORT_CSV_BYTES = 5 * 1024 * 1024
|
||||
_IMPORT_CSV_TYPES = {
|
||||
"text/csv",
|
||||
"text/plain",
|
||||
"application/csv",
|
||||
}
|
||||
# 上传分块读取粒度 (与 ext_data 上传一致)
|
||||
_UPLOAD_CHUNK_BYTES = 1024 * 1024
|
||||
|
||||
|
||||
async def _read_upload_capped(file: UploadFile, max_bytes: int, too_large: str) -> bytes:
|
||||
"""分块读取上传内容, 累计超过 max_bytes 立即拒绝(400), 返回完整字节。
|
||||
|
||||
与 ext_data._write_upload_capped 同类保护: 一次性 `await file.read()` 会先把整个
|
||||
文件读入内存再比较长度, 上限在那之后才生效, 一个远超上限的上传照样把进程内存
|
||||
顶满; 分块读取在越过上限的那一块就停止, 内存占用不超过上限 + 一块。
|
||||
"""
|
||||
chunks: list[bytes] = []
|
||||
total = 0
|
||||
while True:
|
||||
chunk = await file.read(_UPLOAD_CHUNK_BYTES)
|
||||
if not chunk:
|
||||
break
|
||||
total += len(chunk)
|
||||
if total > max_bytes:
|
||||
raise HTTPException(400, too_large)
|
||||
chunks.append(chunk)
|
||||
return b"".join(chunks)
|
||||
|
||||
|
||||
class AddRequest(BaseModel):
|
||||
@@ -43,6 +74,7 @@ class BatchAddRequest(BaseModel):
|
||||
symbols: list[str]
|
||||
note: str = ""
|
||||
group_id: str | None = None
|
||||
group_ids: list[str] | None = None
|
||||
|
||||
|
||||
class GroupNameRequest(BaseModel):
|
||||
@@ -58,6 +90,10 @@ class GroupAssignRequest(BaseModel):
|
||||
group_id: str | None = None
|
||||
|
||||
|
||||
class ImportCodesRequest(BaseModel):
|
||||
text: str
|
||||
|
||||
|
||||
def _with_names(rows: list[dict], request: Request) -> list[dict]:
|
||||
if not rows:
|
||||
return rows
|
||||
@@ -89,7 +125,12 @@ def add_one(req: AddRequest, request: Request):
|
||||
@router.post("/batch")
|
||||
def add_batch(req: BatchAddRequest, request: Request):
|
||||
try:
|
||||
rows, added = watchlist.add_batch(req.symbols, req.note, req.group_id)
|
||||
rows, added = watchlist.add_batch(
|
||||
req.symbols,
|
||||
req.note,
|
||||
group_id=req.group_id,
|
||||
group_ids=req.group_ids,
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(400, str(e)) from e
|
||||
return {"symbols": _with_names(rows, request), "added": added}
|
||||
@@ -167,11 +208,9 @@ async def import_from_image(request: Request, file: UploadFile = File(...)):
|
||||
if not ok_type and not ok_ext:
|
||||
raise HTTPException(400, "仅支持 JPG / PNG / WebP / BMP / GIF 图片")
|
||||
|
||||
data = await file.read()
|
||||
data = await _read_upload_capped(file, _MAX_IMPORT_IMAGE_BYTES, "图片过大(上限 12MB)")
|
||||
if not data:
|
||||
raise HTTPException(400, "空文件")
|
||||
if len(data) > _MAX_IMPORT_IMAGE_BYTES:
|
||||
raise HTTPException(400, "图片过大(上限 12MB)")
|
||||
|
||||
existing = {r["symbol"] for r in watchlist.list_symbols()}
|
||||
data_dir = request.app.state.repo.store.data_dir
|
||||
@@ -194,6 +233,68 @@ async def import_from_image(request: Request, file: UploadFile = File(...)):
|
||||
return result
|
||||
|
||||
|
||||
def _run_candidate_import(parse: Callable[[], dict], empty_msg: str) -> dict:
|
||||
"""执行候选解析:ValueError→400、其他→500、空候选→400、剥离 raw_text。"""
|
||||
try:
|
||||
result = parse()
|
||||
except ValueError as e:
|
||||
raise HTTPException(400, str(e)) from e
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.exception("watchlist import failed")
|
||||
raise HTTPException(500, f"解析失败: {e}") from e
|
||||
if not result["candidates"]:
|
||||
raise HTTPException(400, empty_msg)
|
||||
result.pop("raw_text", None)
|
||||
return result
|
||||
|
||||
|
||||
@router.post("/import-csv")
|
||||
async def import_from_csv(request: Request, file: UploadFile = File(...)):
|
||||
"""从 CSV / TXT 导入自选候选列表(不自动写入自选)。
|
||||
|
||||
兼容同花顺/东财/通达信导出(逗号或 Tab 分隔、UTF-8 或 GBK 编码)。目标分组
|
||||
在候选确认时由前端传入 batch 接口,本端点只做解析与主数据校验。
|
||||
"""
|
||||
content_type = (file.content_type or "").split(";")[0].strip().lower()
|
||||
filename = (file.filename or "").lower()
|
||||
ok_type = content_type in _IMPORT_CSV_TYPES
|
||||
ok_ext = filename.endswith((".csv", ".txt"))
|
||||
if not ok_type and not ok_ext:
|
||||
raise HTTPException(400, "仅支持 CSV / TXT 文件")
|
||||
|
||||
data = await _read_upload_capped(file, _MAX_IMPORT_CSV_BYTES, "文件过大(上限 5MB)")
|
||||
if not data:
|
||||
raise HTTPException(400, "空文件")
|
||||
|
||||
data_dir = request.app.state.repo.store.data_dir
|
||||
# 解码与自选/instruments parquet 读取为同步 CPU/IO,挪线程池避免卡事件循环
|
||||
return await anyio.to_thread.run_sync(
|
||||
lambda: _run_candidate_import(
|
||||
lambda: import_watchlist_csv(
|
||||
data,
|
||||
data_dir,
|
||||
existing_symbols={r["symbol"] for r in watchlist.list_symbols()},
|
||||
),
|
||||
"文件中未识别到股票代码或名称",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@router.post("/import-codes")
|
||||
def import_from_codes(req: ImportCodesRequest, request: Request):
|
||||
"""从粘贴的证券代码导入自选候选列表(不自动写入自选)。"""
|
||||
text = req.text.strip()
|
||||
if not text:
|
||||
raise HTTPException(400, "请输入要导入的股票代码")
|
||||
|
||||
existing = {r["symbol"] for r in watchlist.list_symbols()}
|
||||
data_dir = request.app.state.repo.store.data_dir
|
||||
return _run_candidate_import(
|
||||
lambda: import_watchlist_codes(text, data_dir, existing_symbols=existing),
|
||||
"未识别到股票代码",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/{symbol}/top")
|
||||
def move_one_to_top(symbol: str, request: Request):
|
||||
rows = watchlist.move_to_top(symbol)
|
||||
|
||||
@@ -300,6 +300,14 @@ class PanelCache:
|
||||
return f"{asset_type}:{generation or 'unmanaged'}:{h}:{start}:{end}:{cols}"
|
||||
|
||||
|
||||
# 等待进行中 enriched 发布的上限与轮询间隔。孤儿标记由 get_enriched_generation
|
||||
# 在读取时直接自愈, 因此这里等到的 EnrichedGenerationUnavailableError 意味着
|
||||
# 发布方确实存活 —— 对回测/优化这类长任务, 有界等待优于立即失败。仅用于
|
||||
# worker 任务路径 (矩阵加载), 实时热路径不得调用 data_generation_await。
|
||||
_GENERATION_WAIT_TIMEOUT_S = 300.0
|
||||
_GENERATION_POLL_S = 1.0
|
||||
|
||||
|
||||
# ================================================================
|
||||
# BacktestEngine
|
||||
# ================================================================
|
||||
@@ -317,6 +325,25 @@ class BacktestEngine:
|
||||
loader = getattr(self.repo, "get_matrix_data_generation", None)
|
||||
return loader(asset_type) if callable(loader) else None
|
||||
|
||||
def data_generation_await(
|
||||
self,
|
||||
asset_type: str = "stock",
|
||||
*,
|
||||
cancel_event: threading.Event | None = None,
|
||||
timeout_s: float = _GENERATION_WAIT_TIMEOUT_S,
|
||||
) -> str | None:
|
||||
"""获取 generation; 发布进行中时在超时窗口内轮询, 可被取消事件打断。"""
|
||||
deadline = time.monotonic() + timeout_s
|
||||
while True:
|
||||
try:
|
||||
return self.data_generation(asset_type)
|
||||
except EnrichedGenerationUnavailableError:
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
raise
|
||||
if time.monotonic() >= deadline:
|
||||
raise
|
||||
time.sleep(_GENERATION_POLL_S)
|
||||
|
||||
def assert_data_generation(
|
||||
self,
|
||||
asset_type: str,
|
||||
@@ -506,15 +533,10 @@ class BacktestEngine:
|
||||
if cache_profile is not None
|
||||
else settings.backtest_matrix_cache_max_mb * 1024 * 1024
|
||||
)
|
||||
generation_loader = getattr(self.repo, "get_matrix_data_generation", None)
|
||||
source_generation = (
|
||||
expected_generation
|
||||
if expected_generation is not None
|
||||
else (
|
||||
generation_loader(asset_type)
|
||||
if callable(generation_loader)
|
||||
else None
|
||||
)
|
||||
else self.data_generation_await(asset_type, cancel_event=cancel_event)
|
||||
)
|
||||
attempts = 1 if expected_generation is not None else 2
|
||||
for attempt in range(attempts):
|
||||
@@ -555,7 +577,9 @@ class BacktestEngine:
|
||||
except EnrichedGenerationUnavailableError:
|
||||
if attempt + 1 >= attempts:
|
||||
raise
|
||||
source_generation = self.data_generation(asset_type)
|
||||
source_generation = self.data_generation_await(
|
||||
asset_type, cancel_event=cancel_event
|
||||
)
|
||||
except pa.ArrowException as exc:
|
||||
raise ValueError(f"direct market matrix parquet scan failed: {exc}") from exc
|
||||
raise EnrichedGenerationUnavailableError(
|
||||
|
||||
@@ -17,12 +17,14 @@ from typing import Any, Literal
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
|
||||
from app.backtest import stats_v2
|
||||
from app.backtest.engine import BacktestEngine
|
||||
from app.backtest.fundamentals import (
|
||||
FUNDAMENTAL_FACTOR_NAMES,
|
||||
attach_fundamental_factors,
|
||||
load_fundamental_snapshot,
|
||||
)
|
||||
from app.factors.registry import factor_columns_view as _factor_columns_view
|
||||
from app.strategy.scoring import (
|
||||
VIRTUAL_SCORING_DEPENDENCIES as DERIVED_FACTOR_DEPENDENCIES,
|
||||
)
|
||||
@@ -33,80 +35,8 @@ from app.strategy.scoring import (
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 可研究因子目录。保留历史 ID 兼容已有候选方案; 价格尺度相关指标优先提供归一化版本。
|
||||
FACTOR_COLUMNS: list[dict] = [
|
||||
{"id": "momentum_5d", "label": "5日动量", "group": "动量", "desc": "5个交易日累计收益率"},
|
||||
{"id": "momentum_10d", "label": "10日动量", "group": "动量", "desc": "10个交易日累计收益率"},
|
||||
{"id": "momentum_20d", "label": "20日动量", "group": "动量", "desc": "20个交易日累计收益率"},
|
||||
{"id": "momentum_30d", "label": "30日动量", "group": "动量", "desc": "30个交易日累计收益率"},
|
||||
{"id": "momentum_60d", "label": "60日动量", "group": "动量", "desc": "60个交易日累计收益率"},
|
||||
{"id": "change_pct", "label": "日涨跌幅", "group": "动量", "desc": "当日收盘相对前收盘的收益率"},
|
||||
|
||||
{"id": "ma5_bias", "label": "MA5乖离", "group": "均线偏离", "desc": "收盘价 / MA5 - 1"},
|
||||
{"id": "ma10_bias", "label": "MA10乖离", "group": "均线偏离", "desc": "收盘价 / MA10 - 1"},
|
||||
{"id": "ma20_bias", "label": "MA20乖离", "group": "均线偏离", "desc": "收盘价 / MA20 - 1"},
|
||||
{"id": "ma30_bias", "label": "MA30乖离", "group": "均线偏离", "desc": "收盘价 / MA30 - 1"},
|
||||
{"id": "ma60_bias", "label": "MA60乖离", "group": "均线偏离", "desc": "收盘价 / MA60 - 1"},
|
||||
{"id": "ema5_bias", "label": "EMA5乖离", "group": "均线偏离", "desc": "收盘价 / EMA5 - 1"},
|
||||
{"id": "ema10_bias", "label": "EMA10乖离", "group": "均线偏离", "desc": "收盘价 / EMA10 - 1"},
|
||||
{"id": "ema20_bias", "label": "EMA20乖离", "group": "均线偏离", "desc": "收盘价 / EMA20 - 1"},
|
||||
{"id": "ema30_bias", "label": "EMA30乖离", "group": "均线偏离", "desc": "收盘价 / EMA30 - 1"},
|
||||
{"id": "ema60_bias", "label": "EMA60乖离", "group": "均线偏离", "desc": "收盘价 / EMA60 - 1"},
|
||||
|
||||
{"id": "rsi_6", "label": "RSI(6)", "group": "超买超卖", "desc": "6日相对强弱指标"},
|
||||
{"id": "rsi_14", "label": "RSI(14)", "group": "超买超卖", "desc": "14日相对强弱指标"},
|
||||
{"id": "rsi_24", "label": "RSI(24)", "group": "超买超卖", "desc": "24日相对强弱指标"},
|
||||
|
||||
{"id": "macd_hist", "label": "MACD柱(原值)", "group": "趋势", "desc": "兼容历史研究; 跨股票比较建议优先使用MACD柱强度"},
|
||||
{"id": "macd_dif_pct", "label": "MACD DIF强度", "group": "趋势", "desc": "MACD DIF / 收盘价"},
|
||||
{"id": "macd_dea_pct", "label": "MACD DEA强度", "group": "趋势", "desc": "MACD DEA / 收盘价"},
|
||||
{"id": "macd_hist_pct", "label": "MACD柱强度", "group": "趋势", "desc": "MACD柱 / 收盘价, 消除股价尺度影响"},
|
||||
{"id": "kdj_k", "label": "KDJ-K", "group": "趋势", "desc": "KDJ指标K值"},
|
||||
{"id": "kdj_d", "label": "KDJ-D", "group": "趋势", "desc": "KDJ指标D值"},
|
||||
{"id": "kdj_j", "label": "KDJ-J", "group": "趋势", "desc": "KDJ指标J值"},
|
||||
{"id": "boll_position", "label": "布林位置", "group": "趋势", "desc": "收盘价在布林带下轨到上轨之间的位置"},
|
||||
|
||||
{"id": "annual_vol_20d", "label": "20日波动率", "group": "波动率", "desc": "20日收益率年化标准差"},
|
||||
{"id": "atr_14", "label": "ATR(14)原值", "group": "波动率", "desc": "兼容历史研究; 跨股票比较建议优先使用ATR相对波动"},
|
||||
{"id": "atr_pct", "label": "ATR相对波动", "group": "波动率", "desc": "ATR(14) / 收盘价"},
|
||||
{"id": "amplitude", "label": "日振幅", "group": "波动率", "desc": "当日高低价差 / 前收盘价"},
|
||||
{"id": "boll_width", "label": "布林带宽", "group": "波动率", "desc": "布林带上下轨宽度 / MA20"},
|
||||
|
||||
{"id": "vol_ratio_5d", "label": "5日量比", "group": "量价", "desc": "当日成交量 / 前5日平均成交量"},
|
||||
{"id": "vol_ratio_10d", "label": "10日量比", "group": "量价", "desc": "当日成交量 / 前10日平均成交量"},
|
||||
{"id": "vol_trend_5_10", "label": "成交量趋势", "group": "量价", "desc": "5日平均成交量 / 10日平均成交量 - 1"},
|
||||
{"id": "turnover_rate", "label": "换手率", "group": "量价", "desc": "使用历史时点流通股本计算的当日换手率"},
|
||||
{"id": "turnover_ratio_5d", "label": "换手率放大", "group": "量价", "desc": "当日换手率 / 前5日平均换手率 - 1"},
|
||||
{"id": "log_amount", "label": "成交额对数", "group": "量价", "desc": "ln(成交额 + 1), 降低极端规模影响"},
|
||||
{"id": "amount_ratio_5d", "label": "成交额放大", "group": "量价", "desc": "当日成交额 / 前5日平均成交额 - 1"},
|
||||
|
||||
{"id": "gap_return", "label": "开盘跳空", "group": "价格位置", "desc": "开盘价 / 前收盘价 - 1"},
|
||||
{"id": "intraday_return", "label": "日内收益", "group": "价格位置", "desc": "收盘价 / 开盘价 - 1"},
|
||||
{"id": "close_position", "label": "收盘位置", "group": "价格位置", "desc": "收盘价在当日最低价到最高价之间的位置"},
|
||||
{"id": "distance_to_high_60d", "label": "距60日高点", "group": "价格位置", "desc": "收盘价 / 60日最高收盘价 - 1"},
|
||||
{"id": "distance_from_low_60d", "label": "距60日低点", "group": "价格位置", "desc": "收盘价 / 60日最低收盘价 - 1"},
|
||||
{"id": "vwap_bias", "label": "VWAP乖离", "group": "价格位置", "desc": "收盘价 / 当日成交均价 - 1, 成交均价 = 成交额 / (成交量x100)"},
|
||||
|
||||
{"id": "max_ret_20d", "label": "20日最大单日涨幅", "group": "收益形态", "desc": "近20个交易日单日涨幅最大值(彩票效应, 高值代表博彩型特征强)"},
|
||||
{"id": "ret_skew_20d", "label": "20日收益偏度", "group": "收益形态", "desc": "近20个交易日日收益偏度, 高值代表右偏(偶发大涨)"},
|
||||
{"id": "up_days_20d", "label": "20日上涨天数", "group": "收益形态", "desc": "近20个交易日中上涨天数(0~20)"},
|
||||
|
||||
{"id": "amihud_20d", "label": "20日Amihud非流动性", "group": "流动性", "desc": "近20日平均 |日涨跌幅| / 成交额(亿元), 高值代表流动性差"},
|
||||
{"id": "turnover_z_60d", "label": "换手率60日z分", "group": "流动性", "desc": "(当日换手率 - 前60日均值) / 前60日标准差, 衡量换手异动"},
|
||||
|
||||
{"id": "vol_price_corr_20d", "label": "20日量价相关", "group": "量价", "desc": "近20个交易日日涨跌幅与成交量的相关系数, 高值代表量价同向"},
|
||||
{"id": "vol_trend_5_60", "label": "量能趋势(5/60)", "group": "量价", "desc": "5日平均成交量 / 60日平均成交量 - 1"},
|
||||
|
||||
{"id": "limit_up_count_20d", "label": "涨停基因(20日)", "group": "涨停基因", "desc": "近20个交易日涨停次数"},
|
||||
{"id": "limit_up_count_60d", "label": "涨停基因(60日)", "group": "涨停基因", "desc": "近60个交易日涨停次数"},
|
||||
|
||||
{"id": "pb_latest", "label": "市净率(最新公告)", "group": "财务", "desc": "收盘价 / 最新已公告每股净资产; 无财务数据或公告前为空"},
|
||||
{"id": "roe_latest", "label": "ROE(最新公告)", "group": "财务", "desc": "最新已公告净资产收益率(%); 无财务数据或公告前为空"},
|
||||
{"id": "gross_margin_latest", "label": "毛利率(最新公告)", "group": "财务", "desc": "最新已公告销售毛利率(%)"},
|
||||
{"id": "net_margin_latest", "label": "净利率(最新公告)", "group": "财务", "desc": "最新已公告销售净利率(%)"},
|
||||
{"id": "revenue_yoy_latest", "label": "营收增速(最新公告)", "group": "财务", "desc": "最新已公告营业收入同比(%)"},
|
||||
{"id": "net_income_yoy_latest", "label": "净利增速(最新公告)", "group": "财务", "desc": "最新已公告归母净利润同比(%)"},
|
||||
{"id": "debt_ratio_latest", "label": "资产负债率(最新公告)", "group": "财务", "desc": "最新已公告资产负债率(%)"},
|
||||
]
|
||||
# P1 起目录元数据单一权威来源为 app/factors/registry.py, 本常量为兼容别名 (顺序与键不变)。
|
||||
FACTOR_COLUMNS: list[dict] = _factor_columns_view()
|
||||
|
||||
FACTOR_WARMUP_DAYS = 120
|
||||
FACTOR_METHODOLOGY_VERSION = "factor_v2"
|
||||
@@ -208,6 +138,12 @@ class FactorBatchItem:
|
||||
yearly_ic: list[dict] = field(default_factory=list)
|
||||
ic_decay: list[dict] = field(default_factory=list)
|
||||
regime_stats: list[dict] = field(default_factory=list)
|
||||
# metrics_v2 (P3): NW HAC t 值 (滞后=1, 日频 1 日前瞻) 与 BH-FDR q 值; 样本不足为 None
|
||||
t_naive: float | None = None
|
||||
t_newey_west: float | None = None
|
||||
nw_lag: int | None = None
|
||||
p_value: float | None = None
|
||||
q_value: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -230,6 +166,17 @@ class FactorBacktestService:
|
||||
config: FactorConfig,
|
||||
*,
|
||||
regime_by_date: Mapping[object, Any] | None = None,
|
||||
) -> FactorResult:
|
||||
from app.services.heavy_job_limiter import shared_heavy_job_limiter
|
||||
|
||||
with shared_heavy_job_limiter.slot("exclusive"):
|
||||
return self._run(config, regime_by_date=regime_by_date)
|
||||
|
||||
def _run(
|
||||
self,
|
||||
config: FactorConfig,
|
||||
*,
|
||||
regime_by_date: Mapping[object, Any] | None = None,
|
||||
) -> FactorResult:
|
||||
t0 = time.perf_counter()
|
||||
run_id = uuid.uuid4().hex[:10]
|
||||
@@ -274,6 +221,17 @@ class FactorBacktestService:
|
||||
config: FactorBatchConfig,
|
||||
*,
|
||||
regime_by_date: Mapping[object, Any] | None = None,
|
||||
) -> FactorBatchResult:
|
||||
from app.services.heavy_job_limiter import shared_heavy_job_limiter
|
||||
|
||||
with shared_heavy_job_limiter.slot("exclusive"):
|
||||
return self._run_batch(config, regime_by_date=regime_by_date)
|
||||
|
||||
def _run_batch(
|
||||
self,
|
||||
config: FactorBatchConfig,
|
||||
*,
|
||||
regime_by_date: Mapping[object, Any] | None = None,
|
||||
) -> FactorBatchResult:
|
||||
"""在同一份 Panel 上依次评估多个因子, 避免重复读取和计算指标。"""
|
||||
t0 = time.perf_counter()
|
||||
@@ -357,6 +315,10 @@ class FactorBacktestService:
|
||||
**evaluate_kwargs,
|
||||
)
|
||||
long_short = result.long_short_stats
|
||||
# metrics_v2: 由 IC 序列推导 NW HAC t 值与 p 值 (日频 1 日前瞻 → 滞后 1)
|
||||
ic_values = [row.get("ic") for row in result.ic_series]
|
||||
t_naive = stats_v2.naive_t(ic_values)
|
||||
nw = stats_v2.newey_west_t(ic_values, lag=1)
|
||||
items.append(FactorBatchItem(
|
||||
factor_name=factor_name,
|
||||
label=str(meta.get("label", factor_name)),
|
||||
@@ -377,6 +339,14 @@ class FactorBacktestService:
|
||||
n_dates=result.n_dates,
|
||||
elapsed_ms=result.elapsed_ms,
|
||||
error=result.error,
|
||||
t_naive=t_naive,
|
||||
t_newey_west=nw[0] if nw else None,
|
||||
nw_lag=1 if nw else None,
|
||||
p_value=(
|
||||
stats_v2.normal_two_sided_p(nw[0]) if nw
|
||||
else stats_v2.normal_two_sided_p(t_naive) if t_naive is not None
|
||||
else None
|
||||
),
|
||||
))
|
||||
except Exception as exc: # 单因子失败不能中止整个筛选批次
|
||||
logger.exception("factor batch item failed: %s", factor_name)
|
||||
@@ -390,6 +360,10 @@ class FactorBacktestService:
|
||||
|
||||
n_symbols = max((item.n_symbols for item in items), default=0)
|
||||
n_dates = max((item.n_dates for item in items), default=0)
|
||||
# metrics_v2: 批内 BH-FDR q 值 (m = 可检验因子数, 计算失败项不计)
|
||||
q_values = stats_v2.bh_fdr_qvalues([item.p_value for item in items])
|
||||
for item, q_value in zip(items, q_values, strict=True):
|
||||
item.q_value = q_value
|
||||
return FactorBatchResult(
|
||||
run_id=run_id,
|
||||
config=result_config,
|
||||
@@ -701,8 +675,39 @@ class FactorBacktestService:
|
||||
logger.warning("factors %s cannot be computed, missing columns: %s", factor_cols, missing)
|
||||
return panel
|
||||
|
||||
# 扩展表因子 (ext_ base 条目) = 外部物化列, 指标补算管线不认识;
|
||||
# 请求的因子集合命中时在此按 (symbol, date) 时序对齐注入 (与
|
||||
# compute_signals 同一原语, 历史帧不含快照 → 无未来函数)。
|
||||
from app.factors import ext_factors
|
||||
|
||||
if factor_cols & ext_factors.ext_factor_ids():
|
||||
panel = ext_factors.attach_ext_columns(panel, include_snapshot=False)
|
||||
|
||||
from app.factors.registry import get_factor
|
||||
from app.indicators.pipeline import compute_indicators
|
||||
|
||||
# custom/composite 因子走注册表→DSL/组合 的同一条物化路径 (P3, 与策略评分共用);
|
||||
# 其底层依赖 (如 change_pct/ma20) 先经内置补算路径物化, 再做 DSL/组合物化。
|
||||
registry_names = {
|
||||
name for name in factor_cols
|
||||
if (spec := get_factor(name)) is not None and spec.kind in ("custom", "composite")
|
||||
}
|
||||
if registry_names:
|
||||
base_deps: set[str] = set()
|
||||
for name in registry_names:
|
||||
spec = get_factor(name)
|
||||
if spec is not None:
|
||||
base_deps.update(spec.dependencies)
|
||||
base_deps -= set(panel.columns) | registry_names
|
||||
if base_deps:
|
||||
panel = FactorBacktestService._compute_missing_factors(
|
||||
panel, base_deps, assume_sorted=assume_sorted,
|
||||
)
|
||||
panel = materialize_scoring_columns(panel, registry_names)
|
||||
factor_cols = factor_cols - registry_names
|
||||
if not factor_cols:
|
||||
return panel
|
||||
|
||||
derived = factor_cols & set(DERIVED_FACTOR_DEPENDENCIES)
|
||||
indicator_columns = factor_cols - derived
|
||||
for factor_name in derived:
|
||||
|
||||
@@ -103,11 +103,18 @@ def attach_fundamental_factors(
|
||||
columns = sorted(
|
||||
{FUNDAMENTAL_FACTORS[name]["column"] for name in missing_columns}
|
||||
)
|
||||
right = snapshot.select(["symbol", "_announce", *columns]).sort(["symbol", "_announce"])
|
||||
# asof 键取生效日 (公告日次日) 而非公告日: 直接用公告日回看会在换报告期的
|
||||
# 公告当日取到「尚未生效」的新一期并被门控置 null, 打断上一期的前向填充,
|
||||
# 与矩阵路径 (searchsorted side="right") 不一致。
|
||||
right = (
|
||||
snapshot.select(["symbol", "_announce", *columns])
|
||||
.with_columns(pl.col("_announce").dt.offset_by("1d").alias("_effective"))
|
||||
.sort(["symbol", "_effective"])
|
||||
)
|
||||
joined = panel.join_asof(
|
||||
right,
|
||||
left_on="date",
|
||||
right_on="_announce",
|
||||
right_on="_effective",
|
||||
by="symbol",
|
||||
strategy="backward",
|
||||
check_sortedness=False, # 双侧均已按 (symbol, key) 排序, 免除逐组检查开销
|
||||
@@ -128,7 +135,7 @@ def attach_fundamental_factors(
|
||||
expressions.append(
|
||||
pl.when(announced).then(value).otherwise(None).alias(name)
|
||||
)
|
||||
return joined.with_columns(expressions)
|
||||
return joined.with_columns(expressions).drop("_effective")
|
||||
|
||||
|
||||
def build_fundamental_matrices(
|
||||
@@ -174,9 +181,11 @@ def build_fundamental_matrices(
|
||||
continue
|
||||
for column, target in raw_columns.items():
|
||||
value = snapshot[column][row_index]
|
||||
if value is None or not np.isfinite(float(value)):
|
||||
continue
|
||||
target[start:, column_index] = float(value)
|
||||
numeric = float("nan") if value is None else float(value)
|
||||
# 新一期该指标为空时必须覆盖旧值为 NaN: 跳过写入会让同一行混用两期
|
||||
# 报告 (bps 取新期、roe 停在上一期), 与 polars 侧 join_asof 只认
|
||||
# 最新一期整行的口径不一致。
|
||||
target[start:, column_index] = numeric if np.isfinite(numeric) else np.nan
|
||||
|
||||
for name in requested:
|
||||
spec = FUNDAMENTAL_FACTORS[name]
|
||||
|
||||
@@ -1427,12 +1427,23 @@ def _populate_matrix_derived_arrays(
|
||||
if "turnover_rate" in wanted_fields and "turnover_rate" not in parquet_fields:
|
||||
float_shares = fields.get("float_shares")
|
||||
if float_shares is None:
|
||||
raise ValueError("matrix turnover_rate requires float_shares")
|
||||
_write_turnover_rate_matrix(
|
||||
fields["turnover_rate"],
|
||||
arrays["volume"],
|
||||
float_shares,
|
||||
)
|
||||
# 非股票资产 (etf/index) 无股本数据: instruments 无 float_shares 列,
|
||||
# 也无法从 parquet 读到 turnover_rate (数据源不提供, ETF 无换手率口径)。
|
||||
# 此时矩阵中该字段保持全 NaN 列 (matrix_fields 已占位), 与运行期
|
||||
# _optional_field 的降级语义一致, 供不需要换手率的策略正常回测。
|
||||
# 若本应有股本 (vector_fields 含 float_shares) 却取不到值, 才是数据
|
||||
# 异常, 由 _resolve_matrix_storage_fields 的 vector 装载路径显式失败。
|
||||
if "float_shares" in vector_fields:
|
||||
raise ValueError("matrix turnover_rate requires float_shares")
|
||||
logger.debug(
|
||||
"turnover_rate unavailable (asset has no float_shares); keeping NaN column"
|
||||
)
|
||||
else:
|
||||
_write_turnover_rate_matrix(
|
||||
fields["turnover_rate"],
|
||||
arrays["volume"],
|
||||
float_shares,
|
||||
)
|
||||
return names, latest_limits
|
||||
|
||||
|
||||
@@ -3757,6 +3768,12 @@ _MATRIX_COMPUTED_FEATURES = frozenset({
|
||||
"amihud_20d", "turnover_z_60d", "vol_price_corr_20d",
|
||||
"vwap_bias", "vol_trend_5_60",
|
||||
"limit_up_count_20d", "limit_up_count_60d",
|
||||
# --- 扩充批次 (2026-09-05): 与注册表/scoring 口径一致的 16 个新虚拟因子 ---
|
||||
"log_float_mv", "mom_accel_20_60", "rsi_14_delta_5d",
|
||||
"overnight_ret_20d", "intraday_ret_20d", "downside_vol_20d",
|
||||
"vol_regime_5_60", "amplitude_trend_20_60", "obv_trend_20d",
|
||||
"amount_mean_20d", "turnover_mean_20d", "turnover_std_20d",
|
||||
"position_240d", "distance_to_high_240d", "kdj_kd_diff",
|
||||
})
|
||||
|
||||
|
||||
@@ -4005,6 +4022,81 @@ def _compute_matrix_feature(market: MarketDataMatrix, name: str) -> np.ndarray:
|
||||
hits = np.where(np.isfinite(consecutive) & (consecutive > 0), np.float32(1.0), np.float32(0.0))
|
||||
hits = hits.astype(np.float32)
|
||||
return valid_rolling_sum(hits, close_valid, window)
|
||||
# --- 扩充批次 (2026-09-05): numpy 内核实现, 口径与 strategy/scoring.py 一致 ---
|
||||
if name == "log_float_mv":
|
||||
turnover = market.field("turnover_rate")
|
||||
valid = close_valid & np.isfinite(turnover) & (turnover > 0) & (market.volume > 0)
|
||||
out = np.full(market.shape, np.nan, dtype=np.float32)
|
||||
np.multiply(market.close, market.volume, out=out, where=valid)
|
||||
np.divide(out, turnover, out=out, where=valid)
|
||||
np.log(out, out=out, where=valid)
|
||||
return out
|
||||
if name == "mom_accel_20_60" or name == "kdj_kd_diff":
|
||||
left, right = (
|
||||
(matrix_feature(market, "momentum_20d"), matrix_feature(market, "momentum_60d"))
|
||||
if name == "mom_accel_20_60"
|
||||
else (matrix_feature(market, "kdj_k"), matrix_feature(market, "kdj_d"))
|
||||
)
|
||||
out = np.full(market.shape, np.nan, dtype=np.float32)
|
||||
np.subtract(left, right, out=out, where=np.isfinite(left) & np.isfinite(right))
|
||||
return out
|
||||
if name == "rsi_14_delta_5d":
|
||||
rsi = matrix_feature(market, "rsi_14")
|
||||
return rsi - valid_shift(rsi, 5, np.isfinite(rsi))
|
||||
if name == "overnight_ret_20d":
|
||||
overnight = _matrix_relative(market.open, matrix_feature(market, "prev_close"))
|
||||
return valid_rolling_sum(overnight, np.isfinite(overnight), 20)
|
||||
if name == "intraday_ret_20d":
|
||||
intraday = _matrix_relative(market.close, market.open)
|
||||
return valid_rolling_sum(intraday, np.isfinite(intraday), 20)
|
||||
if name == "downside_vol_20d":
|
||||
daily = matrix_feature(market, "change_pct")
|
||||
downside = np.where(
|
||||
np.isfinite(daily), np.minimum(daily, np.float32(0.0)), np.nan,
|
||||
).astype(np.float32)
|
||||
mean_sq = valid_rolling_mean(np.square(downside, dtype=np.float32), np.isfinite(downside), 20)
|
||||
out = np.full(market.shape, np.nan, dtype=np.float32)
|
||||
np.sqrt(mean_sq, out=out, where=np.isfinite(mean_sq))
|
||||
return out
|
||||
if name == "vol_regime_5_60":
|
||||
daily = matrix_feature(market, "change_pct")
|
||||
valid = np.isfinite(daily)
|
||||
return _matrix_ratio(
|
||||
valid_rolling_std(daily, valid, 5, ddof=1),
|
||||
valid_rolling_std(daily, valid, 60, ddof=1),
|
||||
)
|
||||
if name == "amplitude_trend_20_60":
|
||||
amplitude = matrix_feature(market, "amplitude")
|
||||
valid = np.isfinite(amplitude)
|
||||
return _matrix_relative(
|
||||
valid_rolling_mean(amplitude, valid, 20),
|
||||
valid_rolling_mean(amplitude, valid, 60),
|
||||
)
|
||||
if name == "obv_trend_20d":
|
||||
daily = matrix_feature(market, "change_pct")
|
||||
volume_valid = close_valid & np.isfinite(market.volume)
|
||||
signed = np.where(
|
||||
np.isfinite(daily), np.sign(daily) * market.volume, np.nan,
|
||||
).astype(np.float32)
|
||||
total = valid_rolling_sum(signed, volume_valid & np.isfinite(daily), 20)
|
||||
scale = valid_rolling_mean(market.volume, volume_valid, 20) * np.float32(20.0)
|
||||
return _matrix_ratio(total, scale)
|
||||
if name == "amount_mean_20d":
|
||||
amount = market.field("amount")
|
||||
return valid_rolling_mean(amount / np.float32(1e8), np.isfinite(amount), 20)
|
||||
if name == "turnover_mean_20d" or name == "turnover_std_20d":
|
||||
turnover = market.field("turnover_rate")
|
||||
valid = np.isfinite(turnover)
|
||||
mean = valid_rolling_mean(turnover, valid, 20)
|
||||
if name == "turnover_mean_20d":
|
||||
return mean
|
||||
return _matrix_ratio(valid_rolling_std(turnover, valid, 20, ddof=1), mean)
|
||||
if name == "position_240d":
|
||||
high = valid_rolling_max(market.close, close_valid, 240)
|
||||
low = valid_rolling_min(market.close, close_valid, 240)
|
||||
return _matrix_ratio(market.close - low, high - low)
|
||||
if name == "distance_to_high_240d":
|
||||
return _matrix_relative(market.close, valid_rolling_max(market.close, close_valid, 240))
|
||||
raise ValueError(f"unsupported matrix feature: {name}")
|
||||
|
||||
|
||||
|
||||
@@ -17,7 +17,6 @@ import numpy as np
|
||||
import polars as pl
|
||||
|
||||
from app.backtest.factor import (
|
||||
FACTOR_COLUMNS,
|
||||
FACTOR_METHODOLOGY_VERSION,
|
||||
FACTOR_WARMUP_DAYS,
|
||||
FactorBacktestService,
|
||||
@@ -57,6 +56,7 @@ from app.enriched_generation import (
|
||||
EnrichedGenerationUnavailableError,
|
||||
enriched_publication_incomplete,
|
||||
)
|
||||
from app.factors.registry import factor_columns_view
|
||||
from app.services.mining_jobs import MiningRunStore
|
||||
from app.services.mining_preflight import enriched_partition_dates
|
||||
from app.services.mining_schedule import MINING_ALGORITHM_VERSION
|
||||
@@ -66,7 +66,6 @@ from app.strategy.engine import StrategyEngine
|
||||
ProgressCallback = Callable[[dict[str, Any]], None]
|
||||
CancelCheck = Callable[[], bool] | Any
|
||||
_PROFILE_NAMES = frozenset({"exploratory", "balanced", "strict"})
|
||||
_FACTOR_IDS = frozenset(str(item["id"]) for item in FACTOR_COLUMNS)
|
||||
_MINING_MATRIX_CACHE_BYTES = 32 * 1024 * 1024
|
||||
_RESULT_POLICY = BacktestResultPolicy(
|
||||
required_stats=frozenset({"total_return", "sharpe", "max_drawdown", "n_trades"}),
|
||||
@@ -807,7 +806,9 @@ def _decode_runtime_request(
|
||||
factor_names = tuple(str(value) for value in request.get("factor_names") or ())
|
||||
if not factor_names or len(set(factor_names)) != len(factor_names):
|
||||
raise ValueError("factor_names must be non-empty and unique")
|
||||
unknown_factors = sorted(set(factor_names) - _FACTOR_IDS)
|
||||
# 注册表动态读取: worker 子进程已在入口加载自定义/复合因子
|
||||
known_ids = frozenset(str(item["id"]) for item in factor_columns_view())
|
||||
unknown_factors = sorted(set(factor_names) - known_ids)
|
||||
if unknown_factors:
|
||||
raise ValueError(f"unknown mining factors: {unknown_factors}")
|
||||
if len(factor_names) > 48:
|
||||
@@ -1037,7 +1038,7 @@ def _build_artifacts(
|
||||
candidate.factor_names, candidate.directions, strict=True
|
||||
):
|
||||
direction_by_factor.setdefault(factor_name, int(direction))
|
||||
metadata = {str(item["id"]): item for item in FACTOR_COLUMNS}
|
||||
metadata = {str(item["id"]): item for item in factor_columns_view()}
|
||||
factor_rows = []
|
||||
for factor_name in request.factor_names:
|
||||
metric = latest_metrics[factor_name]
|
||||
|
||||
@@ -60,7 +60,9 @@ def _candidates_for(param_id: str, spec, pmeta: dict) -> list:
|
||||
raise ValueError(f"参数 '{param_id}' 的 max < min")
|
||||
step = float(step)
|
||||
# 整数计数生成候选, 避免浮点累加误差丢端点 (如 0.1/0.1 步长)。
|
||||
n_steps = round((hi - lo) / step)
|
||||
# 步数向下取整: (hi-lo) 不是 step 整数倍时, 四舍五入会多造一个越过 hi 的候选
|
||||
# (1~20 步长 7 → 22), 用户填的上限反而被越界校验拒绝。1e-9 容差保住整除端点。
|
||||
n_steps = int((hi - lo) / step + 1e-9)
|
||||
raw = [round(lo + i * step, 10) for i in range(n_steps + 1)]
|
||||
else:
|
||||
raise ValueError(f"参数 '{param_id}' 的网格 spec 必须是列表或 {{min,max,step}} 字典")
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
"""metrics_v2 统计函数 (P3) — Newey-West HAC t 值 / BH-FDR q 值 / DSR。
|
||||
|
||||
运行时零新增第三方依赖 (后端无 scipy/statsmodels), 全部 numpy 手写;
|
||||
数值测试用固定黄金参考向量锁定 (tests/test_stats_v2.py)。
|
||||
|
||||
口径 (设计文档 factor-system-design.md §6):
|
||||
- IC 序列因 h 日前瞻收益存在 h-1 阶移动平均自相关, 主口径 t 值取 NW HAC, 滞后 L=h。
|
||||
- 多因子批量检验按 Benjamini-Hochberg 步进法控制 FDR。
|
||||
- DSR (Deflated Sharpe Ratio, Bailey & Lopez de Prado 2014) 用于多重试验校正后的
|
||||
夏普显著性; 期望最大夏普 EM = sqrt(V[SR]) * ((1-gamma)Φ^-1(1-1/N) + gammaΦ^-1(1-1/(Ne)))
|
||||
其中 gamma 为欧拉-马歇罗尼常数。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
|
||||
EULER_GAMMA = 0.5772156649015329
|
||||
|
||||
|
||||
def _clean_values(values) -> np.ndarray:
|
||||
array = np.asarray([value for value in values if value is not None and np.isfinite(value)], dtype=float)
|
||||
return array
|
||||
|
||||
|
||||
def newey_west_t(values, lag: int) -> tuple[float, float, float] | None:
|
||||
"""Newey-West HAC 稳健 t 统计量 (Bartlett 核)。
|
||||
|
||||
返回 (t值, 均值, NW标准误); 样本不足 (n <= lag+2) 或方差为零返回 None。
|
||||
"""
|
||||
array = _clean_values(values)
|
||||
n = array.size
|
||||
if n <= lag + 2 or n < 3:
|
||||
return None
|
||||
mean = float(array.mean())
|
||||
centered = array - mean
|
||||
# 长方差 S = gamma0 + 2 Σ_l w_l gamma_l, w_l = 1 - l/(lag+1) (Bartlett)
|
||||
gamma = [float(np.dot(centered[: n - lag_i], centered[lag_i:]) / n) for lag_i in range(lag + 1)]
|
||||
long_variance = gamma[0]
|
||||
for lag_i in range(1, lag + 1):
|
||||
weight = 1.0 - lag_i / (lag + 1)
|
||||
long_variance += 2.0 * weight * gamma[lag_i]
|
||||
long_variance = max(long_variance, 0.0)
|
||||
nw_se = math.sqrt(long_variance / n)
|
||||
if nw_se == 0:
|
||||
return None
|
||||
return (mean - 0.0) / nw_se, mean, nw_se
|
||||
|
||||
|
||||
def naive_t(values) -> float | None:
|
||||
array = _clean_values(values)
|
||||
n = array.size
|
||||
if n < 3:
|
||||
return None
|
||||
std = float(array.std(ddof=1))
|
||||
if std == 0:
|
||||
return None
|
||||
return float(array.mean()) / (std / math.sqrt(n))
|
||||
|
||||
|
||||
def normal_two_sided_p(t_stat: float) -> float:
|
||||
"""标准正态双侧 p 值: erfc(|t|/sqrt(2))。"""
|
||||
return math.erfc(abs(t_stat) / math.sqrt(2.0))
|
||||
|
||||
|
||||
def bh_fdr_qvalues(pvalues: list[float | None]) -> list[float | None]:
|
||||
"""Benjamini-Hochberg 步进法 q 值 (与输入等长, None 透传)。
|
||||
|
||||
m 取可检验假设数 (None 不计入); q_i = min over j>=rank_i { p_j * m / rank_j },
|
||||
从大到小单调回填保证递增约束。
|
||||
"""
|
||||
indexed = [
|
||||
(index, p) for index, p in enumerate(pvalues)
|
||||
if p is not None and np.isfinite(p)
|
||||
]
|
||||
qvalues: list[float | None] = [None] * len(pvalues)
|
||||
if not indexed:
|
||||
return qvalues
|
||||
m = len(indexed)
|
||||
indexed.sort(key=lambda pair: pair[1])
|
||||
running_min = float("inf")
|
||||
for reverse_rank in range(len(indexed) - 1, -1, -1):
|
||||
index, p = indexed[reverse_rank]
|
||||
rank = reverse_rank + 1
|
||||
candidate = p * m / rank
|
||||
running_min = min(running_min, candidate)
|
||||
qvalues[index] = min(1.0, running_min)
|
||||
return qvalues
|
||||
|
||||
|
||||
def _normal_ppf(probability: float) -> float:
|
||||
"""标准正态分位数 Acklam 逆逼近 (相对误差 < 1.15e-9), 零依赖替代 scipy.stats.norm.ppf。"""
|
||||
if not (0.0 < probability < 1.0):
|
||||
raise ValueError("probability 必须在 (0,1) 开区间")
|
||||
a = (-3.969683028665376e+01, 2.209460984245205e+02, -2.759285104469687e+02,
|
||||
1.383577518672690e+02, -3.066479806614716e+01, 2.506628277459239e+00)
|
||||
b = (-5.447609879822406e+01, 1.615858368580409e+02, -1.556989798598866e+02,
|
||||
6.680131188771972e+01, -1.328068155288572e+01)
|
||||
c = (-7.784894002430293e-03, -3.223964580411365e-01, -2.400758277161838e+00,
|
||||
-2.549732539343734e+00, 4.374664141464968e+00, 2.938163982698783e+00)
|
||||
d = (7.784695709041462e-03, 3.224671290700398e-01, 2.445134137142996e+00,
|
||||
3.754408661907416e+00)
|
||||
p_low, p_high = 0.02425, 1 - 0.02425
|
||||
if probability < p_low:
|
||||
q_value = math.sqrt(-2 * math.log(probability))
|
||||
return (((((c[0] * q_value + c[1]) * q_value + c[2]) * q_value + c[3]) * q_value + c[4]) * q_value + c[5]) / \
|
||||
((((d[0] * q_value + d[1]) * q_value + d[2]) * q_value + d[3]) * q_value + 1)
|
||||
if probability <= p_high:
|
||||
q_value = probability - 0.5
|
||||
r = q_value * q_value
|
||||
return (((((a[0] * r + a[1]) * r + a[2]) * r + a[3]) * r + a[4]) * r + a[5]) * q_value / \
|
||||
(((((b[0] * r + b[1]) * r + b[2]) * r + b[3]) * r + b[4]) * r + 1)
|
||||
q_value = math.sqrt(-2 * math.log(1 - probability))
|
||||
return -(((((c[0] * q_value + c[1]) * q_value + c[2]) * q_value + c[3]) * q_value + c[4]) * q_value + c[5]) / \
|
||||
((((d[0] * q_value + d[1]) * q_value + d[2]) * q_value + d[3]) * q_value + 1)
|
||||
|
||||
|
||||
def expected_max_sharpe(n_trials: int, variance_sharpes: float) -> float:
|
||||
"""N 次独立试验的期望最大夏普 EM (方差>0 时); 单次试验不校正。"""
|
||||
if n_trials <= 1 or variance_sharpes <= 0:
|
||||
return 0.0
|
||||
z1 = _normal_ppf(1.0 - 1.0 / n_trials)
|
||||
z2 = _normal_ppf(1.0 - 1.0 / (n_trials * math.e))
|
||||
return math.sqrt(variance_sharpes) * ((1.0 - EULER_GAMMA) * z1 + EULER_GAMMA * z2)
|
||||
|
||||
|
||||
def deflated_sharpe_psr(
|
||||
sharpe: float,
|
||||
n_obs: int,
|
||||
skewness: float | None = None,
|
||||
kurtosis: float | None = None,
|
||||
expected_max_sharpe: float = 0.0,
|
||||
) -> float | None:
|
||||
"""Deflated Sharpe (PSR 对 EM 校正) 概率; 参数不足或退化返回 None。
|
||||
|
||||
PSR = Φ( (SR - SR*) * sqrt(n-1) / sqrt(1 - gamma3 SR + (gamma4-1)/4 SR^2) )
|
||||
"""
|
||||
if n_obs < 5 or not np.isfinite(sharpe):
|
||||
return None
|
||||
skewness = 0.0 if skewness is None else skewness
|
||||
kurtosis = 3.0 if kurtosis is None else kurtosis
|
||||
denominator = 1.0 - skewness * sharpe + (kurtosis - 1.0) / 4.0 * sharpe * sharpe
|
||||
if denominator <= 0:
|
||||
return None
|
||||
statistic = (sharpe - expected_max_sharpe) * math.sqrt(n_obs - 1) / math.sqrt(denominator)
|
||||
return 0.5 * (1.0 + math.erf(statistic / math.sqrt(2.0)))
|
||||
@@ -487,20 +487,42 @@ _SHARE_CAP_FILTER_KEYS = (
|
||||
"float_cap_max",
|
||||
)
|
||||
|
||||
# 换手率界同样依赖股本派生字段 (turnover_rate ← float_shares):
|
||||
# 非股票资产 (etf/index) 没有股本数据, 若保留非 None 的换手率界,
|
||||
# _basic_filter_dependencies 会解析出 turnover_rate 字段需求,
|
||||
# 矩阵缓存档构建时因无 float_shares 而失败 (matrix turnover_rate requires
|
||||
# float_shares)。与市值界同一族问题, 必须一并中和。
|
||||
_TURNOVER_FILTER_KEYS = (
|
||||
"turnover_min",
|
||||
"turnover_max",
|
||||
)
|
||||
|
||||
# 股票专属的价格界与板块过滤对非股票资产同样不可满足 (#215):
|
||||
# ETF 单价普遍 0.5~7 元, 会被 price_min=3 整列误杀; boards 按股票代码
|
||||
# 前缀匹配, ETF 代码不属于任何板块 → 掩码全 False, 静默零信号。
|
||||
_STOCK_ONLY_FILTER_KEYS = (
|
||||
*_SHARE_CAP_FILTER_KEYS,
|
||||
*_TURNOVER_FILTER_KEYS,
|
||||
"price_min",
|
||||
"price_max",
|
||||
"boards",
|
||||
)
|
||||
|
||||
|
||||
def _basic_filter_for_asset(basic_filter: dict, asset_type: str) -> dict:
|
||||
"""非股票资产没有股本数据 (etf/index 维表只有 symbol/name), 市值与流通
|
||||
市值界对它们既无意义也不可满足: 依赖解析前先置 None, 避免解析出
|
||||
total_shares/float_shares 字段需求导致矩阵加载直接失败。
|
||||
"""非股票资产没有股本数据 (etf/index 维表只有 symbol/name), 市值、流通
|
||||
市值与换手率界对它们既无意义也不可满足: 依赖解析与运行期过滤前先置
|
||||
None。价格界 (price_min/max) 与板块过滤 (boards) 是股票专属口径, 对
|
||||
ETF 同样不可满足, 一并中和, 否则入场候选在运行期被静默清零 (#215)。
|
||||
|
||||
运行期过滤无需同步修改 —— polars 侧有列守卫 (engine._basic_filter_expr),
|
||||
矩阵侧 _optional_field 对缺失字段返回全 NaN 且 _apply_bound 跳过全 NaN
|
||||
界, 二者对缺失股本列本就降级为 no-op。
|
||||
置 None 后: 依赖解析不再产出 total_shares/float_shares/turnover_rate
|
||||
需求; polars 侧有列守卫 (engine._basic_filter_expr), 矩阵侧
|
||||
_optional_field 对缺失字段返回全 NaN 且 _apply_bound 跳过全 NaN 界。
|
||||
"""
|
||||
if asset_type == "stock" or not basic_filter:
|
||||
return basic_filter
|
||||
sanitized = dict(basic_filter)
|
||||
for key in _SHARE_CAP_FILTER_KEYS:
|
||||
for key in _STOCK_ONLY_FILTER_KEYS:
|
||||
sanitized[key] = None
|
||||
return sanitized
|
||||
|
||||
@@ -569,6 +591,7 @@ class StrategyBacktestResult:
|
||||
trades: list[dict] = field(default_factory=list)
|
||||
per_symbol_stats: list[dict] = field(default_factory=list)
|
||||
strategy_info: dict = field(default_factory=dict)
|
||||
factor_attribution: dict | None = None
|
||||
elapsed_ms: float = 0.0
|
||||
error: str | None = None
|
||||
|
||||
@@ -633,6 +656,56 @@ class BacktestResultPolicy:
|
||||
return {key: value for key, value in stats.items() if key in keep}
|
||||
|
||||
|
||||
def _factor_attribution_summary(
|
||||
snapshot: pl.DataFrame,
|
||||
trades: list,
|
||||
) -> dict | None:
|
||||
"""v1 因子归因: 入场信号日因子快照 x 成交盈亏, 对比盈利/亏损单因子均值。
|
||||
|
||||
snapshot 来自 _apply_score 物化的候选行 (与评分同一条计算管线), 模拟结束后
|
||||
按 (symbol, 信号日) 关联成交。快照缺失、无可关联行或因子列全空时返回 None,
|
||||
归因失败不影响回测主结果。
|
||||
"""
|
||||
factor_cols = [c for c in snapshot.columns if c not in ("symbol", "date")]
|
||||
if not factor_cols or not trades:
|
||||
return None
|
||||
normalized = snapshot.with_columns(
|
||||
pl.col("date").cast(pl.Utf8).str.slice(0, 10).alias("date")
|
||||
)
|
||||
symbols: list[str] = []
|
||||
days: list[str] = []
|
||||
pnls: list[float] = []
|
||||
for trade in trades:
|
||||
day = trade.entry_signal_date or trade.entry_date
|
||||
if day is None:
|
||||
continue
|
||||
symbols.append(trade.symbol)
|
||||
days.append(str(day)[:10])
|
||||
pnls.append(float(trade.pnl_pct))
|
||||
if not symbols:
|
||||
return None
|
||||
frame = pl.DataFrame({"symbol": symbols, "date": days, "pnl_pct": pnls})
|
||||
joined = frame.join(normalized, on=["symbol", "date"], how="left")
|
||||
win = joined.filter(pl.col("pnl_pct") > 0)
|
||||
lose = joined.filter(pl.col("pnl_pct") <= 0)
|
||||
factors: list[dict] = []
|
||||
for col in factor_cols:
|
||||
win_vals = win.get_column(col).drop_nulls().cast(pl.Float64)
|
||||
lose_vals = lose.get_column(col).drop_nulls().cast(pl.Float64)
|
||||
if win_vals.is_empty() and lose_vals.is_empty():
|
||||
continue
|
||||
factors.append({
|
||||
"factor": col,
|
||||
"win_mean": round(float(win_vals.mean()), 6) if not win_vals.is_empty() else None,
|
||||
"lose_mean": round(float(lose_vals.mean()), 6) if not lose_vals.is_empty() else None,
|
||||
"win_n": int(win_vals.len()),
|
||||
"lose_n": int(lose_vals.len()),
|
||||
})
|
||||
if not factors:
|
||||
return None
|
||||
return {"factors": factors, "n_win": win.height, "n_lose": lose.height}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PreparedMatrixBacktest:
|
||||
"""Job-scoped immutable market data reused by every optimizer trial."""
|
||||
@@ -834,7 +907,11 @@ class StrategyBacktestService:
|
||||
)
|
||||
|
||||
overrides = first.overrides or {}
|
||||
basic_filter = self._effective_basic_filter(strategy, overrides)
|
||||
# 运行期过滤用的也是同一份 basic_filter: 在入口处按资产类型中和,
|
||||
# 否则 boards/price_min 会在掩码阶段静默清零 ETF 候选 (#215)
|
||||
basic_filter = _basic_filter_for_asset(
|
||||
self._effective_basic_filter(strategy, overrides), first.asset_type
|
||||
)
|
||||
entry_signals = self._effective_signals(overrides, "entry_signals", strategy.entry_signals)
|
||||
exit_signals = self._effective_signals(overrides, "exit_signals", strategy.exit_signals)
|
||||
resolver = StrategyDependencyResolver()
|
||||
@@ -986,6 +1063,8 @@ class StrategyBacktestService:
|
||||
t0 = time.perf_counter()
|
||||
run_id = uuid.uuid4().hex[:10]
|
||||
result_policy = result_policy or BacktestResultPolicy()
|
||||
# 因子归因快照容器: 日线路径在 _apply_score 里填充, 其余路径保持空
|
||||
factor_snapshot: dict = {}
|
||||
|
||||
def _err(msg: str) -> StrategyBacktestResult:
|
||||
return StrategyBacktestResult(
|
||||
@@ -1011,7 +1090,10 @@ class StrategyBacktestService:
|
||||
|
||||
params = self._normalize_params(config.params or {}, s)
|
||||
overrides = config.overrides or {}
|
||||
basic_filter = self._effective_basic_filter(s, overrides)
|
||||
# 同回测 run 路径: 挖掘运行期也要按资产类型中和股票专属过滤键 (#215)
|
||||
basic_filter = _basic_filter_for_asset(
|
||||
self._effective_basic_filter(s, overrides), config.asset_type
|
||||
)
|
||||
entry_signals = self._effective_signals(overrides, "entry_signals", s.entry_signals)
|
||||
exit_signals = self._effective_signals(overrides, "exit_signals", s.exit_signals)
|
||||
if config.exit_fill == "signal_next_minute":
|
||||
@@ -1476,7 +1558,7 @@ class StrategyBacktestService:
|
||||
|
||||
candidate_filter_mask = self._build_candidate_filter_mask(panel, s, params)
|
||||
candidate_mask = basic_mask & candidate_filter_mask
|
||||
panel = self._apply_score(panel, s, overrides, universe_mask=candidate_mask)
|
||||
panel = self._apply_score(panel, s, overrides, universe_mask=candidate_mask, factor_snapshot=factor_snapshot)
|
||||
formal_candidate_mask = candidate_mask & formal_range
|
||||
entry_mask = self._build_entry_mask_from_candidate(panel, candidate_mask, s, entry_signals)
|
||||
entry_mask = entry_mask & formal_range
|
||||
@@ -1633,6 +1715,16 @@ class StrategyBacktestService:
|
||||
|
||||
selected_stats = result_policy.select_stats(result.stats)
|
||||
|
||||
# 因子归因 (fail-open): 快照与成交按信号日关联, 失败只记日志不影响结果
|
||||
factor_attribution = None
|
||||
if factor_snapshot and result.trades and result_policy.include_trades:
|
||||
try:
|
||||
factor_attribution = _factor_attribution_summary(
|
||||
factor_snapshot["frame"], result.trades
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("factor attribution failed: %s", exc)
|
||||
|
||||
elapsed = (time.perf_counter() - t0) * 1000
|
||||
|
||||
return StrategyBacktestResult(
|
||||
@@ -1653,6 +1745,7 @@ class StrategyBacktestService:
|
||||
else []
|
||||
),
|
||||
strategy_info=strategy_info,
|
||||
factor_attribution=factor_attribution,
|
||||
elapsed_ms=round(elapsed, 1),
|
||||
)
|
||||
|
||||
@@ -2436,6 +2529,7 @@ class StrategyBacktestService:
|
||||
s: StrategyDef,
|
||||
overrides: dict | None,
|
||||
universe_mask: pl.Series | None = None,
|
||||
factor_snapshot: dict | None = None,
|
||||
) -> pl.DataFrame:
|
||||
scoring = effective_scoring(s.meta.get("scoring"), overrides)
|
||||
directions = effective_scoring_directions(overrides)
|
||||
@@ -2446,6 +2540,18 @@ class StrategyBacktestService:
|
||||
if has_universe:
|
||||
work = work.with_columns(universe_mask.rename("_score_universe"))
|
||||
|
||||
# 因子归因快照: 在临时因子列被 _finish 丢弃前, 截取候选行的
|
||||
# (symbol, date, 因子值)。与评分共用同一份物化结果, 无第二次计算。
|
||||
if factor_snapshot is not None:
|
||||
snapshot_cols = ["symbol", "date"] + [
|
||||
name for name in scoring if name in work.columns
|
||||
]
|
||||
if len(snapshot_cols) > 2:
|
||||
frame = work
|
||||
if has_universe:
|
||||
frame = frame.filter(pl.col("_score_universe"))
|
||||
factor_snapshot["frame"] = frame.select(snapshot_cols)
|
||||
|
||||
def _value_in_universe(value: pl.Expr) -> pl.Expr:
|
||||
if has_universe:
|
||||
return pl.when(pl.col("_score_universe")).then(value).otherwise(None)
|
||||
|
||||
@@ -149,14 +149,21 @@ class WalkForwardService:
|
||||
self.strategy_engine = strategy_engine
|
||||
|
||||
def _prepare_shared_matrix(self, cfg: WalkForwardConfig, folds: list[Fold]):
|
||||
"""Build one immutable superset matrix for every matrix-native fold."""
|
||||
"""Build one immutable superset matrix for every matrix-native fold.
|
||||
|
||||
返回 None 时 run() 走通用路径: 每折独立优化 + OOS 回测, 正确但无共享矩阵加速。
|
||||
python_history_legacy (filter_history) 与 polars_expr (内置) 策略无法装入
|
||||
共享矩阵, 走通用路径; composite / minute_filter 仍不支持, 保持 fail-closed。
|
||||
"""
|
||||
if self.strategy_engine is None or not folds:
|
||||
return None
|
||||
strategy = self.strategy_engine.get(cfg.strategy_id)
|
||||
if strategy.execution_backend != "matrix_native":
|
||||
if strategy.execution_backend in ("python_history_legacy", "polars_expr"):
|
||||
return None
|
||||
raise ValueError(
|
||||
f"步进优化暂仅支持矩阵(matrix_native)策略; "
|
||||
f"{cfg.strategy_id} 是 {strategy.execution_backend}"
|
||||
f"步进优化暂仅支持矩阵(matrix_native)/日线历史(python_history_legacy/"
|
||||
f"polars_expr)策略; {cfg.strategy_id} 是 {strategy.execution_backend}"
|
||||
)
|
||||
|
||||
from app.backtest.optimizer import expand_param_grid
|
||||
|
||||
@@ -166,6 +166,15 @@ def _attach_worker_metrics(
|
||||
result["worker"] = metrics
|
||||
|
||||
|
||||
def _error_message(exc: BaseException) -> str:
|
||||
"""任务级错误文案: enriched 发布类失败对用户是"稍后再试", 不透出原始异常。"""
|
||||
from app.enriched_generation import EnrichedGenerationUnavailableError
|
||||
|
||||
if isinstance(exc, EnrichedGenerationUnavailableError):
|
||||
return "指标数据正在发布更新,请稍后重试"
|
||||
return str(exc)
|
||||
|
||||
|
||||
def _worker_entry(task: dict[str, Any], event_queue, cancel_event) -> None:
|
||||
sampler = _PeakRssSampler()
|
||||
sampler.start()
|
||||
@@ -182,6 +191,12 @@ def _worker_entry(task: dict[str, Any], event_queue, cancel_event) -> None:
|
||||
data_dir = Path(task["data_dir"])
|
||||
store = DataStore(data_dir)
|
||||
repo = KlineRepository(store)
|
||||
# 子进程不继承主进程的因子注册表; 自定义/复合因子 (uf_/cf_) 在任何
|
||||
# 涉及因子物化的 worker 任务里都依赖注册表, 启动时从存储加载。
|
||||
# 单个加载失败只跳过 (fail-open 跳过该因子), 与主进程启动行为一致。
|
||||
from app.factors.store import load_into_registry
|
||||
|
||||
load_into_registry(data_dir)
|
||||
strategy_engine = StrategyEngine(
|
||||
strategy_dirs=_strategy_dirs(data_dir),
|
||||
override_loader=lambda sid: strategy_config.load_override(data_dir, sid),
|
||||
@@ -250,7 +265,7 @@ def _worker_entry(task: dict[str, Any], event_queue, cancel_event) -> None:
|
||||
sampler.stop()
|
||||
event_queue.put({
|
||||
"type": "error",
|
||||
"message": str(exc),
|
||||
"message": _error_message(exc),
|
||||
"traceback": traceback.format_exc(),
|
||||
})
|
||||
finally:
|
||||
|
||||
@@ -112,6 +112,29 @@ class Settings(BaseSettings):
|
||||
backtest_matrix_cache_prewarm: bool = True
|
||||
backtest_matrix_cache_prewarm_years: int = 5
|
||||
|
||||
# polars collect 并发闸 — polars 共享执行器在多线程并发 collect 下存在死锁
|
||||
# (上游 #24448/#25754 同族), 限流并发是社区验证的缓解手段。background 限额
|
||||
# 保证预热/增量等后台计算不占满闸位饿死页面读请求。
|
||||
polars_collect_permits: int = 4
|
||||
polars_collect_background_permits: int = 2
|
||||
|
||||
# 后端自愈看门狗 — 探测 collect 闸与全局写锁, 连续失败即退出交由
|
||||
# supervisor 拉起 (见 app/watchdog.py)。误伤防护靠保守阈值。
|
||||
watchdog_enabled: bool = True
|
||||
watchdog_interval_s: float = 30.0
|
||||
watchdog_probe_timeout_s: float = 15.0
|
||||
watchdog_failure_threshold: int = 2
|
||||
|
||||
# 策略批量执行 (run_all / 策略页全量跑) 的并发 worker 上限。实测 2026-09-07:
|
||||
# polars eager 操作内部已多线程并行, 外层再并发 4 worker 属超订, 41 策略
|
||||
# 299.6s 慢于串行 — 默认 1 (串行)。保留开关供配合 POLARS_MAX_THREADS 调优实验。
|
||||
strategy_run_all_workers: int = 1
|
||||
|
||||
# run_all 渐进式返回: HTTP 同步等待时限 (秒)。策略按历史耗时升序执行,
|
||||
# 到点后已算完的随响应返回, 未算完的转后台继续算并逐个写入策略缓存,
|
||||
# 前端轮询 cached-summary 点亮卡片。0 = 关闭 (整段阻塞, 旧行为)。
|
||||
strategy_run_all_first_return_s: float = 15.0
|
||||
|
||||
# Auth — 首次启动时预置访问密码(明文, 仅用于初始化, 详见 services/auth.bootstrap_from_env)
|
||||
# 公网服务器部署时免去 SSH 端口转发设密码的麻烦。写入 auth.json(哈希)后即不再读取。
|
||||
auth_password: str = ""
|
||||
@@ -140,6 +163,20 @@ class Settings(BaseSettings):
|
||||
raise ValueError("ai_max_output_tokens must be positive")
|
||||
if self.ai_context_window <= 0:
|
||||
raise ValueError("ai_context_window must be positive")
|
||||
if self.polars_collect_permits < 2:
|
||||
raise ValueError("polars_collect_permits must be >= 2")
|
||||
if not 1 <= self.polars_collect_background_permits < self.polars_collect_permits:
|
||||
raise ValueError(
|
||||
"polars_collect_background_permits must be in [1, polars_collect_permits)"
|
||||
)
|
||||
if self.watchdog_interval_s <= 0 or self.watchdog_probe_timeout_s <= 0:
|
||||
raise ValueError("watchdog intervals must be positive")
|
||||
if self.watchdog_failure_threshold < 1:
|
||||
raise ValueError("watchdog_failure_threshold must be >= 1")
|
||||
if self.strategy_run_all_workers < 1:
|
||||
raise ValueError("strategy_run_all_workers must be >= 1")
|
||||
if self.strategy_run_all_first_return_s < 0:
|
||||
raise ValueError("strategy_run_all_first_return_s must be >= 0")
|
||||
return self
|
||||
|
||||
@property
|
||||
|
||||
@@ -23,6 +23,7 @@ class ProviderCapabilities:
|
||||
adj_factor: bool = False
|
||||
minute: bool = False
|
||||
realtime: bool = False
|
||||
depth5: bool = False
|
||||
financial: bool = False
|
||||
|
||||
|
||||
@@ -73,3 +74,6 @@ class MarketDataProvider(Protocol):
|
||||
symbols: list[str] | None = None,
|
||||
) -> pl.DataFrame:
|
||||
"""Return normalized realtime quotes. Implementations may return empty."""
|
||||
|
||||
def get_depth_batch(self, symbols: list[str]) -> dict[str, dict]:
|
||||
"""Return five-level order books keyed by symbol."""
|
||||
|
||||
@@ -3,8 +3,8 @@
|
||||
能力 (capability) = 一个标准化数据集 (CONTRIBUTING「数据源插件化要求」):
|
||||
daily / adj_factor / realtime / minute / depth5 / financial (注册表顺序即设置页卡片顺序)。注册表集中声明每个
|
||||
能力的展示元数据、路由偏好字段与 TickFlow 档位要求, 前端设置页不再各自硬编码。
|
||||
depth5 目前仅 TickFlow 供 (插件数据集白名单未开放, 见 loader), 仍进矩阵是为了
|
||||
可用性门控诚实: 五档不可用时连板梯队封单/看板封单缺数据应有提示。
|
||||
depth5 与其他数据集一样可由插件声明并独立路由; 五档不可用时连板梯队封单/
|
||||
看板封单通过 usable 给出缺数据提示。
|
||||
|
||||
build_capability_matrix 把注册表、插件/自定义源的能力声明 (datasets) 和当前
|
||||
路由偏好合并为一个矩阵, 供设置页一次拉全。当前偏好由 API 层注入
|
||||
@@ -67,7 +67,6 @@ CAPABILITY_REGISTRY: list[dict] = [
|
||||
"field": "depth5_data_provider",
|
||||
"default": "tickflow",
|
||||
"tf_tier": "pro",
|
||||
# 插件契约暂未开放 depth5 数据集 (loader 白名单), 当前仅 TickFlow 供
|
||||
},
|
||||
{
|
||||
"id": "financial",
|
||||
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
@@ -132,6 +133,23 @@ class GenericHTTPProvider:
|
||||
)
|
||||
return errors
|
||||
|
||||
def _request_rows_retry(
|
||||
self, cfg, symbols: list[str], *, start_time=None, end_time=None, retries: int = 1
|
||||
) -> list[dict]:
|
||||
"""单批请求 + 短退避重试。仍失败抛出, 由调用方决定隔离粒度 (#226)。"""
|
||||
last: Exception | None = None
|
||||
for attempt in range(retries + 1):
|
||||
try:
|
||||
return self._request_rows(
|
||||
cfg, symbols=symbols, start_time=start_time, end_time=end_time
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
last = e
|
||||
if attempt < retries:
|
||||
time.sleep(1.0 * (attempt + 1))
|
||||
assert last is not None
|
||||
raise last
|
||||
|
||||
def get_daily(
|
||||
self,
|
||||
symbols: list[str],
|
||||
@@ -143,15 +161,35 @@ class GenericHTTPProvider:
|
||||
cfg = self._dataset("daily")
|
||||
frames: list[pl.DataFrame] = []
|
||||
chunks = chunked(symbols, cfg.batch)
|
||||
failed: list[str] = []
|
||||
for i, chunk in enumerate(chunks):
|
||||
sleep_between_batches(i, cfg.rpm)
|
||||
rows = self._request_rows(cfg, symbols=chunk, start_time=start_time, end_time=end_time)
|
||||
try:
|
||||
rows = self._request_rows_retry(
|
||||
cfg, chunk, start_time=start_time, end_time=end_time
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
# 单批失败只隔离该批 (#226): 之前任一批 502 会让整个 stage
|
||||
# 抛异常, 已成功批次的结果留在内存里全部丢弃
|
||||
failed.extend(chunk)
|
||||
logger.warning(
|
||||
"custom daily: batch %d/%d failed (%d symbols), skipped: %s",
|
||||
i + 1, len(chunks), len(chunk), e,
|
||||
)
|
||||
if on_chunk_done:
|
||||
on_chunk_done(i + 1, len(chunks))
|
||||
continue
|
||||
df = self._mapped_frame(cfg, rows)
|
||||
df = normalize_daily(df, source=self.name)
|
||||
if not df.is_empty():
|
||||
frames.append(df)
|
||||
if on_chunk_done:
|
||||
on_chunk_done(i + 1, len(chunks))
|
||||
if failed:
|
||||
logger.warning(
|
||||
"custom daily: %d/%d symbols missing due to batch failures: %s",
|
||||
len(failed), len(symbols), ", ".join(failed[:20]),
|
||||
)
|
||||
return pl.concat(frames, how="diagonal_relaxed") if frames else pl.DataFrame()
|
||||
|
||||
def get_adj_factors(
|
||||
@@ -165,15 +203,33 @@ class GenericHTTPProvider:
|
||||
cfg = self._dataset("adj_factor")
|
||||
frames: list[pl.DataFrame] = []
|
||||
chunks = chunked(symbols, cfg.batch)
|
||||
failed: list[str] = []
|
||||
for i, chunk in enumerate(chunks):
|
||||
sleep_between_batches(i, cfg.rpm)
|
||||
rows = self._request_rows(cfg, symbols=chunk, start_time=start_time, end_time=end_time)
|
||||
try:
|
||||
rows = self._request_rows_retry(
|
||||
cfg, chunk, start_time=start_time, end_time=end_time
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
failed.extend(chunk)
|
||||
logger.warning(
|
||||
"custom adj_factor: batch %d/%d failed (%d symbols), skipped: %s",
|
||||
i + 1, len(chunks), len(chunk), e,
|
||||
)
|
||||
if on_chunk_done:
|
||||
on_chunk_done(i + 1, len(chunks))
|
||||
continue
|
||||
df = self._mapped_frame(cfg, rows)
|
||||
df = normalize_adj_factors(df, source=self.name)
|
||||
if not df.is_empty():
|
||||
frames.append(df)
|
||||
if on_chunk_done:
|
||||
on_chunk_done(i + 1, len(chunks))
|
||||
if failed:
|
||||
logger.warning(
|
||||
"custom adj_factor: %d/%d symbols missing due to batch failures: %s",
|
||||
len(failed), len(symbols), ", ".join(failed[:20]),
|
||||
)
|
||||
return pl.concat(frames, how="diagonal_relaxed") if frames else pl.DataFrame()
|
||||
|
||||
def get_realtime(self) -> list[dict]:
|
||||
@@ -293,12 +349,18 @@ class GenericHTTPProvider:
|
||||
return pl.DataFrame()
|
||||
return pl.concat(frames, how="diagonal_relaxed")
|
||||
|
||||
@staticmethod
|
||||
def _normalize_minute(df: pl.DataFrame) -> pl.DataFrame:
|
||||
@classmethod
|
||||
def _normalize_minute(cls, df: pl.DataFrame) -> pl.DataFrame:
|
||||
"""把映射后的 df 规范成 minute canonical 列。"""
|
||||
if df.is_empty():
|
||||
return df
|
||||
if "datetime" in df.columns and df.schema["datetime"] != pl.Datetime("us"):
|
||||
if df.schema["datetime"] == pl.Utf8:
|
||||
# 字符串 datetime 直接 cast 会整体置 null (polars 不做字符串解析);
|
||||
# 先解析再对齐微秒精度 (#225, 参照
|
||||
# kline_sync._enforce_minute_beijing_wallclock 的处理)。
|
||||
# Series 级立即解析: 表达式错误要到 collect 才抛, 无法按格式回退
|
||||
df = df.with_columns(cls._parse_datetime_series(df["datetime"]))
|
||||
df = df.with_columns(pl.col("datetime").cast(pl.Datetime("us"), strict=False))
|
||||
for col in ("open", "high", "low", "close", "volume", "amount"):
|
||||
if col in df.columns:
|
||||
@@ -306,6 +368,27 @@ class GenericHTTPProvider:
|
||||
keep = [c for c in ("symbol", "datetime", "open", "high", "low", "close", "volume", "amount") if c in df.columns]
|
||||
return df.select(keep) if keep else pl.DataFrame()
|
||||
|
||||
_DATETIME_STR_FORMATS = (
|
||||
None, # 自动推断
|
||||
"%Y-%m-%d %H:%M:%S",
|
||||
"%Y-%m-%dT%H:%M:%S",
|
||||
"%Y/%m/%d %H:%M:%S",
|
||||
"%Y-%m-%d %H:%M",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _parse_datetime_series(cls, s: pl.Series) -> pl.Series:
|
||||
"""逐格式尝试解析字符串 datetime; 均失败返回全 null (宽松语义)。"""
|
||||
for fmt in cls._DATETIME_STR_FORMATS:
|
||||
try:
|
||||
return (
|
||||
s.str.to_datetime(strict=False, format=fmt)
|
||||
if fmt else s.str.to_datetime(strict=False)
|
||||
)
|
||||
except Exception: # noqa: BLE001 — 该格式不适用, 换下一个
|
||||
continue
|
||||
return pl.Series("datetime", [None] * s.len(), dtype=pl.Datetime("us"))
|
||||
|
||||
def test_dataset(self, dataset: str, symbols: list[str] | None = None) -> dict:
|
||||
cfg = self._dataset(dataset)
|
||||
test_symbols = symbols or ["000001.SZ"]
|
||||
|
||||
@@ -24,6 +24,7 @@ class TickFlowProvider:
|
||||
adj_factor=True,
|
||||
minute=True,
|
||||
realtime=True,
|
||||
depth5=True,
|
||||
financial=True,
|
||||
)
|
||||
|
||||
@@ -119,3 +120,9 @@ class TickFlowProvider:
|
||||
else:
|
||||
return pl.DataFrame()
|
||||
return pl.DataFrame(resp or [])
|
||||
|
||||
def get_depth_batch(self, symbols: list[str]) -> dict[str, dict]:
|
||||
if not symbols:
|
||||
return {}
|
||||
data = get_client().depth.batch(symbols)
|
||||
return data if isinstance(data, dict) else {}
|
||||
|
||||
@@ -154,6 +154,45 @@ def _ready_payload(generation: str) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _is_ready_payload(payload: dict[str, Any]) -> bool:
|
||||
generation = payload.get("generation")
|
||||
return (
|
||||
payload.get("state", "ready") == "ready"
|
||||
and isinstance(generation, str)
|
||||
and bool(generation)
|
||||
)
|
||||
|
||||
|
||||
def _publication_claim_is_running(payload: dict[str, Any]) -> bool:
|
||||
"""标记指向的发布是否仍在推进: 进程内活跃对象存在, 或属主进程仍存活。
|
||||
|
||||
owner_pid 等于当前进程但无活跃对象视为可接管 (同进程上一次尝试的遗留),
|
||||
与写入方 recover 接管的判定一致。
|
||||
"""
|
||||
if _ACTIVE_PUBLICATIONS.get(str(payload.get("publication_id"))) is not None:
|
||||
return True
|
||||
owner_pid = payload.get("owner_pid")
|
||||
return owner_pid != os.getpid() and _process_is_alive(owner_pid)
|
||||
|
||||
|
||||
def _orphaned_publishing_claim(payload: dict[str, Any]) -> bool:
|
||||
"""标记是否指向确定已死的发布: 属主是其他进程且已退出。
|
||||
|
||||
owner_pid 等于当前进程但无活跃对象时保守不判孤儿 —— 同进程异常遗留的
|
||||
publishing 标记意味着磁盘可能处于部分修改状态 (如清库删了一半), 读取方
|
||||
恢复 ready 会放行读取半修改数据; 必须由下一个写入方接管重发布。
|
||||
"""
|
||||
if _ACTIVE_PUBLICATIONS.get(str(payload.get("publication_id"))) is not None:
|
||||
return False
|
||||
owner_pid = payload.get("owner_pid")
|
||||
return (
|
||||
isinstance(owner_pid, int)
|
||||
and owner_pid > 0
|
||||
and owner_pid != os.getpid()
|
||||
and not _process_is_alive(owner_pid)
|
||||
)
|
||||
|
||||
|
||||
def get_enriched_generation(
|
||||
data_dir: Path,
|
||||
asset_type: str = "stock",
|
||||
@@ -167,19 +206,33 @@ def get_enriched_generation(
|
||||
raise EnrichedGenerationUnavailableError(
|
||||
"enriched data generation marker is unavailable"
|
||||
)
|
||||
with _exclusive_generation_lock(data_dir, asset_type):
|
||||
payload = _read_marker(path)
|
||||
if payload is None:
|
||||
generation = uuid.uuid4().hex
|
||||
_write_marker(path, _ready_payload(generation))
|
||||
return generation
|
||||
state = payload.get("state", "ready")
|
||||
generation = payload.get("generation")
|
||||
if state != "ready" or not isinstance(generation, str) or not generation:
|
||||
elif _is_ready_payload(payload):
|
||||
return payload["generation"]
|
||||
elif not _orphaned_publishing_claim(payload):
|
||||
# 发布仍在推进, 或为同进程异常遗留 (无法证明属主已死): 读取保持 fail-closed。
|
||||
raise EnrichedGenerationUnavailableError(
|
||||
"enriched data is being published; retry after the update finishes"
|
||||
)
|
||||
return generation
|
||||
# 指向已死发布的僵死标记: 在独占锁内二次确认后恢复 ready。
|
||||
with _exclusive_generation_lock(data_dir, asset_type):
|
||||
payload = _read_marker(path)
|
||||
if payload is None:
|
||||
generation = uuid.uuid4().hex
|
||||
_write_marker(path, _ready_payload(generation))
|
||||
return generation
|
||||
if _is_ready_payload(payload):
|
||||
return payload["generation"]
|
||||
if not _orphaned_publishing_claim(payload):
|
||||
raise EnrichedGenerationUnavailableError(
|
||||
"enriched data is being published; retry after the update finishes"
|
||||
)
|
||||
# 属主已死的 publishing 标记永远不会 commit, 读取方持续失败直到某个
|
||||
# 写入方碰巧接管 (dev 热重载杀掉发布进程即产生这种孤儿)。恢复为 ready
|
||||
# 并换新 generation: 磁盘可能残留部分替换的文件, 新 generation 让按代
|
||||
# 缓存全部失效, 避免把混合状态混入旧快照 —— 与写入方 recover 接管同语义。
|
||||
generation = uuid.uuid4().hex
|
||||
_write_marker(path, _ready_payload(generation))
|
||||
return generation
|
||||
|
||||
|
||||
def enriched_publication_incomplete(
|
||||
@@ -296,12 +349,7 @@ class EnrichedPublication:
|
||||
return
|
||||
_ACTIVE_PUBLICATIONS[self._publication_id] = self
|
||||
if current is not None and current.get("state", "ready") != "ready":
|
||||
current_id = current.get("publication_id")
|
||||
current_owner = _ACTIVE_PUBLICATIONS.get(str(current_id))
|
||||
owner_pid = current.get("owner_pid")
|
||||
if current_owner is not None or (
|
||||
owner_pid != os.getpid() and _process_is_alive(owner_pid)
|
||||
):
|
||||
if _publication_claim_is_running(current):
|
||||
raise EnrichedGenerationUnavailableError(
|
||||
"another enriched publication is active"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,749 @@
|
||||
"""因子公式 DSL 编译器 (P2)。
|
||||
|
||||
流水线: text → tokenizer → 递归下降解析(EBNF 见设计文档 §3.4) → AST → 语义检查
|
||||
→ 依赖/预热推导 → Polars Expr。编译失败返回结构化错误 (E001-E016), 不抛裸异常。
|
||||
|
||||
窗口纪律 (Polars 嵌套窗口会静默产出全 null, 必须在编译期杜绝):
|
||||
- 所有 ts_* 算子只向后看 (负 shift 常量层强制 E005)。
|
||||
- 时序子树仅在离开时序上下文时挂一次 over("symbol"); 截面算子挂 over("date")。
|
||||
- 截面算子消费含窗口的子树时, 编译为两阶段: 先把该子树物化为临时列 (单层 over),
|
||||
再对临时列做截面运算 —— frame_transform 负责按依赖顺序执行全部阶段。
|
||||
- 截面算子嵌在时序窗口内 (如 ts_mean(rank(x), n)) v1 不支持, 编译期 E009 拒绝。
|
||||
- 引用的注册因子(含 virtual)不内联表达式: 调用方用 materialize_scoring_columns
|
||||
物化成列, 编译产物统一以 pl.col(name) 引用; 运行期缺列即 fail-closed。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from functools import lru_cache
|
||||
from typing import Any
|
||||
|
||||
import polars as pl
|
||||
|
||||
from app.factors.registry import factor_dependencies, get_factor
|
||||
|
||||
FACTOR_COLUMN = "__dsl_factor__"
|
||||
|
||||
# 基准列 (设计文档 §3.1); 指标列 = 注册表 base 因子, 已注册因子 id 经注册表解析。
|
||||
BASE_COLUMNS: frozenset[str] = frozenset({
|
||||
"open", "high", "low", "close", "volume", "amount",
|
||||
"turnover_rate", "prev_close", "raw_close",
|
||||
})
|
||||
|
||||
MAX_AST_DEPTH = 12
|
||||
MAX_TOKENS = 200
|
||||
WINDOW_MIN, WINDOW_MAX = 2, 512
|
||||
DELAY_MAX = 512
|
||||
POWER_ABS_MAX = 4.0
|
||||
WINSORIZE_K_RANGE = (1.0, 6.0)
|
||||
|
||||
# 算子表: 名 -> (表达式参数个数, 常量参数名元组); 常量参数必须是数字字面量 (E003)。
|
||||
OPERATORS: dict[str, tuple[int, tuple[str, ...]]] = {
|
||||
"ts_mean": (1, ("n",)),
|
||||
"ts_std": (1, ("n",)),
|
||||
"ts_sum": (1, ("n",)),
|
||||
"ts_max": (1, ("n",)),
|
||||
"ts_min": (1, ("n",)),
|
||||
"ts_delay": (1, ("n",)),
|
||||
"ts_delta": (1, ("n",)),
|
||||
"ts_rank": (1, ("n",)),
|
||||
"ts_zscore": (1, ("n",)),
|
||||
"ts_corr": (2, ("n",)),
|
||||
"ts_cov": (2, ("n",)),
|
||||
"ts_quantile": (1, ("n", "q")),
|
||||
"decay_linear": (1, ("n",)),
|
||||
"rank": (1, ()),
|
||||
"zscore": (1, ()),
|
||||
"winsorize": (1, ("k",)), # k 可省略, 默认 3
|
||||
"power": (1, ("c",)),
|
||||
"clamp": (1, ("lo", "hi")),
|
||||
"if_else": (3, ()),
|
||||
"min": (2, ()),
|
||||
"max": (2, ()),
|
||||
"log": (1, ()),
|
||||
"abs": (1, ()),
|
||||
"sign": (1, ()),
|
||||
"sqrt": (1, ()),
|
||||
}
|
||||
TS_OPERATORS = frozenset({
|
||||
"ts_mean", "ts_std", "ts_sum", "ts_max", "ts_min", "ts_delay", "ts_delta",
|
||||
"ts_rank", "ts_zscore", "ts_corr", "ts_cov", "ts_quantile", "decay_linear",
|
||||
})
|
||||
CROSS_OPERATORS = frozenset({"rank", "zscore", "winsorize"})
|
||||
|
||||
|
||||
@dataclass
|
||||
class DslError:
|
||||
code: str
|
||||
message: str
|
||||
offset: int = 0
|
||||
detail: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"code": self.code,
|
||||
"message": self.message,
|
||||
"position": {"offset": self.offset, "line": 1},
|
||||
"detail": self.detail,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class CompiledFormula:
|
||||
ok: bool
|
||||
errors: list[DslError] = field(default_factory=list)
|
||||
frame_transform: Any | None = None # (frame: pl.DataFrame) -> pl.DataFrame | None (缺列 None = E013)
|
||||
dependencies: frozenset[str] = frozenset() # 展开到 enriched base 列
|
||||
referenced_factors: frozenset[str] = frozenset() # 引用的注册因子 id (含 virtual, 需物化)
|
||||
warmup_bars: int = 1
|
||||
cross_sectional: bool = False
|
||||
formula_text: str = ""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- tokenizer
|
||||
|
||||
_TOKEN_RE = re.compile(
|
||||
r"\s*(?:(?P<num>\d+(?:\.\d+)?)|(?P<ident>[A-Za-z_][A-Za-z0-9_]*)|(?P<op>>=|<=|==|!=|[+\-*/><(),]))"
|
||||
)
|
||||
_KEYWORDS = frozenset({"and", "or", "not"})
|
||||
|
||||
|
||||
def _tokenize(text: str) -> tuple[list[tuple[str, Any, int]], DslError | None]:
|
||||
tokens: list[tuple[str, Any, int]] = []
|
||||
pos = 0
|
||||
while pos < len(text):
|
||||
match = _TOKEN_RE.match(text, pos)
|
||||
if match is None or match.end() == pos:
|
||||
rest = text[pos:].strip()
|
||||
if not rest:
|
||||
break
|
||||
return [], DslError("E014", f"语法错误: 无法识别的字符 '{rest[0]}'", offset=pos)
|
||||
if match.group("num") is not None:
|
||||
tokens.append(("num", float(match.group("num")), match.start("num")))
|
||||
elif match.group("ident") is not None:
|
||||
tokens.append(("ident", match.group("ident"), match.start("ident")))
|
||||
else:
|
||||
tokens.append(("op", match.group("op"), match.start("op")))
|
||||
pos = match.end()
|
||||
return tokens, None
|
||||
|
||||
|
||||
# ------------------------------------------------------------------- parser
|
||||
# AST 节点: dict(kind, value, children, offset[, _constants])
|
||||
|
||||
|
||||
class _Parser:
|
||||
_CMP = frozenset({">", ">=", "<", "<=", "==", "!="})
|
||||
|
||||
def __init__(self, tokens: list[tuple[str, Any, int]], text: str) -> None:
|
||||
self.tokens = tokens
|
||||
self.text = text
|
||||
self.index = 0
|
||||
|
||||
def _peek(self) -> tuple[str, Any, int] | None:
|
||||
return self.tokens[self.index] if self.index < len(self.tokens) else None
|
||||
|
||||
def _next(self) -> tuple[str, Any, int]:
|
||||
token = self.tokens[self.index]
|
||||
self.index += 1
|
||||
return token
|
||||
|
||||
def parse(self) -> tuple[dict | None, DslError | None]:
|
||||
if not self.tokens:
|
||||
return None, DslError("E014", "语法错误: 表达式为空", offset=0)
|
||||
node, error = self._or_expr()
|
||||
if error:
|
||||
return None, error
|
||||
if self._peek() is not None:
|
||||
_, value, offset = self._peek()
|
||||
return None, DslError("E014", f"语法错误: 多余的记号 '{value}'", offset=offset)
|
||||
return node, None
|
||||
|
||||
def _or_expr(self):
|
||||
left, error = self._and_expr()
|
||||
if error:
|
||||
return None, error
|
||||
while (token := self._peek()) and token[0] == "ident" and token[1] == "or":
|
||||
self._next()
|
||||
right, error = self._and_expr()
|
||||
if error:
|
||||
return None, error
|
||||
left = {"kind": "bin", "value": "or", "children": [left, right], "offset": token[2]}
|
||||
return left, None
|
||||
|
||||
def _and_expr(self):
|
||||
left, error = self._cmp_expr()
|
||||
if error:
|
||||
return None, error
|
||||
while (token := self._peek()) and token[0] == "ident" and token[1] == "and":
|
||||
self._next()
|
||||
right, error = self._cmp_expr()
|
||||
if error:
|
||||
return None, error
|
||||
left = {"kind": "bin", "value": "and", "children": [left, right], "offset": token[2]}
|
||||
return left, None
|
||||
|
||||
def _cmp_expr(self):
|
||||
left, error = self._add_expr()
|
||||
if error:
|
||||
return None, error
|
||||
while (token := self._peek()) and token[0] == "op" and token[1] in self._CMP:
|
||||
self._next()
|
||||
right, error = self._add_expr()
|
||||
if error:
|
||||
return None, error
|
||||
left = {"kind": "bin", "value": token[1], "children": [left, right], "offset": token[2]}
|
||||
return left, None
|
||||
|
||||
def _add_expr(self):
|
||||
left, error = self._mul_expr()
|
||||
if error:
|
||||
return None, error
|
||||
while (token := self._peek()) and token[0] == "op" and token[1] in ("+", "-"):
|
||||
self._next()
|
||||
right, error = self._mul_expr()
|
||||
if error:
|
||||
return None, error
|
||||
left = {"kind": "bin", "value": token[1], "children": [left, right], "offset": token[2]}
|
||||
return left, None
|
||||
|
||||
def _mul_expr(self):
|
||||
left, error = self._unary()
|
||||
if error:
|
||||
return None, error
|
||||
while (token := self._peek()) and token[0] == "op" and token[1] in ("*", "/"):
|
||||
self._next()
|
||||
right, error = self._unary()
|
||||
if error:
|
||||
return None, error
|
||||
left = {"kind": "bin", "value": token[1], "children": [left, right], "offset": token[2]}
|
||||
return left, None
|
||||
|
||||
def _unary(self):
|
||||
token = self._peek()
|
||||
if token and token[0] == "op" and token[1] == "-":
|
||||
self._next()
|
||||
operand, error = self._unary()
|
||||
if error:
|
||||
return None, error
|
||||
return {"kind": "unary", "value": "-", "children": [operand], "offset": token[2]}, None
|
||||
return self._primary()
|
||||
|
||||
def _primary(self):
|
||||
token = self._peek()
|
||||
if token is None:
|
||||
return None, DslError("E014", "语法错误: 表达式意外结束", offset=len(self.text))
|
||||
kind, value, offset = self._next()
|
||||
if kind == "num":
|
||||
return {"kind": "num", "value": value, "children": [], "offset": offset}, None
|
||||
if kind == "ident":
|
||||
if value in _KEYWORDS:
|
||||
return None, DslError("E014", f"语法错误: 关键字 '{value}' 不能作为操作数", offset=offset)
|
||||
nxt = self._peek()
|
||||
if nxt and nxt[0] == "op" and nxt[1] == "(":
|
||||
return self._call(value, offset)
|
||||
return {"kind": "col", "value": value, "children": [], "offset": offset}, None
|
||||
if kind == "op" and value == "(":
|
||||
inner, error = self._or_expr()
|
||||
if error:
|
||||
return None, error
|
||||
closing = self._peek()
|
||||
if not (closing and closing[0] == "op" and closing[1] == ")"):
|
||||
return None, DslError("E014", "语法错误: 缺少右括号 ')'", offset=offset)
|
||||
self._next()
|
||||
return inner, None
|
||||
return None, DslError("E014", f"语法错误: 意外的记号 '{value}'", offset=offset)
|
||||
|
||||
def _call(self, name: str, offset: int):
|
||||
self._next() # consume '('
|
||||
args: list[dict] = []
|
||||
token = self._peek()
|
||||
if not (token and token[0] == "op" and token[1] == ")"):
|
||||
while True:
|
||||
arg, error = self._or_expr()
|
||||
if error:
|
||||
return None, error
|
||||
args.append(arg)
|
||||
token = self._peek()
|
||||
if token and token[0] == "op" and token[1] == ",":
|
||||
self._next()
|
||||
continue
|
||||
break
|
||||
closing = self._peek()
|
||||
if not (closing and closing[0] == "op" and closing[1] == ")"):
|
||||
return None, DslError("E014", f"语法错误: 函数 '{name}' 缺少右括号", offset=offset)
|
||||
self._next()
|
||||
return {"kind": "call", "value": name, "children": args, "offset": offset}, None
|
||||
|
||||
|
||||
# ---------------------------------------------------------- semantic checks
|
||||
|
||||
|
||||
def _ast_depth(node: dict) -> int:
|
||||
if not node["children"]:
|
||||
return 1
|
||||
return 1 + max(_ast_depth(child) for child in node["children"])
|
||||
|
||||
|
||||
def _collect_identifiers(node: dict, found: set[str]) -> None:
|
||||
if node["kind"] == "col":
|
||||
found.add(node["value"])
|
||||
for child in node["children"]:
|
||||
_collect_identifiers(child, found)
|
||||
|
||||
|
||||
def _const_value(node: dict) -> float | None:
|
||||
if node["kind"] == "num":
|
||||
return float(node["value"])
|
||||
if node["kind"] == "unary" and node["value"] == "-" and node["children"][0]["kind"] == "num":
|
||||
return -float(node["children"][0]["value"])
|
||||
return None
|
||||
|
||||
|
||||
def _check_call(node: dict, errors: list[DslError]) -> dict[str, float]:
|
||||
"""检查函数签名与常量参数范围; 返回解析出的常量参数表。"""
|
||||
name = node["value"]
|
||||
args = node["children"]
|
||||
if name not in OPERATORS:
|
||||
errors.append(DslError("E002", f"未知函数: {name}", offset=node["offset"], detail={"name": name}))
|
||||
return {}
|
||||
n_expr, const_names = OPERATORS[name]
|
||||
has_optional_k = name == "winsorize"
|
||||
total_min, total_max = n_expr + (0 if has_optional_k else len(const_names)), n_expr + len(const_names)
|
||||
if not (total_min <= len(args) <= total_max):
|
||||
errors.append(DslError(
|
||||
"E003", f"函数 {name} 参数数量不符: 期望 {total_min}~{total_max} 个, 实际 {len(args)}",
|
||||
offset=node["offset"], detail={"name": name, "args": len(args)},
|
||||
))
|
||||
return {}
|
||||
constants: dict[str, float] = {}
|
||||
for index, const_name in enumerate(const_names):
|
||||
arg = args[n_expr + index]
|
||||
value = _const_value(arg)
|
||||
if value is None:
|
||||
errors.append(DslError(
|
||||
"E003", f"函数 {name} 的参数 {const_name} 必须是数字常量",
|
||||
offset=arg["offset"], detail={"name": name, "param": const_name},
|
||||
))
|
||||
continue
|
||||
constants[const_name] = value
|
||||
if "n" in constants:
|
||||
n_value = constants["n"]
|
||||
if n_value != int(n_value):
|
||||
errors.append(DslError("E004", "窗口参数必须是整数", offset=node["offset"], detail={"n": n_value}))
|
||||
else:
|
||||
n_int = int(n_value)
|
||||
if n_int < 0 and name in ("ts_delay", "ts_delta"):
|
||||
errors.append(DslError(
|
||||
"E005", f"负 shift: {name} 的 n 必须 ≥ 0 (负数即未来函数)",
|
||||
offset=node["offset"], detail={"n": n_int},
|
||||
))
|
||||
elif name == "ts_delay" and not (1 <= n_int <= DELAY_MAX):
|
||||
errors.append(DslError("E004", f"ts_delay 的 n 必须在 [1,{DELAY_MAX}] 内", offset=node["offset"], detail={"n": n_int}))
|
||||
elif name == "ts_delta" and not (0 <= n_int <= DELAY_MAX):
|
||||
errors.append(DslError("E004", f"ts_delta 的 n 必须在 [0,{DELAY_MAX}] 内", offset=node["offset"], detail={"n": n_int}))
|
||||
elif name not in ("ts_delay", "ts_delta") and not (WINDOW_MIN <= n_int <= WINDOW_MAX):
|
||||
errors.append(DslError(
|
||||
"E004", f"窗口 n 必须在 [{WINDOW_MIN},{WINDOW_MAX}] 内", offset=node["offset"], detail={"n": n_int},
|
||||
))
|
||||
if "q" in constants and not (0.0 < constants["q"] < 1.0):
|
||||
errors.append(DslError("E004", "ts_quantile 的 q 必须在 (0,1) 开区间内", offset=node["offset"], detail={"q": constants["q"]}))
|
||||
if "c" in constants and abs(constants["c"]) > POWER_ABS_MAX:
|
||||
errors.append(DslError("E010", f"power 指数 |c| ≤ {POWER_ABS_MAX}", offset=node["offset"], detail={"c": constants["c"]}))
|
||||
if "k" in constants and not (WINSORIZE_K_RANGE[0] <= constants["k"] <= WINSORIZE_K_RANGE[1]):
|
||||
errors.append(DslError("E011", "winsorize 的 k 必须在 [1,6] 内", offset=node["offset"], detail={"k": constants["k"]}))
|
||||
if "lo" in constants and "hi" in constants and constants["lo"] > constants["hi"]:
|
||||
errors.append(DslError("E003", "clamp 的 lo 不能大于 hi", offset=node["offset"]))
|
||||
return constants
|
||||
|
||||
|
||||
def _semantic_walk(node: dict, errors: list[DslError], constants_by_call: dict[int, dict]) -> None:
|
||||
if node["kind"] == "call":
|
||||
constants_by_call[id(node)] = _check_call(node, errors)
|
||||
for child in node["children"]:
|
||||
_semantic_walk(child, errors, constants_by_call)
|
||||
return
|
||||
if node["kind"] == "bin" and node["value"] == "/":
|
||||
right = node["children"][1]
|
||||
if _const_value(right) == 0:
|
||||
errors.append(DslError("E008", "静态除零: 分母为常量 0", offset=right["offset"]))
|
||||
for child in node["children"]:
|
||||
_semantic_walk(child, errors, constants_by_call)
|
||||
|
||||
|
||||
# ------------------------------------------------------------- code generation
|
||||
|
||||
_CMP_METHOD = {">": "gt", ">=": "ge", "<": "lt", "<=": "le", "==": "eq", "!=": "ne"}
|
||||
|
||||
|
||||
def _safe_div(numerator: pl.Expr, denominator: pl.Expr) -> pl.Expr:
|
||||
return (
|
||||
pl.when(denominator.is_not_null() & (denominator != 0))
|
||||
.then(numerator / denominator)
|
||||
.otherwise(None)
|
||||
)
|
||||
|
||||
|
||||
def _rolling_apply(inner: pl.Expr, op: str, n: int, extra: dict[str, float]) -> pl.Expr:
|
||||
"""对无 over 的内层序列应用窗口逻辑; 返回值同样不挂 over。"""
|
||||
if op == "ts_mean":
|
||||
return inner.rolling_mean(n, min_samples=n)
|
||||
if op == "ts_std":
|
||||
return inner.rolling_std(n, min_samples=n)
|
||||
if op == "ts_sum":
|
||||
return inner.rolling_sum(n, min_samples=n)
|
||||
if op == "ts_max":
|
||||
return inner.rolling_max(n, min_samples=n)
|
||||
if op == "ts_min":
|
||||
return inner.rolling_min(n, min_samples=n)
|
||||
if op == "ts_delay":
|
||||
return inner.shift(n)
|
||||
if op == "ts_delta":
|
||||
return inner - inner.shift(n)
|
||||
if op == "ts_rank":
|
||||
return inner.rolling_rank(n, min_samples=n)
|
||||
if op == "ts_zscore":
|
||||
mean = inner.rolling_mean(n, min_samples=n)
|
||||
std = inner.rolling_std(n, min_samples=n)
|
||||
return pl.when(std > 0).then((inner - mean) / std).otherwise(None)
|
||||
if op == "ts_quantile":
|
||||
return inner.rolling_quantile(extra.get("q", 0.5), window_size=n, min_samples=n)
|
||||
if op == "decay_linear":
|
||||
# 近端权重大: 权重 n, n-1, ..., 1, 总权 n(n+1)/2
|
||||
weighted = None
|
||||
for i in range(n):
|
||||
term = (n - i) * inner.shift(i)
|
||||
weighted = term if weighted is None else weighted + term
|
||||
assert weighted is not None
|
||||
return _safe_div(weighted, pl.lit(float(n * (n + 1) / 2)))
|
||||
raise AssertionError(op)
|
||||
|
||||
|
||||
def _compile_node(node: dict) -> tuple[pl.Expr | None, bool, bool]:
|
||||
"""返回 (expr, needs_symbol_window, is_bool)。
|
||||
|
||||
needs_symbol_window=True 表示该子树含 ts 窗口逻辑但尚未挂 over;
|
||||
由非时序上下文的调用方挂 over("symbol"), 时序上下文继续向内传递。
|
||||
"""
|
||||
kind = node["kind"]
|
||||
if kind == "num":
|
||||
return pl.lit(node["value"]), False, False
|
||||
if kind == "col":
|
||||
# 基准列/base 因子/虚拟因子统一以列引用; 虚拟因子由调用方物化 (运行期缺列 fail-closed)
|
||||
return pl.col(node["value"]), False, False
|
||||
if kind == "unary":
|
||||
operand, needs_window, _ = _compile_node(node["children"][0])
|
||||
if operand is None:
|
||||
return None, False, False
|
||||
return -operand, needs_window, False
|
||||
if kind == "bin":
|
||||
op = node["value"]
|
||||
left, left_window, _ = _compile_node(node["children"][0])
|
||||
right, right_window, _ = _compile_node(node["children"][1])
|
||||
if left is None or right is None:
|
||||
return None, False, False
|
||||
if left_window:
|
||||
left = left.over("symbol")
|
||||
if right_window:
|
||||
right = right.over("symbol")
|
||||
if op == "+":
|
||||
return left + right, False, False
|
||||
if op == "-":
|
||||
return left - right, False, False
|
||||
if op == "*":
|
||||
return left * right, False, False
|
||||
if op == "/":
|
||||
return _safe_div(left, right), False, False
|
||||
if op in _CMP_METHOD:
|
||||
return getattr(left, _CMP_METHOD[op])(right), False, True
|
||||
if op == "and":
|
||||
return left & right, False, True
|
||||
if op == "or":
|
||||
return left | right, False, True
|
||||
return None, False, False
|
||||
if kind == "call":
|
||||
return _compile_call(node)
|
||||
return None, False, False
|
||||
|
||||
|
||||
def _compile_call(node: dict) -> tuple[pl.Expr | None, bool, bool]:
|
||||
name = node["value"]
|
||||
children = node["children"]
|
||||
constants: dict[str, float] = node.get("_constants", {})
|
||||
n_expr, _ = OPERATORS[name]
|
||||
|
||||
if name in TS_OPERATORS:
|
||||
inner, _, _ = _compile_node(children[0])
|
||||
if inner is None:
|
||||
return None, False, False
|
||||
if name in ("ts_corr", "ts_cov"):
|
||||
second, _, _ = _compile_node(children[1])
|
||||
if second is None:
|
||||
return None, False, False
|
||||
n = int(constants.get("n", 0))
|
||||
expr = (
|
||||
pl.rolling_corr(inner, second, window_size=n)
|
||||
if name == "ts_corr"
|
||||
else pl.rolling_cov(inner, second, window_size=n)
|
||||
)
|
||||
return expr, True, False
|
||||
expr = _rolling_apply(inner, name, int(constants.get("n", 0)), constants)
|
||||
return expr, True, False
|
||||
|
||||
if name in CROSS_OPERATORS:
|
||||
inner, inner_window, _ = _compile_node(children[0])
|
||||
if inner is None:
|
||||
return None, False, False
|
||||
if inner_window:
|
||||
inner = inner.over("symbol")
|
||||
if name == "rank":
|
||||
count = inner.count().over("date")
|
||||
return inner.rank(method="average").over("date") / count, False, False
|
||||
if name == "zscore":
|
||||
mean = inner.mean().over("date")
|
||||
std = inner.std().over("date")
|
||||
return pl.when(std > 0).then((inner - mean) / std).otherwise(None), False, False
|
||||
k = constants.get("k", 3.0)
|
||||
mean = inner.mean().over("date")
|
||||
std = inner.std().over("date")
|
||||
return inner.clip(mean - k * std, mean + k * std), False, False
|
||||
|
||||
if name == "if_else":
|
||||
cond, cond_window, _ = _compile_node(children[0])
|
||||
then_expr, then_window, _ = _compile_node(children[1])
|
||||
else_expr, else_window, _ = _compile_node(children[2])
|
||||
if cond is None or then_expr is None or else_expr is None:
|
||||
return None, False, False
|
||||
if cond_window:
|
||||
cond = cond.over("symbol")
|
||||
if then_window:
|
||||
then_expr = then_expr.over("symbol")
|
||||
if else_window:
|
||||
else_expr = else_expr.over("symbol")
|
||||
return pl.when(cond).then(then_expr).otherwise(else_expr), False, False
|
||||
|
||||
args: list[pl.Expr | None] = []
|
||||
arg_windows: list[bool] = []
|
||||
for index in range(n_expr):
|
||||
arg, arg_window, _ = _compile_node(children[index])
|
||||
args.append(arg)
|
||||
arg_windows.append(arg_window)
|
||||
if any(arg is None for arg in args):
|
||||
return None, False, False
|
||||
resolved: list[pl.Expr] = []
|
||||
for arg, arg_window in zip(args, arg_windows, strict=True):
|
||||
resolved.append(arg.over("symbol") if arg_window else arg)
|
||||
first = resolved[0]
|
||||
if name == "log":
|
||||
return pl.when(first > 0).then(first.log()).otherwise(None), False, False
|
||||
if name == "abs":
|
||||
return first.abs(), False, False
|
||||
if name == "sign":
|
||||
return first.sign(), False, False
|
||||
if name == "sqrt":
|
||||
return pl.when(first >= 0).then(first.sqrt()).otherwise(None), False, False
|
||||
if name == "power":
|
||||
return first.pow(constants.get("c", 1.0)), False, False
|
||||
if name == "clamp":
|
||||
return first.clip(constants.get("lo"), constants.get("hi")), False, False
|
||||
if name == "min":
|
||||
return pl.min_horizontal(*resolved), False, False
|
||||
if name == "max":
|
||||
return pl.max_horizontal(*resolved), False, False
|
||||
return None, False, False
|
||||
|
||||
|
||||
def compile_formula(text: str) -> CompiledFormula:
|
||||
"""编译公式文本; 永不抛异常, 失败以 errors 表达 (fail-closed)。"""
|
||||
if not isinstance(text, str) or not text.strip():
|
||||
return CompiledFormula(ok=False, errors=[DslError("E014", "语法错误: 表达式为空")], formula_text=text)
|
||||
|
||||
tokens, tokenize_error = _tokenize(text)
|
||||
errors: list[DslError] = [tokenize_error] if tokenize_error else []
|
||||
if len(tokens) > MAX_TOKENS:
|
||||
errors.append(DslError("E007", f"规模超限: token 数 {len(tokens)} > {MAX_TOKENS}"))
|
||||
if errors:
|
||||
return CompiledFormula(ok=False, errors=errors, formula_text=text)
|
||||
|
||||
ast, parse_error = _Parser(tokens, text).parse()
|
||||
if parse_error:
|
||||
return CompiledFormula(ok=False, errors=[parse_error], formula_text=text)
|
||||
|
||||
if _ast_depth(ast) > MAX_AST_DEPTH:
|
||||
errors.append(DslError("E006", f"嵌套深度超限: AST 深度 {_ast_depth(ast)} > {MAX_AST_DEPTH}"))
|
||||
|
||||
identifiers: set[str] = set()
|
||||
_collect_identifiers(ast, identifiers)
|
||||
if not identifiers:
|
||||
errors.append(DslError("E016", "常量表达式: 公式必须引用至少一个数据列或因子"))
|
||||
|
||||
for name in sorted(identifiers):
|
||||
if name not in BASE_COLUMNS and get_factor(name) is None:
|
||||
errors.append(DslError("E001", f"未知标识符: {name}", detail={"name": name}))
|
||||
|
||||
constants_by_call: dict[int, dict] = {}
|
||||
_semantic_walk(ast, errors, constants_by_call)
|
||||
|
||||
dependencies: set[str] = set()
|
||||
referenced_factors: set[str] = set()
|
||||
warmup = 1
|
||||
cross_sectional = False
|
||||
for name in identifiers:
|
||||
if name in BASE_COLUMNS:
|
||||
dependencies.add(name)
|
||||
continue
|
||||
spec = get_factor(name)
|
||||
if spec is None:
|
||||
continue
|
||||
referenced_factors.add(name)
|
||||
dependencies.update(factor_dependencies([name]))
|
||||
warmup = max(warmup, spec.warmup_bars)
|
||||
|
||||
for node_constants in constants_by_call.values():
|
||||
n_value = node_constants.get("n")
|
||||
if n_value is not None and n_value == int(n_value) and int(n_value) > 0:
|
||||
warmup = max(warmup, int(n_value) + 1)
|
||||
|
||||
def _find_cross(node: dict) -> None:
|
||||
nonlocal cross_sectional
|
||||
if node["kind"] == "call" and node["value"] in CROSS_OPERATORS:
|
||||
cross_sectional = True
|
||||
for child in node["children"]:
|
||||
_find_cross(child)
|
||||
|
||||
_find_cross(ast)
|
||||
|
||||
if errors:
|
||||
return CompiledFormula(
|
||||
ok=False, errors=errors, dependencies=frozenset(dependencies),
|
||||
referenced_factors=frozenset(referenced_factors),
|
||||
warmup_bars=warmup, cross_sectional=cross_sectional, formula_text=text,
|
||||
)
|
||||
|
||||
# 挂常量表必须在任何 deepcopy 之前 (deepcopy 携带 _constants; 事后按 id() 重挂会失联)
|
||||
def _attach(node: dict) -> None:
|
||||
if node["kind"] == "call":
|
||||
node["_constants"] = constants_by_call.get(id(node), {})
|
||||
for child in node["children"]:
|
||||
_attach(child)
|
||||
|
||||
_attach(ast)
|
||||
|
||||
# 阶段一: 校验并拒绝"截面算子嵌在时序窗口内" (无法单层 over 表达)
|
||||
def _contains_cross(node: dict) -> bool:
|
||||
if node["kind"] == "call" and node["value"] in CROSS_OPERATORS:
|
||||
return True
|
||||
return any(_contains_cross(child) for child in node["children"])
|
||||
|
||||
def _reject_cross_in_ts(node: dict) -> None:
|
||||
if node["kind"] == "call" and node["value"] in TS_OPERATORS:
|
||||
for child in node["children"]:
|
||||
if _contains_cross(child):
|
||||
errors.append(DslError(
|
||||
"E009",
|
||||
f"截面算子不能嵌在时序窗口内: {node['value']}(...) 的参数含 rank/zscore/winsorize",
|
||||
offset=node["offset"],
|
||||
))
|
||||
return
|
||||
for child in node["children"]:
|
||||
_reject_cross_in_ts(child)
|
||||
|
||||
_reject_cross_in_ts(ast)
|
||||
if errors:
|
||||
return CompiledFormula(
|
||||
ok=False, errors=errors, dependencies=frozenset(dependencies),
|
||||
referenced_factors=frozenset(referenced_factors),
|
||||
warmup_bars=warmup, cross_sectional=cross_sectional, formula_text=text,
|
||||
)
|
||||
|
||||
# 阶段二: 提取截面算子的含窗口子树为临时列 (Polars 嵌套窗口会静默全 null)
|
||||
# worklist 逐层下钻; temps 后进先出反转即依赖顺序 (深层先算)。
|
||||
def _needs_symbol_window(node: dict) -> bool:
|
||||
kind = node["kind"]
|
||||
if kind in ("num", "col"):
|
||||
return False
|
||||
if kind == "unary":
|
||||
return _needs_symbol_window(node["children"][0])
|
||||
if node["kind"] == "call" and node["value"] in TS_OPERATORS:
|
||||
return True
|
||||
return any(_needs_symbol_window(child) for child in node["children"])
|
||||
|
||||
def _has_any_over(node: dict) -> bool:
|
||||
# 含时序窗口 或 含截面算子(编译后自带 over("date")) 的子树都不能直接进截面上下文
|
||||
return _needs_symbol_window(node) or _contains_cross(node)
|
||||
|
||||
temp_roots: list[dict] = []
|
||||
pending: list[dict] = [ast]
|
||||
while pending:
|
||||
current = pending.pop(0)
|
||||
if current.get("kind") == "call" and current.get("value") in CROSS_OPERATORS:
|
||||
operand = current["children"][0]
|
||||
if _has_any_over(operand):
|
||||
alias = f"__tsfx_{len(temp_roots)}__"
|
||||
current["children"][0] = {"kind": "col", "value": alias, "children": [], "offset": operand["offset"]}
|
||||
temp_roots.append({"alias": alias, "root": copy.deepcopy(operand)})
|
||||
pending.append(temp_roots[-1]["root"])
|
||||
continue # 操作数已替换为临时列, 不再下钻原子树
|
||||
pending.extend(current.get("children", []))
|
||||
|
||||
# 阶段三: 编译最终表达式与临时列表达式 (按依赖顺序: 深层在前)
|
||||
# _constants 已在 deepcopy 前挂载并被复制携带, 不得按 id() 重挂 (复制后 id 失联)
|
||||
temp_exprs: list[pl.Expr] = []
|
||||
for item in reversed(temp_roots):
|
||||
root = copy.deepcopy(item["root"])
|
||||
expr, needs_window, _ = _compile_node(root)
|
||||
if expr is None:
|
||||
errors.append(DslError("E009", f"无法编译临时列: {item['alias']}"))
|
||||
continue
|
||||
if needs_window:
|
||||
expr = expr.over("symbol")
|
||||
temp_exprs.append(expr.alias(item["alias"]))
|
||||
|
||||
final_ast = copy.deepcopy(ast)
|
||||
compiled, needs_window, is_bool = _compile_node(final_ast)
|
||||
if compiled is None or errors:
|
||||
return CompiledFormula(
|
||||
ok=False,
|
||||
errors=errors or [DslError("E009", "产出类型非法: 无法编译为数值表达式")],
|
||||
dependencies=frozenset(dependencies),
|
||||
referenced_factors=frozenset(referenced_factors),
|
||||
warmup_bars=warmup, cross_sectional=cross_sectional, formula_text=text,
|
||||
)
|
||||
if needs_window:
|
||||
compiled = compiled.over("symbol")
|
||||
if is_bool:
|
||||
compiled = compiled.cast(pl.Float64)
|
||||
|
||||
# 运行期帧变换: 检查全部引用列 (基准依赖 + 引用因子) 存在, 否则 None (E013 fail-closed)
|
||||
required_columns = set(dependencies) | set(referenced_factors)
|
||||
staged_exprs = temp_exprs # 依赖顺序已排
|
||||
|
||||
def frame_transform(frame: pl.DataFrame) -> pl.DataFrame | None:
|
||||
if not required_columns.issubset(set(frame.columns)):
|
||||
return None
|
||||
result = frame
|
||||
if staged_exprs:
|
||||
result = result.with_columns(staged_exprs)
|
||||
return result.with_columns(compiled.alias(FACTOR_COLUMN))
|
||||
|
||||
return CompiledFormula(
|
||||
ok=True,
|
||||
errors=[],
|
||||
frame_transform=frame_transform,
|
||||
dependencies=frozenset(dependencies),
|
||||
referenced_factors=frozenset(referenced_factors),
|
||||
warmup_bars=warmup,
|
||||
cross_sectional=cross_sectional,
|
||||
formula_text=text,
|
||||
)
|
||||
|
||||
|
||||
@lru_cache(maxsize=256)
|
||||
def compile_formula_cached(text: str) -> CompiledFormula:
|
||||
"""带 LRU 缓存的编译入口 (公式文本 → 编译产物, 设计文档 §3.3)。
|
||||
|
||||
CompiledFormula 为不可变值对象 (frame_transform 闭包只读), 缓存共享安全。
|
||||
"""
|
||||
return compile_formula(text)
|
||||
@@ -0,0 +1,354 @@
|
||||
"""扩展表字段 → 因子/信号接入 (单一原语, 两个消费方)。
|
||||
|
||||
扩展数据 (data/ext_data/{config_id}) 的数值字段在 enriched 帧组装时 join 到帧上,
|
||||
并以 kind="base" (空依赖 = 已物化列自身) 注册进因子注册表:
|
||||
|
||||
- 自定义信号: custom_signals.allowed_fields() 并入注册表因子, 扩展列出现在
|
||||
信号条件字段下拉中 (all_factors → ensure_synced 惰性同步);
|
||||
- 因子/评分/检验: scoring_value_expr 对帧上已有列直接 pl.col 引用,
|
||||
注册表条目让扩展字段同时出现在因子库列表与 AI 提示词中。
|
||||
|
||||
口径与边界 (金融契约, 见 CONTRIBUTING §3/§5.3):
|
||||
- timeseries 模式: 按 (symbol, date) 分区日期精确对齐, 历史帧无未来函数;
|
||||
- snapshot 模式: 代表"最新值", 仅在单日帧 (compute_enriched_today 盘中/当日)
|
||||
注入; 多日历史帧跳过, 否则回测/历史回看会引入未来数据;
|
||||
- 数值字段 (int/float, 统一 Float64): 因子 + 信号双通道 (注册表 base 条目);
|
||||
- string 字段: 仅信号条件通道 (contains/==/!= 字符串运算符, 概念/行业归属
|
||||
筛选), 不注册为因子 —— 因子 IC/排序是数值口径; bool 不参与。
|
||||
|
||||
缓存与失效 (CONTRIBUTING §6.1):
|
||||
- 配置清单复用 ExtConfigStore.load_all 的目录签名缓存;
|
||||
- 已加载的扩展帧按 (config 目录/分区签名) 缓存, 数据/配置变更后由
|
||||
invalidate_ext_caches 清除 (写入端 write_ext_parquet / upsert / delete 自动调用),
|
||||
同时清策略结果缓存 —— 策略历史窗口与 enriched 内存缓存里的帧含旧扩展列。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
import polars as pl
|
||||
|
||||
from app.factors.registry import FactorSpec, get_factor, register_factor, unregister_factor
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
EXT_PREFIX = "ext_"
|
||||
_NUMERIC_DTYPES = frozenset({"int", "float"})
|
||||
# 信号通道支持的 dtype: 数值 (Float64) + 字符串 (Utf8, contains/==/!=)
|
||||
_SIGNAL_DTYPES = _NUMERIC_DTYPES | {"string"}
|
||||
|
||||
# 帧缓存: (data_dir, config_id, mode) -> (目录/分区签名, DataFrame)
|
||||
_frame_cache: dict[tuple[str, str, str], tuple[tuple, pl.DataFrame]] = {}
|
||||
# 注册同步状态: (data_dir, 配置签名); None/失配 → 下次调用重新同步。
|
||||
# 已注册集合以注册表为权威 (ext_ 前缀条目), 不单独记账 —— 失效入口清空
|
||||
# 状态后, 重新同步仍能从注册表注销已移除的扩展因子。
|
||||
_sync_state: tuple | None = None
|
||||
|
||||
|
||||
def ext_column_name(config_id: str, field_name: str) -> str:
|
||||
"""扩展字段在帧/信号中的列名: ext_{config_id}_{field}。
|
||||
|
||||
保留中日韩文字 (\w 含 unicode 字母) —— 预设表的字段名多为中文
|
||||
(所属概念/股票简称), 全部折叠为 ASCII 会互相碰撞。非单词字符转下划线。
|
||||
"""
|
||||
sanitized = re.sub(r"[^\w]+", "_", field_name, flags=re.UNICODE).strip("_") or "f"
|
||||
return f"{EXT_PREFIX}{config_id}_{sanitized}"
|
||||
|
||||
|
||||
def _resolve_dir(data_dir: Path | None) -> Path:
|
||||
if data_dir is not None:
|
||||
return Path(data_dir)
|
||||
from app.config import settings
|
||||
|
||||
return Path(settings.data_dir)
|
||||
|
||||
|
||||
def _load_configs(data_dir: Path):
|
||||
from app.services.ext_data import ExtConfigStore
|
||||
|
||||
return ExtConfigStore(data_dir).load_all()
|
||||
|
||||
|
||||
def _numeric_fields(config) -> list:
|
||||
return [f for f in config.fields if f.dtype in _NUMERIC_DTYPES]
|
||||
|
||||
|
||||
def _signal_fields(config) -> list:
|
||||
"""帧 join / 信号条件可用的字段 (数值 + 字符串)。"""
|
||||
return [f for f in config.fields if f.dtype in _SIGNAL_DTYPES]
|
||||
|
||||
|
||||
def ext_string_fields(data_dir: Path | None = None) -> frozenset[str]:
|
||||
"""string 扩展字段的列名集合 (仅供信号条件, 不注册为因子)。"""
|
||||
return frozenset(e["key"] for e in ext_string_field_entries(data_dir))
|
||||
|
||||
|
||||
def ext_string_field_entries(data_dir: Path | None = None) -> list[dict[str, str]]:
|
||||
"""string 扩展字段条目 [{key, label}], 供 /options 与 AI 提示词展示。"""
|
||||
root = _resolve_dir(data_dir)
|
||||
return [
|
||||
{"key": ext_column_name(cfg.id, f.name), "label": f"{cfg.label}·{f.label or f.name}"[:40]}
|
||||
for cfg in _load_configs(root)
|
||||
for f in cfg.fields
|
||||
if f.dtype == "string"
|
||||
]
|
||||
|
||||
|
||||
def ext_factor_specs(data_dir: Path | None = None) -> list[FactorSpec]:
|
||||
"""扩展表数值字段的 base 因子条目 (列自身即值, 无依赖)。
|
||||
|
||||
id 含非 ASCII (中文字段名) 的字段跳过注册: DSL 公式标识符是
|
||||
ASCII-only, 注册一个公式里写不出来的因子只会误导; 该列仍参与
|
||||
帧 join, 信号条件 (数值比较) 照常可用。
|
||||
"""
|
||||
root = _resolve_dir(data_dir)
|
||||
specs: list[FactorSpec] = []
|
||||
for cfg in _load_configs(root):
|
||||
for f in _numeric_fields(cfg):
|
||||
fid = ext_column_name(cfg.id, f.name)
|
||||
if not fid.isascii():
|
||||
continue
|
||||
specs.append(FactorSpec(
|
||||
id=fid,
|
||||
label=f"{cfg.label}·{f.label}"[:32],
|
||||
group="扩展数据",
|
||||
formula_text=(
|
||||
f"扩展表「{cfg.label}」字段 {f.name} "
|
||||
f"({'时序·按交易日对齐' if cfg.mode == 'timeseries' else '最新快照·仅当日帧'})"
|
||||
),
|
||||
kind="base",
|
||||
warmup_bars=1,
|
||||
scale_free=False,
|
||||
tags=("ext", cfg.id),
|
||||
))
|
||||
return specs
|
||||
|
||||
|
||||
def ext_factor_ids(data_dir: Path | None = None) -> frozenset[str]:
|
||||
"""当前扩展因子 id 集合 (供补算入口判断是否需要注入扩展列)。"""
|
||||
return frozenset(s.id for s in ext_factor_specs(data_dir))
|
||||
|
||||
|
||||
def ensure_synced(data_dir: Path | None = None) -> None:
|
||||
"""把扩展因子同步进注册表 (幂等, 按配置目录签名跳过)。
|
||||
|
||||
以注册表中已存在的 ext_ 前缀条目为权威做增删 —— 不触碰内置目录与
|
||||
用户自定义因子 (uf_/cf_)。重复注册采用"先注销再注册"模式
|
||||
(与 api/factors.py 状态迁移一致), 避免版本未提升时的 fail-closed 拒绝。
|
||||
"""
|
||||
global _sync_state
|
||||
root = _resolve_dir(data_dir)
|
||||
from app.services.ext_data import _ext_config_dir_signature
|
||||
|
||||
ext_base = root / "ext_data"
|
||||
# 目录不存在 = 明确的"无配置" (空签名, 继续同步以清理残留注册);
|
||||
# 目录存在但扫描失败才跳过 (fail-open, 不清空已注册条目)。
|
||||
if not ext_base.exists():
|
||||
sig: tuple | None = ()
|
||||
else:
|
||||
sig = _ext_config_dir_signature(ext_base)
|
||||
if sig is None:
|
||||
return
|
||||
key = (str(root), sig)
|
||||
if _sync_state == key:
|
||||
return
|
||||
desired = ext_factor_specs(root)
|
||||
desired_ids = {s.id for s in desired}
|
||||
from app.factors.registry import _REGISTRY
|
||||
|
||||
for fid in [f for f in list(_REGISTRY) if f.startswith(EXT_PREFIX) and f not in desired_ids]:
|
||||
try:
|
||||
unregister_factor(fid)
|
||||
except ValueError:
|
||||
logger.warning("扩展因子注销失败: %s", fid)
|
||||
for spec in desired:
|
||||
if get_factor(spec.id) is not None:
|
||||
with contextlib.suppress(ValueError):
|
||||
unregister_factor(spec.id)
|
||||
register_factor(spec)
|
||||
_sync_state = key
|
||||
|
||||
|
||||
def _timeseries_signature(ts_dir: Path) -> tuple | None:
|
||||
"""时序分区签名: (分区目录名, part.parquet mtime_ns, size)。"""
|
||||
try:
|
||||
sig = []
|
||||
for d in sorted(ts_dir.glob("date=*")):
|
||||
part = d / "part.parquet"
|
||||
if d.is_dir() and part.exists():
|
||||
st = part.stat()
|
||||
sig.append((d.name, st.st_mtime_ns, st.st_size))
|
||||
return tuple(sig)
|
||||
except OSError:
|
||||
return None
|
||||
|
||||
|
||||
def _select_fields(df: pl.DataFrame, config, fields: list, *, with_date: str | None) -> pl.DataFrame:
|
||||
"""选列 + 统一 dtype: int/float → Float64 (数值阈值), string → Utf8 (contains)。"""
|
||||
exprs = [pl.col("symbol").cast(pl.Utf8)]
|
||||
for f in fields:
|
||||
name = ext_column_name(config.id, f.name)
|
||||
if f.name not in df.columns:
|
||||
continue # 分区 schema 漂移: 缺列以 null 补 (diagonal concat)
|
||||
dtype = pl.Float64 if f.dtype in _NUMERIC_DTYPES else pl.Utf8
|
||||
exprs.append(pl.col(f.name).cast(dtype).alias(name))
|
||||
if len(exprs) == 1:
|
||||
return pl.DataFrame()
|
||||
out = df.select(exprs)
|
||||
if with_date is not None:
|
||||
out = out.with_columns(pl.lit(with_date).alias("_ext_date"))
|
||||
return out
|
||||
|
||||
|
||||
def _timeseries_frame(root: Path, config, fields: list) -> pl.DataFrame:
|
||||
"""全量时序扩展帧 (symbol, _ext_date, ext 列); 按分区签名缓存。
|
||||
|
||||
缓存不过滤日期范围: 调用方用帧自身日期范围在 join 后自然裁剪,
|
||||
避免按日期范围缓存导致的键膨胀。
|
||||
"""
|
||||
ts_dir = root / "ext_data" / config.id / "timeseries"
|
||||
sig = _timeseries_signature(ts_dir)
|
||||
if sig is not None and not sig:
|
||||
return pl.DataFrame()
|
||||
key = (str(root), config.id, "timeseries")
|
||||
if sig is not None:
|
||||
cached = _frame_cache.get(key)
|
||||
if cached is not None and cached[0] == sig:
|
||||
return cached[1]
|
||||
parts: list[pl.DataFrame] = []
|
||||
if sig is not None:
|
||||
for d in sorted(ts_dir.glob("date=*")):
|
||||
part = d / "part.parquet"
|
||||
if not (d.is_dir() and part.exists()):
|
||||
continue
|
||||
try:
|
||||
raw = pl.read_parquet(part)
|
||||
except Exception as e:
|
||||
logger.warning("扩展表 %s 分区 %s 读取失败, 跳过: %s", config.id, d.name, e)
|
||||
continue
|
||||
frag = _select_fields(raw, config, fields, with_date=d.name[5:])
|
||||
if not frag.is_empty():
|
||||
parts.append(frag)
|
||||
frame = (
|
||||
pl.concat(parts, how="diagonal").unique(subset=["symbol", "_ext_date"], keep="last")
|
||||
if parts else pl.DataFrame()
|
||||
)
|
||||
if sig is not None:
|
||||
_frame_cache[key] = (sig, frame)
|
||||
return frame
|
||||
|
||||
|
||||
def _snapshot_frame(root: Path, config, fields: list) -> pl.DataFrame:
|
||||
"""快照扩展帧 (symbol, ext 列); 按 part.parquet (mtime, size) 签名缓存。"""
|
||||
path = root / "ext_data" / config.id / "part.parquet"
|
||||
try:
|
||||
sig = None
|
||||
if path.exists():
|
||||
st = path.stat()
|
||||
sig = (st.st_mtime_ns, st.st_size)
|
||||
if sig is None:
|
||||
return pl.DataFrame()
|
||||
key = (str(root), config.id, "snapshot")
|
||||
cached = _frame_cache.get(key)
|
||||
if cached is not None and cached[0] == sig:
|
||||
return cached[1]
|
||||
frame = _select_fields(pl.read_parquet(path), config, fields, with_date=None)
|
||||
if not frame.is_empty():
|
||||
frame = frame.unique(subset=["symbol"], keep="last")
|
||||
_frame_cache[key] = (sig, frame)
|
||||
return frame
|
||||
except Exception as e:
|
||||
logger.warning("扩展表 %s 快照读取失败, 跳过: %s", config.id, e)
|
||||
return pl.DataFrame()
|
||||
|
||||
|
||||
def attach_ext_columns(
|
||||
df: pl.DataFrame,
|
||||
*,
|
||||
include_snapshot: bool,
|
||||
data_dir: Path | None = None,
|
||||
) -> pl.DataFrame:
|
||||
"""把扩展表信号列 (数值 + 字符串) join 到 enriched 帧上 (无配置/无匹配时原样返回)。
|
||||
|
||||
include_snapshot 仅应由单日帧 (当日/盘中) 路径传 True; 多日历史帧
|
||||
传 False 以规避快照"最新值"造成的未来函数。单个配置失败只跳过该配置。
|
||||
"""
|
||||
if df.is_empty() or "symbol" not in df.columns:
|
||||
return df
|
||||
root = _resolve_dir(data_dir)
|
||||
configs = _load_configs(root)
|
||||
if not configs:
|
||||
return df
|
||||
|
||||
if "_ext_date" in df.columns: # pragma: no cover - 防御内部临时列名被占用
|
||||
return df
|
||||
has_date = "date" in df.columns
|
||||
tmp_date = False
|
||||
try:
|
||||
for cfg in configs:
|
||||
fields = _signal_fields(cfg)
|
||||
if not fields:
|
||||
continue
|
||||
try:
|
||||
if cfg.mode == "timeseries":
|
||||
if not has_date:
|
||||
continue # 无日期列无法 PIT 对齐, 跳过 (ETF/指数单行帧等)
|
||||
ext = _timeseries_frame(root, cfg, fields)
|
||||
if ext.is_empty():
|
||||
continue
|
||||
if not tmp_date:
|
||||
df = df.with_columns(pl.col("date").cast(pl.Utf8).alias("_ext_date"))
|
||||
tmp_date = True
|
||||
new_cols = [c for c in ext.columns if c not in df.columns and c != "_ext_date"]
|
||||
if not new_cols:
|
||||
continue
|
||||
df = df.join(
|
||||
ext.select(["symbol", "_ext_date", *new_cols]),
|
||||
on=["symbol", "_ext_date"],
|
||||
how="left",
|
||||
)
|
||||
elif include_snapshot:
|
||||
snap = _snapshot_frame(root, cfg, fields)
|
||||
if snap.is_empty():
|
||||
continue
|
||||
new_cols = [c for c in snap.columns if c not in df.columns]
|
||||
if not new_cols:
|
||||
continue
|
||||
df = df.join(snap.select(["symbol", *new_cols]), on="symbol", how="left")
|
||||
except Exception as e:
|
||||
logger.warning("扩展表 %s 列注入失败, 跳过该表: %s", cfg.id, e)
|
||||
finally:
|
||||
if tmp_date:
|
||||
df = df.drop("_ext_date")
|
||||
return df
|
||||
|
||||
|
||||
def invalidate_ext_caches(data_dir: Path | None = None, *, keep_strategy_cache: bool = False) -> None:
|
||||
"""扩展数据/配置变更后的失效入口 (写入端自动调用)。
|
||||
|
||||
清扩展帧缓存与注册同步状态 (下次读取重新加载), 并清策略结果缓存 ——
|
||||
策略历史窗口磁盘缓存里已含旧扩展列。repo 内存 enriched 缓存由
|
||||
API 层 (repo.clear_cache) 补充清理。
|
||||
|
||||
keep_strategy_cache=True: 例行数据刷新 (定时拉取) 只失效帧缓存 —— 下次
|
||||
策略运行自然读到新值, 但不销毁已算好的结果。周期性清空会让策略页在两次
|
||||
重算之间整页空白 (小服务器上全量重算需分钟级), 例行刷新的取舍是保留旧
|
||||
结果 (页面秒加载) 而非黑屏; 手动上传/配置变更仍走全清。
|
||||
"""
|
||||
global _sync_state
|
||||
root_key = str(_resolve_dir(data_dir))
|
||||
for key in [k for k in _frame_cache if k[0] == root_key]:
|
||||
_frame_cache.pop(key, None)
|
||||
_sync_state = None
|
||||
if keep_strategy_cache:
|
||||
return
|
||||
from app.config import settings as _settings
|
||||
from app.services import strategy_cache
|
||||
|
||||
try:
|
||||
strategy_cache.clear_cache(Path(data_dir) if data_dir else Path(_settings.data_dir))
|
||||
except Exception as e:
|
||||
logger.warning("扩展数据变更后策略缓存清理失败: %s", e)
|
||||
@@ -0,0 +1,400 @@
|
||||
"""因子注册表 (L-REG) — 因子元数据的单一权威来源。
|
||||
|
||||
P1 收口范围: 目录元数据 (id/label/group/公式)、虚拟因子依赖声明、评分预热窗口。
|
||||
三处历史清单在此合一:
|
||||
- backtest/factor.py FACTOR_COLUMNS (由 factor_columns_view() 生成兼容别名)
|
||||
- strategy/scoring.py VIRTUAL_SCORING_DEPENDENCIES (由 virtual_dependencies() 生成)
|
||||
- strategy/scoring.py _ROLLING_SCORING_WARMUP (由 scoring_warmups() 生成)
|
||||
|
||||
P1 边界 (诚实声明):
|
||||
- scoring_value_expr 的表达式分发仍留在 scoring.py, 注册表不含计算逻辑;
|
||||
复合/自定义因子 (composite/custom) 与 DSL 在 P2/P3 接入后再收口。
|
||||
- unit 字段 P1 统一 "none": 单位口径涉及金融数据契约 (CONTRIBUTING §3),
|
||||
未经逐因子核对禁止猜测填充; 前端 P1 也不按 unit 格式化。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Literal
|
||||
|
||||
Kind = Literal["base", "virtual", "composite", "custom"]
|
||||
Direction = Literal["high", "low", "none"]
|
||||
Unit = Literal["ratio", "pct", "score", "count", "days", "currency", "none"]
|
||||
PitSource = Literal["financial_announce", "share_capital_announce", "none"]
|
||||
Stability = Literal["stable", "experimental", "deprecated"]
|
||||
|
||||
_ALL_ASSETS = frozenset({"stock", "etf"})
|
||||
_STOCK_ONLY = frozenset({"stock"})
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FactorSpec:
|
||||
id: str
|
||||
label: str
|
||||
group: str
|
||||
formula_text: str
|
||||
kind: Kind = "base"
|
||||
version: int = 1
|
||||
# base: 空集合 = 已物化列自身; virtual: 展开到 enriched base 列
|
||||
dependencies: frozenset[str] = field(default_factory=frozenset)
|
||||
direction: Direction = "none" # P1 不预填: 方向以最近检验 IC 符号为准 (见平台方案 §3.6)
|
||||
unit: Unit = "none"
|
||||
warmup_bars: int = 1
|
||||
pit: bool = False
|
||||
pit_source: PitSource = "none"
|
||||
asset_types: frozenset[str] = _ALL_ASSETS
|
||||
incremental_safe: bool = True
|
||||
scale_free: bool = True
|
||||
null_policy: Literal["keep", "drop_row"] = "keep"
|
||||
stability: Stability = "stable"
|
||||
tags: tuple[str, ...] = ()
|
||||
# composite 专用: ((成员 id, 权重), ...); 其余类型为空
|
||||
components: tuple[tuple[str, float], ...] = ()
|
||||
|
||||
def column_view(self) -> dict:
|
||||
"""历史 FACTOR_COLUMNS 条目视图 (键与顺序兼容)。"""
|
||||
return {"id": self.id, "label": self.label, "group": self.group, "desc": self.formula_text}
|
||||
|
||||
|
||||
def _base(fid: str, label: str, group: str, desc: str, **overrides) -> FactorSpec:
|
||||
return FactorSpec(id=fid, label=label, group=group, formula_text=desc, kind="base", **overrides)
|
||||
|
||||
|
||||
def _virtual(fid: str, label: str, group: str, desc: str, deps: frozenset[str], **overrides) -> FactorSpec:
|
||||
return FactorSpec(
|
||||
id=fid, label=label, group=group, formula_text=desc,
|
||||
kind="virtual", dependencies=deps, **overrides,
|
||||
)
|
||||
|
||||
|
||||
def _financial(fid: str, label: str, desc: str) -> FactorSpec:
|
||||
return FactorSpec(
|
||||
id=fid, label=label, group="财务", formula_text=desc,
|
||||
kind="base", pit=True, pit_source="financial_announce", asset_types=_STOCK_ONLY,
|
||||
)
|
||||
|
||||
|
||||
# 顺序即历史 FACTOR_COLUMNS 顺序 (mining_schedule 取前 48 个, 不得重排)。
|
||||
_CATALOG: tuple[FactorSpec, ...] = (
|
||||
# --- 动量 ---
|
||||
_base("momentum_5d", "5日动量", "动量", "5个交易日累计收益率"),
|
||||
_base("momentum_10d", "10日动量", "动量", "10个交易日累计收益率"),
|
||||
_base("momentum_20d", "20日动量", "动量", "20个交易日累计收益率"),
|
||||
_base("momentum_30d", "30日动量", "动量", "30个交易日累计收益率"),
|
||||
_base("momentum_60d", "60日动量", "动量", "60个交易日累计收益率"),
|
||||
_base("change_pct", "日涨跌幅", "动量", "当日收盘相对前收盘的收益率"),
|
||||
# --- 均线偏离 (虚拟) ---
|
||||
*(
|
||||
_virtual(
|
||||
f"ma{period}_bias", f"MA{period}乖离", "均线偏离", f"收盘价 / MA{period} - 1",
|
||||
deps=frozenset({"close", f"ma{period}"}),
|
||||
)
|
||||
for period in (5, 10, 20, 30, 60)
|
||||
),
|
||||
*(
|
||||
_virtual(
|
||||
f"ema{period}_bias", f"EMA{period}乖离", "均线偏离", f"收盘价 / EMA{period} - 1",
|
||||
deps=frozenset({"close", f"ema{period}"}),
|
||||
)
|
||||
for period in (5, 10, 20, 30, 60)
|
||||
),
|
||||
# --- 超买超卖 ---
|
||||
_base("rsi_6", "RSI(6)", "超买超卖", "6日相对强弱指标"),
|
||||
_base("rsi_14", "RSI(14)", "超买超卖", "14日相对强弱指标"),
|
||||
_base("rsi_24", "RSI(24)", "超买超卖", "24日相对强弱指标"),
|
||||
# --- 趋势 ---
|
||||
_base(
|
||||
"macd_hist", "MACD柱(原值)", "趋势",
|
||||
"兼容历史研究; 跨股票比较建议优先使用MACD柱强度",
|
||||
scale_free=False,
|
||||
),
|
||||
_virtual("macd_dif_pct", "MACD DIF强度", "趋势", "MACD DIF / 收盘价", deps=frozenset({"close", "macd_dif"})),
|
||||
_virtual("macd_dea_pct", "MACD DEA强度", "趋势", "MACD DEA / 收盘价", deps=frozenset({"close", "macd_dea"})),
|
||||
_virtual("macd_hist_pct", "MACD柱强度", "趋势", "MACD柱 / 收盘价, 消除股价尺度影响", deps=frozenset({"close", "macd_hist"})),
|
||||
_base("kdj_k", "KDJ-K", "趋势", "KDJ指标K值"),
|
||||
_base("kdj_d", "KDJ-D", "趋势", "KDJ指标D值"),
|
||||
_base("kdj_j", "KDJ-J", "趋势", "KDJ指标J值"),
|
||||
_virtual(
|
||||
"boll_position", "布林位置", "趋势", "收盘价在布林带下轨到上轨之间的位置",
|
||||
deps=frozenset({"close", "boll_upper", "boll_lower"}),
|
||||
),
|
||||
# --- 波动率 ---
|
||||
_base("annual_vol_20d", "20日波动率", "波动率", "20日收益率年化标准差"),
|
||||
_base("atr_14", "ATR(14)原值", "波动率", "兼容历史研究; 跨股票比较建议优先使用ATR相对波动", scale_free=False),
|
||||
_virtual("atr_pct", "ATR相对波动", "波动率", "ATR(14) / 收盘价", deps=frozenset({"close", "atr_14"})),
|
||||
_base("amplitude", "日振幅", "波动率", "当日高低价差 / 前收盘价"),
|
||||
_virtual(
|
||||
"boll_width", "布林带宽", "波动率", "布林带上下轨宽度 / MA20",
|
||||
deps=frozenset({"ma20", "boll_upper", "boll_lower"}),
|
||||
),
|
||||
# --- 量价 ---
|
||||
_base("vol_ratio_5d", "5日量比", "量价", "当日成交量 / 前5日平均成交量"),
|
||||
_virtual(
|
||||
"vol_ratio_10d", "10日量比", "量价", "当日成交量 / 前10日平均成交量",
|
||||
deps=frozenset({"volume"}), warmup_bars=11,
|
||||
),
|
||||
_virtual(
|
||||
"vol_trend_5_10", "成交量趋势", "量价", "5日平均成交量 / 10日平均成交量 - 1",
|
||||
deps=frozenset({"vol_ma5", "vol_ma10"}),
|
||||
),
|
||||
_base("turnover_rate", "换手率", "量价", "使用历史时点流通股本计算的当日换手率"),
|
||||
_virtual(
|
||||
"turnover_ratio_5d", "换手率放大", "量价", "当日换手率 / 前5日平均换手率 - 1",
|
||||
deps=frozenset({"turnover_rate"}), warmup_bars=6,
|
||||
),
|
||||
_virtual(
|
||||
"log_amount", "成交额对数", "量价", "ln(成交额 + 1), 降低极端规模影响",
|
||||
deps=frozenset({"amount"}), scale_free=False,
|
||||
),
|
||||
_virtual(
|
||||
"amount_ratio_5d", "成交额放大", "量价", "当日成交额 / 前5日平均成交额 - 1",
|
||||
deps=frozenset({"amount"}), warmup_bars=6,
|
||||
),
|
||||
# --- 价格位置 ---
|
||||
_virtual("gap_return", "开盘跳空", "价格位置", "开盘价 / 前收盘价 - 1", deps=frozenset({"open", "prev_close"})),
|
||||
_virtual("intraday_return", "日内收益", "价格位置", "收盘价 / 开盘价 - 1", deps=frozenset({"open", "close"})),
|
||||
_virtual(
|
||||
"close_position", "收盘位置", "价格位置", "收盘价在当日最低价到最高价之间的位置",
|
||||
deps=frozenset({"high", "low", "close"}),
|
||||
),
|
||||
_virtual(
|
||||
"distance_to_high_60d", "距60日高点", "价格位置", "收盘价 / 60日最高收盘价 - 1",
|
||||
deps=frozenset({"close", "high_60d"}),
|
||||
),
|
||||
_virtual(
|
||||
"distance_from_low_60d", "距60日低点", "价格位置", "收盘价 / 60日最低收盘价 - 1",
|
||||
deps=frozenset({"close", "low_60d"}),
|
||||
),
|
||||
_virtual(
|
||||
"vwap_bias", "VWAP乖离", "价格位置", "收盘价 / 当日成交均价 - 1, 成交均价 = 成交额 / (成交量x100)",
|
||||
deps=frozenset({"close", "volume", "amount"}),
|
||||
),
|
||||
# --- 收益形态 (虚拟, 滚动窗口) ---
|
||||
_virtual(
|
||||
"max_ret_20d", "20日最大单日涨幅", "收益形态", "近20个交易日单日涨幅最大值(彩票效应, 高值代表博彩型特征强)",
|
||||
deps=frozenset({"close"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"ret_skew_20d", "20日收益偏度", "收益形态", "近20个交易日日收益偏度, 高值代表右偏(偶发大涨)",
|
||||
deps=frozenset({"close"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"up_days_20d", "20日上涨天数", "收益形态", "近20个交易日中上涨天数(0~20)",
|
||||
deps=frozenset({"close"}), warmup_bars=21,
|
||||
),
|
||||
# --- 流动性 (虚拟) ---
|
||||
_virtual(
|
||||
"amihud_20d", "20日Amihud非流动性", "流动性", "近20日平均 |日涨跌幅| / 成交额(亿元), 高值代表流动性差",
|
||||
deps=frozenset({"close", "amount"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"turnover_z_60d", "换手率60日z分", "流动性", "(当日换手率 - 前60日均值) / 前60日标准差, 衡量换手异动",
|
||||
deps=frozenset({"turnover_rate"}), warmup_bars=61,
|
||||
),
|
||||
# --- 量价 (续) ---
|
||||
_virtual(
|
||||
"vol_price_corr_20d", "20日量价相关", "量价", "近20个交易日日涨跌幅与成交量的相关系数, 高值代表量价同向",
|
||||
deps=frozenset({"close", "volume"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"vol_trend_5_60", "量能趋势(5/60)", "量价", "5日平均成交量 / 60日平均成交量 - 1",
|
||||
deps=frozenset({"volume"}), warmup_bars=60,
|
||||
),
|
||||
# --- 涨停基因 (虚拟) ---
|
||||
_virtual(
|
||||
"limit_up_count_20d", "涨停基因(20日)", "涨停基因", "近20个交易日涨停次数",
|
||||
deps=frozenset({"consecutive_limit_ups"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"limit_up_count_60d", "涨停基因(60日)", "涨停基因", "近60个交易日涨停次数",
|
||||
deps=frozenset({"consecutive_limit_ups"}), warmup_bars=61,
|
||||
),
|
||||
# --- 财务 (点时, 仅股票) ---
|
||||
_financial("pb_latest", "市净率(最新公告)", "收盘价 / 最新已公告每股净资产; 无财务数据或公告前为空"),
|
||||
_financial("roe_latest", "ROE(最新公告)", "最新已公告净资产收益率(%); 无财务数据或公告前为空"),
|
||||
_financial("gross_margin_latest", "毛利率(最新公告)", "最新已公告销售毛利率(%)"),
|
||||
_financial("net_margin_latest", "净利率(最新公告)", "最新已公告销售净利率(%)"),
|
||||
_financial("revenue_yoy_latest", "营收增速(最新公告)", "最新已公告营业收入同比(%)"),
|
||||
_financial("net_income_yoy_latest", "净利增速(最新公告)", "最新已公告归母净利润同比(%)"),
|
||||
_financial("debt_ratio_latest", "资产负债率(最新公告)", "最新已公告资产负债率(%)"),
|
||||
# --- 扩充批次 (2026-09-05): 规模/收益分解/长窗口/下行风险/量能潮/换手水平 ---
|
||||
_virtual(
|
||||
"log_float_mv", "流通市值对数", "规模",
|
||||
"ln(收盘价 x 当日成交量 / 换手率), 由换手率反推流通股本, 高值代表大盘",
|
||||
deps=frozenset({"close", "volume", "turnover_rate"}), scale_free=False,
|
||||
),
|
||||
_virtual(
|
||||
"momentum_120d", "120日动量", "动量",
|
||||
"120个交易日累计收益率 (中期动量, 与短窗口互补)",
|
||||
deps=frozenset({"close"}), warmup_bars=121,
|
||||
),
|
||||
_virtual(
|
||||
"mom_accel_20_60", "动量加速度", "动量",
|
||||
"20日动量 - 60日动量, 衡量近期动量相对中期是否增强",
|
||||
deps=frozenset({"momentum_20d", "momentum_60d"}),
|
||||
),
|
||||
_virtual(
|
||||
"rsi_14_delta_5d", "RSI五日变化", "超买超卖",
|
||||
"RSI(14) - 5日前的RSI(14), 衡量强弱指标的边际变化",
|
||||
deps=frozenset({"rsi_14"}), warmup_bars=6,
|
||||
),
|
||||
_virtual(
|
||||
"overnight_ret_20d", "20日隔夜收益", "收益形态",
|
||||
"近20日累计隔夜收益(开盘价/前收盘-1求和), A股隔夜与日内收益的定价机制不同",
|
||||
deps=frozenset({"open", "prev_close"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"intraday_ret_20d", "20日日内收益", "收益形态",
|
||||
"近20日累计日内收益(收盘价/开盘价-1求和), 与隔夜收益构成收益分解",
|
||||
deps=frozenset({"open", "close"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"downside_vol_20d", "20日下行波动", "波动率",
|
||||
"sqrt(近20日 min(日收益,0)^2 均值), 只度量下跌侧风险",
|
||||
deps=frozenset({"close"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"vol_regime_5_60", "波动率状态(5/60)", "波动率",
|
||||
"5日收益标准差 / 60日收益标准差, 高值代表波动骤然放大",
|
||||
deps=frozenset({"close"}), warmup_bars=61,
|
||||
),
|
||||
_virtual(
|
||||
"amplitude_trend_20_60", "振幅趋势(20/60)", "波动率",
|
||||
"20日平均振幅 / 60日平均振幅 - 1",
|
||||
deps=frozenset({"amplitude"}), warmup_bars=61,
|
||||
),
|
||||
_virtual(
|
||||
"obv_trend_20d", "20日量能潮", "量价",
|
||||
"近20日 sign(日收益)x成交量 之和 / (20日均量x20), 有界[-1,1], 净买入方向的一致性",
|
||||
deps=frozenset({"close", "volume"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"amount_mean_20d", "20日均成交额(亿)", "量价",
|
||||
"近20日平均成交额(亿元), 规模/流动性水平量",
|
||||
deps=frozenset({"amount"}), warmup_bars=21, scale_free=False,
|
||||
),
|
||||
_virtual(
|
||||
"turnover_mean_20d", "20日均换手", "流动性",
|
||||
"近20日平均换手率, A股经典低换手溢价因子",
|
||||
deps=frozenset({"turnover_rate"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"turnover_std_20d", "20日换手波动", "流动性",
|
||||
"近20日换手率标准差 / 均值 (变异系数), 衡量交易活跃的稳定性",
|
||||
deps=frozenset({"turnover_rate"}), warmup_bars=21,
|
||||
),
|
||||
_virtual(
|
||||
"position_240d", "一年价格位置", "价格位置",
|
||||
"收盘价在近240个交易日最低价到最高价之间的位置 (0~1)",
|
||||
deps=frozenset({"close"}), warmup_bars=241,
|
||||
),
|
||||
_virtual(
|
||||
"distance_to_high_240d", "距一年高点", "价格位置",
|
||||
"收盘价 / 240日最高收盘价 - 1, 接近0代表贴近一年新高",
|
||||
deps=frozenset({"close"}), warmup_bars=241,
|
||||
),
|
||||
_virtual(
|
||||
"kdj_kd_diff", "KDJ K-D差", "趋势",
|
||||
"KDJ K值 - D值, 正值代表快线在慢线上方",
|
||||
deps=frozenset({"kdj_k", "kdj_d"}),
|
||||
),
|
||||
)
|
||||
|
||||
_REGISTRY: dict[str, FactorSpec] = {}
|
||||
|
||||
|
||||
def register_factor(spec: FactorSpec) -> None:
|
||||
"""注册因子; 重复 id 且版本未增时拒绝 (fail-closed)。"""
|
||||
existing = _REGISTRY.get(spec.id)
|
||||
if existing is not None and existing.version >= spec.version:
|
||||
raise ValueError(f"factor id 已注册且版本未提升: {spec.id}")
|
||||
_REGISTRY[spec.id] = spec
|
||||
|
||||
|
||||
for _spec in _CATALOG:
|
||||
register_factor(_spec)
|
||||
|
||||
|
||||
def get_factor(fid: str) -> FactorSpec | None:
|
||||
return _REGISTRY.get(fid)
|
||||
|
||||
|
||||
def unregister_factor(fid: str) -> FactorSpec | None:
|
||||
"""注销动态注册的因子 (内置目录因子不可注销, fail-closed)。"""
|
||||
if any(spec.id == fid for spec in _CATALOG):
|
||||
raise ValueError(f"内置因子不可注销: {fid}")
|
||||
return _REGISTRY.pop(fid, None)
|
||||
|
||||
|
||||
def _ordered_specs() -> list[FactorSpec]:
|
||||
"""内置目录顺序在前, 动态注册因子 (custom/composite) 按注册顺序追加。"""
|
||||
ordered: list[FactorSpec] = list(_CATALOG)
|
||||
known = {spec.id for spec in _CATALOG}
|
||||
ordered.extend(spec for fid, spec in _REGISTRY.items() if fid not in known)
|
||||
return ordered
|
||||
|
||||
|
||||
def _ensure_ext_factors() -> None:
|
||||
"""扩展表字段惰性同步 (配置目录签名幂等); 失败不阻断注册表读取。"""
|
||||
try:
|
||||
from app.factors.ext_factors import ensure_synced
|
||||
|
||||
ensure_synced()
|
||||
except Exception:
|
||||
import logging
|
||||
|
||||
logging.getLogger(__name__).debug("ext factor sync skipped", exc_info=True)
|
||||
|
||||
|
||||
def all_factors(
|
||||
asset_type: str | None = None,
|
||||
stable_only: bool = False,
|
||||
) -> list[FactorSpec]:
|
||||
"""按目录顺序返回因子; asset_type 过滤适用资产, stable_only 过滤实验/废弃因子。
|
||||
|
||||
返回前惰性同步扩展表因子 (ext_ 前缀 base 条目), 使信号字段白名单、
|
||||
因子库列表和 AI 提示词看到同一份扩展字段清单。
|
||||
"""
|
||||
_ensure_ext_factors()
|
||||
return [
|
||||
spec for spec in _ordered_specs()
|
||||
if (asset_type is None or asset_type in spec.asset_types)
|
||||
and (not stable_only or spec.stability == "stable")
|
||||
]
|
||||
|
||||
|
||||
def factor_dependencies(fids) -> frozenset[str]:
|
||||
"""递归展开依赖到 enriched base 列; 未知 id 原样保留 (与 scoring_dependencies 历史语义一致)。"""
|
||||
resolved: set[str] = set()
|
||||
for fid in fids:
|
||||
spec = _REGISTRY.get(str(fid))
|
||||
if spec is None:
|
||||
resolved.add(str(fid))
|
||||
elif spec.dependencies:
|
||||
resolved.update(spec.dependencies)
|
||||
else:
|
||||
resolved.add(spec.id)
|
||||
return frozenset(resolved)
|
||||
|
||||
|
||||
def factor_columns_view() -> list[dict]:
|
||||
"""历史 FACTOR_COLUMNS 兼容视图 (顺序、键一致; 动态注册因子追加在末尾)。"""
|
||||
return [spec.column_view() for spec in _ordered_specs()]
|
||||
|
||||
|
||||
def virtual_dependencies() -> dict[str, frozenset[str]]:
|
||||
"""历史 VIRTUAL_SCORING_DEPENDENCIES 兼容视图。"""
|
||||
return {
|
||||
spec.id: spec.dependencies
|
||||
for spec in _CATALOG
|
||||
if spec.kind == "virtual" and spec.dependencies
|
||||
}
|
||||
|
||||
|
||||
def scoring_warmups() -> dict[str, int]:
|
||||
"""历史 _ROLLING_SCORING_WARMUP 兼容视图 (仅滚动窗口虚拟因子)。"""
|
||||
return {
|
||||
spec.id: spec.warmup_bars
|
||||
for spec in _CATALOG
|
||||
if spec.kind == "virtual" and spec.warmup_bars > 1
|
||||
}
|
||||
@@ -0,0 +1,190 @@
|
||||
"""自定义/复合因子存储 (P3) — data/user_data/custom_factors/*.json。
|
||||
|
||||
镜像 custom_signals 的持久化写法; 单文件损坏只禁用该因子并告警, 不影响启动
|
||||
(对齐 CONTRIBUTING 第 4 节插件隔离要求)。生命周期状态: draft → active →
|
||||
watch → retired (P4 状态机, 存储字段就绪, 迁移逻辑见巡检设计)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from app.factors.dsl import compile_formula
|
||||
from app.factors.registry import FactorSpec, factor_dependencies, get_factor, register_factor
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
CUSTOM_ID_PATTERN = re.compile(r"^uf_[a-z0-9_]{1,40}$")
|
||||
COMPOSITE_ID_PATTERN = re.compile(r"^cf_[a-z0-9_]{1,40}$")
|
||||
MAX_COMPOSITE_MEMBERS = 8
|
||||
STATUSES = frozenset({"draft", "active", "watch", "retired"})
|
||||
|
||||
|
||||
def _dir(data_dir: Path) -> Path:
|
||||
directory = data_dir / "user_data" / "custom_factors"
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
return directory
|
||||
|
||||
|
||||
def _path(data_dir: Path, factor_id: str) -> Path:
|
||||
return _dir(data_dir) / f"{factor_id}.json"
|
||||
|
||||
|
||||
def load_all(data_dir: Path) -> list[dict]:
|
||||
"""读取全部自定义/复合因子定义; 损坏文件跳过。"""
|
||||
out: list[dict] = []
|
||||
for file in sorted(_dir(data_dir).glob("*.json")):
|
||||
try:
|
||||
out.append(json.loads(file.read_text(encoding="utf-8")))
|
||||
except Exception as exc:
|
||||
logger.warning("custom factor load failed %s: %s", file.name, exc)
|
||||
return out
|
||||
|
||||
|
||||
def save_one(data_dir: Path, definition: dict) -> None:
|
||||
target = _path(data_dir, str(definition["id"]))
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
target.write_text(json.dumps(definition, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
|
||||
|
||||
def delete_one(data_dir: Path, factor_id: str) -> bool:
|
||||
target = _path(data_dir, factor_id)
|
||||
if target.exists():
|
||||
target.unlink()
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _now() -> str:
|
||||
return datetime.now().isoformat(timespec="seconds")
|
||||
|
||||
|
||||
def to_spec(definition: dict) -> FactorSpec:
|
||||
"""定义 → FactorSpec; 校验失败抛 ValueError (调用方 fail-closed)。
|
||||
|
||||
custom: 依赖/预热由 DSL 编译推导 (编译失败即拒绝注册)。
|
||||
composite: 依赖 = 成员递归展开; 预热 = 成员最大值; 循环引用拒绝。
|
||||
"""
|
||||
kind = str(definition.get("kind", "custom"))
|
||||
factor_id = str(definition.get("id", ""))
|
||||
label = str(definition.get("label", "")).strip()
|
||||
if not label:
|
||||
raise ValueError("label 不能为空")
|
||||
pattern = COMPOSITE_ID_PATTERN if kind == "composite" else CUSTOM_ID_PATTERN
|
||||
if not pattern.match(factor_id):
|
||||
raise ValueError(f"id 必须匹配 {pattern.pattern}")
|
||||
status = str(definition.get("status", "draft"))
|
||||
if status not in STATUSES:
|
||||
raise ValueError(f"status 必须是 {sorted(STATUSES)} 之一")
|
||||
|
||||
if kind == "custom":
|
||||
formula = str(definition.get("formula", ""))
|
||||
compiled = compile_formula(formula)
|
||||
if not compiled.ok:
|
||||
first = compiled.errors[0]
|
||||
raise ValueError(f"公式无效 [{first.code}]: {first.message}")
|
||||
return FactorSpec(
|
||||
id=factor_id,
|
||||
label=label,
|
||||
group=str(definition.get("group", "自定义")),
|
||||
formula_text=formula,
|
||||
kind="custom",
|
||||
version=int(definition.get("version", 1)),
|
||||
dependencies=frozenset(compiled.dependencies),
|
||||
warmup_bars=compiled.warmup_bars,
|
||||
direction=str(definition.get("direction", "none")), # type: ignore[arg-type]
|
||||
stability="stable" if status == "active" else "experimental",
|
||||
)
|
||||
|
||||
if kind != "composite":
|
||||
raise ValueError(f"未知 kind: {kind}")
|
||||
members_raw = definition.get("members")
|
||||
if not isinstance(members_raw, dict) or not (2 <= len(members_raw) <= MAX_COMPOSITE_MEMBERS):
|
||||
raise ValueError(f"composite 成员必须是 {2}~{MAX_COMPOSITE_MEMBERS} 个")
|
||||
from app.factors.dsl import BASE_COLUMNS
|
||||
|
||||
components: list[tuple[str, float]] = []
|
||||
for member_id, weight in members_raw.items():
|
||||
member_id = str(member_id)
|
||||
if member_id == factor_id:
|
||||
raise ValueError("composite 不能引用自身")
|
||||
try:
|
||||
weight = float(weight)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(f"成员 {member_id} 权重必须是数字") from exc
|
||||
if not weight:
|
||||
raise ValueError(f"成员 {member_id} 权重不能为 0")
|
||||
# 成员 = 注册表因子 或 enriched 基准列 (已物化, 可直接参与组合)
|
||||
if get_factor(member_id) is None and member_id not in BASE_COLUMNS:
|
||||
raise ValueError(f"未知成员因子: {member_id}")
|
||||
components.append((member_id, weight))
|
||||
# 环检测沿 components 链走 (依赖已展开, 看不到链路成员)
|
||||
seen = {factor_id}
|
||||
frontier = [member_id for member_id, _ in components]
|
||||
while frontier:
|
||||
current = frontier.pop()
|
||||
if current in seen:
|
||||
raise ValueError("composite 成员存在循环引用")
|
||||
seen.add(current)
|
||||
current_spec = get_factor(current)
|
||||
if current_spec is not None and current_spec.kind == "composite":
|
||||
frontier.extend(member_id for member_id, _ in current_spec.components)
|
||||
dependencies = factor_dependencies([member_id for member_id, _ in components])
|
||||
warmup = max(
|
||||
((get_factor(member_id).warmup_bars if get_factor(member_id) else 1) for member_id, _ in components),
|
||||
default=1,
|
||||
)
|
||||
formula_text = " + ".join(
|
||||
f"{weight:g}*zscore({member_id})" for member_id, weight in components
|
||||
)
|
||||
return FactorSpec(
|
||||
id=factor_id,
|
||||
label=label,
|
||||
group=str(definition.get("group", "组合")),
|
||||
formula_text=formula_text,
|
||||
kind="composite",
|
||||
version=int(definition.get("version", 1)),
|
||||
dependencies=dependencies,
|
||||
warmup_bars=warmup,
|
||||
direction=str(definition.get("direction", "none")), # type: ignore[arg-type]
|
||||
components=tuple(components),
|
||||
stability="stable" if status == "active" else "experimental",
|
||||
)
|
||||
|
||||
|
||||
def register_definition(definition: dict) -> FactorSpec:
|
||||
"""定义 → spec → 注册 (重复 id 版本未升时由注册表拒绝)。"""
|
||||
spec = to_spec(definition)
|
||||
register_factor(spec)
|
||||
return spec
|
||||
|
||||
|
||||
def load_into_registry(data_dir: Path) -> list[str]:
|
||||
"""启动期把存储中的因子注册进注册表; 单个失败只跳过并告警。
|
||||
|
||||
多轮加载: composite 成员可能引用尚未加载的 custom/其他 composite (文件按
|
||||
字母序加载, cf_* 先于 uf_*), 失败的 composite 延后重试, 覆盖链式引用;
|
||||
重试用尽仍失败的只告警不阻塞启动。
|
||||
"""
|
||||
loaded: list[str] = []
|
||||
pending = list(load_all(data_dir))
|
||||
for round_index in range(3):
|
||||
deferred: list[dict] = []
|
||||
for definition in pending:
|
||||
try:
|
||||
register_definition(definition)
|
||||
loaded.append(str(definition["id"]))
|
||||
except ValueError as exc:
|
||||
if round_index < 2 and str(definition.get("kind")) == "composite":
|
||||
deferred.append(definition)
|
||||
else:
|
||||
logger.warning("custom factor 注册失败 %s: %s", definition.get("id"), exc)
|
||||
except Exception as exc:
|
||||
logger.warning("custom factor 注册失败 %s: %s", definition.get("id"), exc)
|
||||
if not deferred:
|
||||
break
|
||||
pending = deferred
|
||||
return loaded
|
||||
+385
-143
@@ -16,6 +16,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import shutil
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
|
||||
@@ -464,16 +467,21 @@ def compute_indicators(
|
||||
|
||||
# Pass 3: KDJ
|
||||
if "kdj_k" in want:
|
||||
_kdj_rsv = (
|
||||
100 * (pl.col("close") - pl.col("_kdj_ln"))
|
||||
/ (pl.col("_kdj_hn") - pl.col("_kdj_ln")).fill_null(1e-12)
|
||||
# 9 日内最高价=最低价 (场内货币 ETF、长期无成交标的) 时分母是 0 而不是空值,
|
||||
# fill_null 拦不住: 0/0 得到 NaN, 再被 ewm 递推永久传染。与矩阵路径口径一致 ——
|
||||
# 该日 RSV 置空, EWM 跳过空值后继续递推。
|
||||
_kdj_range = pl.col("_kdj_hn") - pl.col("_kdj_ln")
|
||||
_kdj_rsv = pl.when(_kdj_range > 0).then(
|
||||
100 * (pl.col("close") - pl.col("_kdj_ln")) / _kdj_range
|
||||
)
|
||||
df = df.with_columns([
|
||||
_kdj_rsv.ewm_mean(alpha=1.0 / 3, adjust=False).over("symbol").alias("kdj_k"),
|
||||
_kdj_rsv.ewm_mean(alpha=1.0 / 3, adjust=False, ignore_nulls=True)
|
||||
.over("symbol").alias("kdj_k"),
|
||||
])
|
||||
if "kdj_d" in want:
|
||||
df = df.with_columns([
|
||||
pl.col("kdj_k").ewm_mean(alpha=1.0 / 3, adjust=False).over("symbol").alias("kdj_d"),
|
||||
pl.col("kdj_k").ewm_mean(alpha=1.0 / 3, adjust=False, ignore_nulls=True)
|
||||
.over("symbol").alias("kdj_d"),
|
||||
])
|
||||
if "kdj_j" in want:
|
||||
df = df.with_columns([
|
||||
@@ -663,9 +671,17 @@ def compute_signals(df: pl.DataFrame, needed: set[str] | None = None) -> pl.Data
|
||||
if want:
|
||||
df = df.with_columns([expressions[name] for name in SIGNAL_DEPENDENCIES if name in want])
|
||||
|
||||
# 自定义信号(用户配置的字段+运算符+值组合,编译为布尔列)
|
||||
# 自定义信号(用户配置的字段+运算符+值组合,编译为布尔列)。
|
||||
# 扩展表数值列先行 join (ext_ 因子列 = 帧上已有列): 信号条件与评分引用
|
||||
# 都按列存在性解析。历史多日帧仅注入时序模式 —— 快照代表"最新值",
|
||||
# 历史回看注入会引入未来数据 (CONTRIBUTING §5.3)。
|
||||
from app.factors import ext_factors
|
||||
df = ext_factors.attach_ext_columns(df, include_snapshot=False)
|
||||
# 条件引用的注册表因子列先复用评分物化管线补算 (虚拟/自定义/复合均可)。
|
||||
from app.strategy import custom_signals
|
||||
df = custom_signals.inject(df, _get_custom_signal_exprs(), needed=needed)
|
||||
exprs = _get_custom_signal_exprs()
|
||||
df = custom_signals.materialize_factor_columns(df, exprs, needed=needed)
|
||||
df = custom_signals.inject(df, exprs, needed=needed)
|
||||
|
||||
return df
|
||||
|
||||
@@ -792,9 +808,13 @@ def compute_limit_signals(
|
||||
else:
|
||||
authoritative_date = pl.col("date") == pl.col("date").max()
|
||||
if "limit_up" in df.columns:
|
||||
# >0 与实时路径 (_compute_limit_signals_today) 同守卫: 维表 limit_up 为 0
|
||||
# (数据源未提供该字段的占位值) 不是权威价, 直接采用会让 raw_close >= -0.005
|
||||
# 恒成立, 全部标的被判涨停。
|
||||
effective_limit_up = pl.when(
|
||||
authoritative_date
|
||||
& pl.col("limit_up").is_not_null()
|
||||
& (pl.col("limit_up") > 0)
|
||||
& (pl.col("limit_up") < _SENTINEL)
|
||||
).then(pl.col("limit_up")).otherwise(pl.col("_theoretical_limit_up"))
|
||||
else:
|
||||
@@ -803,6 +823,7 @@ def compute_limit_signals(
|
||||
effective_limit_down = pl.when(
|
||||
authoritative_date
|
||||
& pl.col("limit_down").is_not_null()
|
||||
& (pl.col("limit_down") > 0)
|
||||
& (pl.col("limit_down") < _SENTINEL)
|
||||
).then(pl.col("limit_down")).otherwise(pl.col("_theoretical_limit_down"))
|
||||
else:
|
||||
@@ -972,11 +993,16 @@ def filter_halt_days(df: pl.DataFrame) -> pl.DataFrame:
|
||||
|
||||
停牌日的 open/high 必然为 0 (无集合竞价)。注意 close 可能被数据源
|
||||
填充为前收盘价而非 0, 因此不能用 "OHLC 全零" 判断, 否则会漏过这类
|
||||
停牌记录 (如 *ST 撤销风险警示的停牌日), 污染 MA/ATR 等指标。
|
||||
停牌记录 (如 *ST 撤销风险警示的停牌日), 污染 MA/ATR 等指标。旧版实时
|
||||
落盘还会先把 open/high=0 填成 close, 对这类历史数据用零成交量和零成交额
|
||||
作为兼容判据。
|
||||
"""
|
||||
if df.is_empty() or "open" not in df.columns or "high" not in df.columns:
|
||||
return df
|
||||
return df.filter(~((pl.col("open") == 0) & (pl.col("high") == 0)))
|
||||
halted = (pl.col("open") == 0) & (pl.col("high") == 0)
|
||||
if "volume" in df.columns and "amount" in df.columns:
|
||||
halted = halted | ((pl.col("volume") == 0) & (pl.col("amount") == 0))
|
||||
return df.filter(~halted)
|
||||
|
||||
|
||||
# ================================================================
|
||||
@@ -1045,13 +1071,28 @@ def _select_storage_cols(df: pl.DataFrame) -> pl.DataFrame:
|
||||
|
||||
DEVIATION_WINDOWS: tuple[int, ...] = (3, 10, 30)
|
||||
|
||||
# 各交易所基准指数 (偏离值规则的「对应指数」近似): 优先分类指数, 缺失时回退
|
||||
# 各板块基准指数 (偏离值规则的「对应指数」, 按交易所官方口径): 优先首选, 缺失时回退
|
||||
# - 沪主板: 上证A指 → 上证指数 (两者差异可忽略)
|
||||
# - 科创板: 科创50 (上交所《交易规则》2026修订 6.12 指定基准) → 上证A指
|
||||
# - 深主板: 深证A指 → 深证成指 (深交所投教口径)
|
||||
# - 创业板: 创业板综合指数 → 深证A指 (深交所投教口径)
|
||||
# - 北交所: 北证50 → 上证指数 (北交所《交易规则》5.4.4)
|
||||
_BENCHMARK_PREFERENCE: dict[str, list[str]] = {
|
||||
"SH": ["000002.SH", "000001.SH"], # 上证A指 → 上证指数
|
||||
"SZ": ["399107.SZ", "399001.SZ"], # 深证A指 → 深证成指
|
||||
"BJ": ["899050.BJ", "000001.SH"], # 北证50 → 上证指数
|
||||
"SH": ["000002.SH", "000001.SH"],
|
||||
"STAR": ["000688.SH", "000002.SH"],
|
||||
"SZ": ["399107.SZ", "399001.SZ"],
|
||||
"GEM": ["399102.SZ", "399107.SZ"],
|
||||
"BJ": ["899050.BJ", "000001.SH"],
|
||||
}
|
||||
|
||||
# 偏离值计算需要的全部基准指数 (quote_service 并入实时显式拉取, 不依赖监控规则)
|
||||
BENCHMARK_INDEX_SYMBOLS: frozenset[str] = frozenset(
|
||||
sym for cands in _BENCHMARK_PREFERENCE.values() for sym in cands
|
||||
)
|
||||
|
||||
# 全部板块基准键 (SH/STAR/SZ/GEM/BJ)
|
||||
BENCH_KEYS: tuple[str, ...] = tuple(_BENCHMARK_PREFERENCE)
|
||||
|
||||
_benchmark_cache: dict[str, tuple[float, pl.DataFrame | None]] = {}
|
||||
_BENCHMARK_CACHE_TTL = 600.0
|
||||
|
||||
@@ -1059,7 +1100,8 @@ _BENCHMARK_CACHE_TTL = 600.0
|
||||
def load_benchmark_momentum(data_dir: Path) -> pl.DataFrame | None:
|
||||
"""读取指数日K, 计算各基准指数的滚动 N 日涨跌幅。
|
||||
|
||||
返回长表: date, bench_exchange, bench_close, bench_mom3d, bench_mom10d, bench_mom30d。
|
||||
返回长表: date, bench_key, bench_close, bench_mom3d, bench_mom10d, bench_mom30d。
|
||||
bench_key 为板块基准键 (SH/STAR/SZ/GEM/BJ, 见 _BENCHMARK_PREFERENCE)。
|
||||
bench_close 供盘中路径外推今日基准动量 (benchmark_momentum_today)。
|
||||
无可用指数数据时返回 None (偏离列置 null, 不阻塞主流程)。
|
||||
进程内按 data_dir 缓存 (TTL 10 分钟)。
|
||||
@@ -1077,11 +1119,11 @@ def load_benchmark_momentum(data_dir: Path) -> pl.DataFrame | None:
|
||||
index_glob = str(Path(data_dir) / "kline_index_daily" / "**" / "*.parquet")
|
||||
wanted: list[str] = []
|
||||
bench_of: dict[str, str] = {}
|
||||
for exchange, candidates in _BENCHMARK_PREFERENCE.items():
|
||||
for bench_key, candidates in _BENCHMARK_PREFERENCE.items():
|
||||
for sym in candidates:
|
||||
if sym not in bench_of:
|
||||
wanted.append(sym)
|
||||
bench_of[sym] = exchange
|
||||
bench_of[sym] = bench_key
|
||||
lf = scan_daily_parquet(
|
||||
index_glob, cast_options=pl.ScanCastOptions(integer_cast="allow-float")
|
||||
)
|
||||
@@ -1094,15 +1136,15 @@ def load_benchmark_momentum(data_dir: Path) -> pl.DataFrame | None:
|
||||
if not df_idx.is_empty():
|
||||
available = set(df_idx["symbol"].to_list())
|
||||
picked = [s for s in wanted if s in available]
|
||||
# 每个交易所取优先级最高的可用基准; 全缺时回退到任一可用基准。
|
||||
# 同一基准可服务多个交易所 (如北证50 缺失时北交所回退上证指数)。
|
||||
# 每个板块取优先级最高的可用基准; 全缺时回退到任一可用基准。
|
||||
# 同一基准可服务多个板块 (如科创50 缺失时科创板回退上证A指)。
|
||||
pairs: list[tuple[str, str]] = []
|
||||
for exchange, candidates in _BENCHMARK_PREFERENCE.items():
|
||||
for bench_key, candidates in _BENCHMARK_PREFERENCE.items():
|
||||
hit = next((s for s in candidates if s in available), None)
|
||||
if hit is None and picked:
|
||||
hit = picked[0]
|
||||
if hit is not None:
|
||||
pairs.append((hit, exchange))
|
||||
pairs.append((hit, bench_key))
|
||||
df_bench = df_idx.filter(pl.col("symbol").is_in([p[0] for p in pairs]))
|
||||
if not df_bench.is_empty():
|
||||
df_bench = df_bench.with_columns(
|
||||
@@ -1111,16 +1153,16 @@ def load_benchmark_momentum(data_dir: Path) -> pl.DataFrame | None:
|
||||
(pl.col("close") / pl.col("close").shift(n).over("symbol") - 1).alias(f"_bm{n}")
|
||||
for n in DEVIATION_WINDOWS
|
||||
]).rename({f"_bm{n}": f"bench_mom{n}d" for n in DEVIATION_WINDOWS})
|
||||
exchange_map = pl.DataFrame({
|
||||
key_map = pl.DataFrame({
|
||||
"symbol": [p[0] for p in pairs],
|
||||
"bench_exchange": [p[1] for p in pairs],
|
||||
"bench_key": [p[1] for p in pairs],
|
||||
})
|
||||
frame = (
|
||||
df_bench.join(exchange_map, on="symbol", how="inner")
|
||||
.select(["date", "bench_exchange", "close",
|
||||
df_bench.join(key_map, on="symbol", how="inner")
|
||||
.select(["date", "bench_key", "close",
|
||||
*[f"bench_mom{n}d" for n in DEVIATION_WINDOWS]])
|
||||
.rename({"close": "bench_close"})
|
||||
.unique(subset=["date", "bench_exchange"])
|
||||
.unique(subset=["date", "bench_key"])
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("基准指数偏离数据加载失败: %s", exc)
|
||||
@@ -1130,14 +1172,21 @@ def load_benchmark_momentum(data_dir: Path) -> pl.DataFrame | None:
|
||||
return frame
|
||||
|
||||
|
||||
def _bench_exchange_expr() -> pl.Expr:
|
||||
"""symbol 后缀 → 交易所 (SH/SZ/BJ), 无法识别时 null。"""
|
||||
def _bench_key_expr() -> pl.Expr:
|
||||
"""symbol → 板块基准键 (SH/STAR/SZ/GEM/BJ), 无法识别时 null。
|
||||
|
||||
北交所按后缀; 沪市按 68 前缀区分科创板; 深市按 30 前缀区分创业板。
|
||||
与 abnormal_moves.board_of 的板块判定同口径。
|
||||
"""
|
||||
code = pl.col("symbol").str.slice(0, 6)
|
||||
suffix = pl.col("symbol").str.slice(-2).str.to_uppercase()
|
||||
return (
|
||||
pl.col("symbol").str.slice(-2).str.to_uppercase().replace(
|
||||
{ex: ex for ex in _BENCHMARK_PREFERENCE},
|
||||
default=None,
|
||||
return_dtype=pl.Utf8,
|
||||
)
|
||||
pl.when(suffix == "BJ").then(pl.lit("BJ"))
|
||||
.when((suffix == "SH") & code.str.starts_with("68")).then(pl.lit("STAR"))
|
||||
.when(suffix == "SH").then(pl.lit("SH"))
|
||||
.when((suffix == "SZ") & code.str.starts_with("30")).then(pl.lit("GEM"))
|
||||
.when(suffix == "SZ").then(pl.lit("SZ"))
|
||||
.otherwise(pl.lit(None, dtype=pl.Utf8))
|
||||
)
|
||||
|
||||
|
||||
@@ -1164,8 +1213,8 @@ def attach_deviation_columns(df: pl.DataFrame, data_dir: Path) -> pl.DataFrame:
|
||||
for n in missing
|
||||
])
|
||||
out = (
|
||||
df.with_columns(_bench_exchange_expr().alias("_bench_ex"))
|
||||
.join(bench, left_on=["_bench_ex", "date"], right_on=["bench_exchange", "date"], how="left")
|
||||
df.with_columns(_bench_key_expr().alias("_bench_ex"))
|
||||
.join(bench, left_on=["_bench_ex", "date"], right_on=["bench_key", "date"], how="left")
|
||||
.with_columns([
|
||||
(pl.col(f"momentum_{n}d") - pl.col(f"bench_mom{n}d")).alias(f"deviate_{n}d")
|
||||
for n in DEVIATION_WINDOWS
|
||||
@@ -1176,7 +1225,12 @@ def attach_deviation_columns(df: pl.DataFrame, data_dir: Path) -> pl.DataFrame:
|
||||
|
||||
|
||||
def _bench_rt_pct_of(index_quotes: pl.DataFrame | None, candidates: list[str]) -> float:
|
||||
"""从实时指数行情取某交易所首选基准的今日涨跌, 缺数据时 0。"""
|
||||
"""从实时指数行情取某交易所首选基准的今日涨跌 (小数制), 缺数据时 0。
|
||||
|
||||
入参 index_quotes 来自 quote_service 的指数展示缓存, 其 change_pct/pct/pct_change
|
||||
列为百分数口径 (CONTRIBUTING §3.1), 消费前必须显式 /100 (#232);
|
||||
close/prev_close 兜底路径本身就是小数, 不转换。
|
||||
"""
|
||||
if index_quotes is None or index_quotes.is_empty():
|
||||
return 0.0
|
||||
df = index_quotes.filter(pl.col("symbol").is_in(candidates))
|
||||
@@ -1191,22 +1245,27 @@ def _bench_rt_pct_of(index_quotes: pl.DataFrame | None, candidates: list[str]) -
|
||||
for col in ("change_pct", "pct", "pct_change"):
|
||||
v = row.get(col)
|
||||
if v is not None:
|
||||
return float(v)
|
||||
return float(v) / 100.0
|
||||
if row.get("close") is not None and row.get("prev_close") is not None and row["prev_close"]:
|
||||
return float(row["close"] / row["prev_close"] - 1)
|
||||
return 0.0
|
||||
|
||||
|
||||
def bench_rt_pct_for(index_quotes: pl.DataFrame | None, bench_key: str) -> float:
|
||||
"""板块基准键的指数今日实时涨跌 (小数制), 供异动总览实时叠加等外部消费。"""
|
||||
return _bench_rt_pct_of(index_quotes, _BENCHMARK_PREFERENCE.get(bench_key, []))
|
||||
|
||||
|
||||
def benchmark_momentum_today(
|
||||
data_dir: Path,
|
||||
index_quotes: pl.DataFrame | None = None,
|
||||
) -> pl.DataFrame | None:
|
||||
"""各交易所基准指数的「今日」N 日动量 (盘中实时外推)。
|
||||
"""各板块基准指数的「今日」N 日动量 (盘中实时外推)。
|
||||
|
||||
基准日K parquet 盘中不含今日, 今日基准收盘 = 昨收 × (1 + 实时涨跌)。
|
||||
N 日动量 = 今日基准收盘 / N 个交易日前的收盘 - 1; 交易所与
|
||||
N 日动量 = 今日基准收盘 / N 个交易日前的收盘 - 1; 板块与
|
||||
load_benchmark_momentum 的选基逻辑一致 (同一 TTL 缓存帧)。
|
||||
返回小表: bench_exchange, bench_mom3d, bench_mom10d, bench_mom30d。
|
||||
返回小表: bench_key, bench_mom3d, bench_mom10d, bench_mom30d。
|
||||
无基准数据时 None。
|
||||
"""
|
||||
bench = load_benchmark_momentum(data_dir)
|
||||
@@ -1219,15 +1278,15 @@ def benchmark_momentum_today(
|
||||
if bench.is_empty():
|
||||
return None
|
||||
rows: list[dict[str, float | str]] = []
|
||||
for ex in sorted(bench["bench_exchange"].unique().to_list()):
|
||||
sub = bench.filter(pl.col("bench_exchange") == ex).sort("date")
|
||||
for k in sorted(bench["bench_key"].unique().to_list()):
|
||||
sub = bench.filter(pl.col("bench_key") == k).sort("date")
|
||||
closes = sub["bench_close"]
|
||||
if closes.len() == 0:
|
||||
continue
|
||||
yesterday_close = closes[-1]
|
||||
rt = _bench_rt_pct_of(index_quotes, _BENCHMARK_PREFERENCE.get(ex, []))
|
||||
rt = _bench_rt_pct_of(index_quotes, _BENCHMARK_PREFERENCE.get(k, []))
|
||||
row: dict[str, float | str] = {
|
||||
"bench_exchange": ex,
|
||||
"bench_key": k,
|
||||
}
|
||||
for n in DEVIATION_WINDOWS:
|
||||
base = closes[-n] if closes.len() >= n else None # N 个交易日前 (不含今日)
|
||||
@@ -1239,7 +1298,7 @@ def benchmark_momentum_today(
|
||||
rows.append(row)
|
||||
if not rows:
|
||||
return None
|
||||
schema = {"bench_exchange": pl.Utf8, **{f"bench_mom{n}d": pl.Float64 for n in DEVIATION_WINDOWS}}
|
||||
schema = {"bench_key": pl.Utf8, **{f"bench_mom{n}d": pl.Float64 for n in DEVIATION_WINDOWS}}
|
||||
return pl.DataFrame(rows, schema=schema)
|
||||
|
||||
|
||||
@@ -1270,13 +1329,156 @@ def attach_deviation_columns_today(
|
||||
for n in DEVIATION_WINDOWS
|
||||
]
|
||||
return (
|
||||
df.with_columns(_bench_exchange_expr().alias("_bench_ex"))
|
||||
.join(bench, left_on="_bench_ex", right_on="bench_exchange", how="left")
|
||||
df.with_columns(_bench_key_expr().alias("_bench_ex"))
|
||||
.join(bench, left_on="_bench_ex", right_on="bench_key", how="left")
|
||||
.with_columns(exprs)
|
||||
.drop(["_bench_ex", *[f"bench_mom{n}d" for n in DEVIATION_WINDOWS]])
|
||||
)
|
||||
|
||||
|
||||
# ================================================================
|
||||
# 全量重建流式暂存 + 自适应批次 (#208/#174)
|
||||
#
|
||||
# 旧全量模式把所有批次结果累积在内存 date_buffers 直到统一写盘:
|
||||
# 延长历史后 (5年 × 5500 只 ≈ 800 万行) 「全表驻留 + 单批宽表」双双
|
||||
# 超出小内存机器上限, 重建必然 OOM。现改为:
|
||||
# - 每批结果立即写暂存文件 (enriched 树外的隐藏目录 —— polars/duckdb
|
||||
# 的 **/*.parquet glob 均会匹配点目录, 树内暂存会被业务读取扫到),
|
||||
# 最后按日期分块流式合并、逐分区原子替换;
|
||||
# - 批次大小按单批目标行数自适应收缩 (指标/信号全部 over("symbol")
|
||||
# 分组, symbol 级分批不改变计算结果, 只约束单批宽表峰值)。
|
||||
# 任一时刻峰值内存 = 单批计算 + 单个日期块合并, 与总历史长度无关。
|
||||
# ================================================================
|
||||
|
||||
_STAGING_ROOT = Path(".staging") / "enriched_rebuild"
|
||||
_STALE_STAGING_MAX_AGE_S = 24 * 3600
|
||||
_RAM_LARGE_BYTES = 8 * 1024 ** 3 # ≥8GB 视为内存充裕, 批次保持用户设置
|
||||
_BATCH_TARGET_ROWS = 150_000 # 小内存单批目标行数 (宽表 ~60-80MB)
|
||||
_BATCH_MIN_SYMBOLS = 50
|
||||
_MERGE_DATE_CHUNKS = 15 # 最终合并按日期切 15 块流式执行
|
||||
|
||||
_ram_bytes_cache: int | None | bool = False # False = 未探测
|
||||
|
||||
|
||||
def _total_ram_bytes() -> int | None:
|
||||
global _ram_bytes_cache
|
||||
if _ram_bytes_cache is False:
|
||||
try:
|
||||
import psutil
|
||||
_ram_bytes_cache = psutil.virtual_memory().total
|
||||
except Exception:
|
||||
_ram_bytes_cache = None
|
||||
return _ram_bytes_cache # type: ignore[return-value]
|
||||
|
||||
|
||||
def _adaptive_sym_batch(default_batch: int, rows_per_symbol: int) -> int:
|
||||
"""小内存机器按单批目标行数收缩批次; 大内存机器保持原值 (#208)。"""
|
||||
if (_total_ram_bytes() or 0) >= _RAM_LARGE_BYTES:
|
||||
return default_batch
|
||||
return max(
|
||||
_BATCH_MIN_SYMBOLS,
|
||||
min(default_batch, _BATCH_TARGET_ROWS // max(rows_per_symbol, 1)),
|
||||
)
|
||||
|
||||
|
||||
def _sweep_stale_staging(data_dir: Path) -> None:
|
||||
"""清理崩溃/取消运行残留的暂存目录 (按 mtime 判定, 不碰活跃目录)。"""
|
||||
root = data_dir / _STAGING_ROOT
|
||||
if not root.exists():
|
||||
return
|
||||
cutoff = time.time() - _STALE_STAGING_MAX_AGE_S
|
||||
for run_dir in root.iterdir():
|
||||
try:
|
||||
if run_dir.is_dir() and run_dir.stat().st_mtime < cutoff:
|
||||
shutil.rmtree(run_dir, ignore_errors=True)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def compute_enriched_history_window(
|
||||
df_hist: pl.DataFrame,
|
||||
data_dir: Path,
|
||||
instruments: pl.DataFrame | None = None,
|
||||
historical_shares: pl.DataFrame | None = None,
|
||||
sym_batch: int | None = None,
|
||||
*,
|
||||
include_instrument_metadata: bool = False,
|
||||
) -> pl.DataFrame:
|
||||
"""按 symbol 分批执行历史窗口计算: 指标 → 偏离列 → 信号 → 涨跌停。
|
||||
|
||||
分批约束计算中间表, 最终完整历史仍常驻内存。排序和可选元数据关联
|
||||
均在单批完成, 避免合并后再复制整张宽表。
|
||||
sym_batch 显式传入时跳过自适应 (测试用)。
|
||||
"""
|
||||
if df_hist.is_empty() or "symbol" not in df_hist.columns:
|
||||
return df_hist
|
||||
symbols = df_hist["symbol"].unique().sort().to_list()
|
||||
if sym_batch is None:
|
||||
rows_per_sym = max(1, df_hist.height // max(len(symbols), 1))
|
||||
sym_batch = _adaptive_sym_batch(2000, rows_per_sym)
|
||||
parts: list[pl.DataFrame] = []
|
||||
for bs in range(0, len(symbols), sym_batch):
|
||||
batch = symbols[bs:bs + sym_batch]
|
||||
part = df_hist.filter(pl.col("symbol").is_in(batch)).sort(["symbol", "date"])
|
||||
part = compute_indicators(part)
|
||||
part = attach_deviation_columns(part, data_dir)
|
||||
part = compute_signals(part)
|
||||
if instruments is not None and not instruments.is_empty():
|
||||
inst_batch = instruments.filter(pl.col("symbol").is_in(batch))
|
||||
shares_batch = (
|
||||
historical_shares.filter(pl.col("symbol").is_in(batch))
|
||||
if historical_shares is not None and not historical_shares.is_empty()
|
||||
else historical_shares
|
||||
)
|
||||
part = compute_limit_signals(part, inst_batch, historical_shares=shares_batch)
|
||||
if include_instrument_metadata:
|
||||
inst_cols = [c for c in ("name", "total_shares", "float_shares")
|
||||
if c in inst_batch.columns and c not in part.columns]
|
||||
if inst_cols:
|
||||
part = part.join(
|
||||
inst_batch.select("symbol", *inst_cols).unique(subset=["symbol"]),
|
||||
on="symbol", how="left",
|
||||
)
|
||||
# 连续、互不重叠的已排序 symbol 批次, 拼接后天然有序。
|
||||
parts.append(part.sort(["symbol", "date"]))
|
||||
return parts[0] if len(parts) == 1 else pl.concat(parts, how="diagonal_relaxed", rechunk=False)
|
||||
|
||||
|
||||
def _compute_storage_batches(
|
||||
raw: pl.DataFrame,
|
||||
*,
|
||||
factors: pl.DataFrame,
|
||||
instruments: pl.DataFrame,
|
||||
historical_shares: pl.DataFrame,
|
||||
) -> pl.DataFrame:
|
||||
"""保留完整标的历史输入, 单批计算宽表后仅累积落盘窄表。"""
|
||||
from app.services import preferences
|
||||
|
||||
if raw.is_empty():
|
||||
return _select_storage_cols(raw)
|
||||
symbols = raw["symbol"].unique().sort().to_list()
|
||||
rows_per_sym = max(1, -(-raw.height // len(symbols)))
|
||||
batch_size = _adaptive_sym_batch(preferences.get_enriched_batch_size(), rows_per_sym)
|
||||
parts = []
|
||||
for start in range(0, len(symbols), batch_size):
|
||||
batch = symbols[start:start + batch_size]
|
||||
part = compute_enriched(
|
||||
raw.filter(pl.col("symbol").is_in(batch)),
|
||||
factors=factors.filter(pl.col("symbol").is_in(batch)) if not factors.is_empty() else factors,
|
||||
instruments=(instruments.filter(pl.col("symbol").is_in(batch))
|
||||
if not instruments.is_empty() else instruments),
|
||||
historical_shares=(historical_shares.filter(pl.col("symbol").is_in(batch))
|
||||
if not historical_shares.is_empty() else historical_shares),
|
||||
)
|
||||
# 下一批开始前释放宽表; 分区发布仍在所有计算批次成功之后。
|
||||
if not part.is_empty():
|
||||
parts.append(_select_storage_cols(part))
|
||||
del part
|
||||
if not parts:
|
||||
return _select_storage_cols(raw.head(0))
|
||||
return pl.concat(parts, how="diagonal_relaxed", rechunk=False)
|
||||
|
||||
|
||||
def run_pipeline(data_dir: Path | None = None,
|
||||
symbols: list[str] | None = None,
|
||||
new_dates_only: bool = False,
|
||||
@@ -1370,7 +1572,7 @@ def run_pipeline(data_dir: Path | None = None,
|
||||
else:
|
||||
raw_full = raw_new
|
||||
|
||||
enriched_new = compute_enriched(
|
||||
enriched_new = _compute_storage_batches(
|
||||
raw_full,
|
||||
factors=factors,
|
||||
instruments=instruments,
|
||||
@@ -1401,6 +1603,7 @@ def run_pipeline(data_dir: Path | None = None,
|
||||
written += date_df.height
|
||||
t_write_new = _t.perf_counter()
|
||||
logger.info("增量写入: %.2fs, %d 行", t_write_new - t_new, written)
|
||||
del raw_new, hist_df, raw_full, enriched_new
|
||||
|
||||
# 3. 受除权因子影响的个股: 重算全部已有日期 (累积因子链变了)
|
||||
if symbols:
|
||||
@@ -1412,7 +1615,7 @@ def run_pipeline(data_dir: Path | None = None,
|
||||
factors_sym = factors.filter(pl.col("symbol").is_in(list(sym_set))) if not factors.is_empty() else factors
|
||||
inst_sym = instruments.filter(pl.col("symbol").is_in(list(sym_set))) if not instruments.is_empty() else instruments
|
||||
shares_sym = historical_shares.filter(pl.col("symbol").is_in(list(sym_set))) if not historical_shares.is_empty() else historical_shares
|
||||
enriched_sym = compute_enriched(
|
||||
enriched_sym = _compute_storage_batches(
|
||||
raw_sym,
|
||||
factors=factors_sym,
|
||||
instruments=inst_sym,
|
||||
@@ -1450,9 +1653,10 @@ def run_pipeline(data_dir: Path | None = None,
|
||||
|
||||
import gc
|
||||
|
||||
# ── 按 symbol 分批处理: 每只股只有 ~244 行, 无冗余计算 ──
|
||||
# 先获取全部 symbol 列表
|
||||
lf_all = scan_daily_parquet(daily_glob, cast_options=_cast)
|
||||
# ── 按 symbol 分批处理: 指标全部 over("symbol") 分组, 分批不改变结果 ──
|
||||
# 文件列表只收集一次, 批间复用 (避免每批重新展开 glob)
|
||||
daily_files = sorted(str(p) for p in daily_dir.rglob("*.parquet"))
|
||||
lf_all = scan_daily_parquet(daily_files, cast_options=_cast)
|
||||
if symbols:
|
||||
sym_set = set(symbols)
|
||||
lf_all = lf_all.filter(pl.col("symbol").is_in(list(sym_set)))
|
||||
@@ -1466,7 +1670,6 @@ def run_pipeline(data_dir: Path | None = None,
|
||||
return 0
|
||||
|
||||
total_syms = len(all_symbols)
|
||||
logger.info("全量计算: %d 只标的, 按 symbol 分批 [%s]", total_syms, mode)
|
||||
|
||||
if not factors.is_empty() and symbols:
|
||||
factors = factors.filter(pl.col("symbol").is_in(list(sym_set)))
|
||||
@@ -1476,106 +1679,139 @@ def run_pipeline(data_dir: Path | None = None,
|
||||
inst_use = instruments.filter(pl.col("symbol").is_in(list(sym_set)))
|
||||
|
||||
from app.services import preferences as prefs_mod
|
||||
SYM_BATCH = prefs_mod.get_enriched_batch_size() # 每批 N 只 × ~244 天, 可在设置中调整
|
||||
# 自适应批次 (#208): 单批体积按目标行数恒定, 与总历史长度解耦;
|
||||
# 小内存机器自动收缩, 大内存机器保持用户设置
|
||||
total_rows = lf_all.select(pl.len()).collect(streaming=True).item()
|
||||
rows_per_sym = max(1, -(-int(total_rows) // total_syms))
|
||||
SYM_BATCH = _adaptive_sym_batch(prefs_mod.get_enriched_batch_size(), rows_per_sym)
|
||||
total_batches = (total_syms + SYM_BATCH - 1) // SYM_BATCH
|
||||
logger.info("全量计算: %d 只标的 (%d 行, ~%d 行/只), symbol 分批 %d 只/批, %d 批 [%s]",
|
||||
total_syms, total_rows, rows_per_sym, SYM_BATCH, total_batches, mode)
|
||||
|
||||
# 全量模式: 收集所有批次结果, 最后按日期分区覆盖写入
|
||||
from collections import defaultdict
|
||||
date_buffers: dict[str, list[pl.DataFrame]] = defaultdict(list)
|
||||
# 全量模式: 流式暂存发布 (#208) —— 每批落盘暂存文件, 不再内存累积;
|
||||
# 暂存目录在 enriched 树外, 不会被任何 **/*.parquet 业务 glob 扫到
|
||||
staging_dir: Path | None = None
|
||||
staging_files: list[str] = []
|
||||
if not symbols:
|
||||
_sweep_stale_staging(d)
|
||||
staging_dir = d / _STAGING_ROOT / uuid.uuid4().hex
|
||||
staging_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for batch_start in range(0, total_syms, SYM_BATCH):
|
||||
batch_end = min(batch_start + SYM_BATCH, total_syms)
|
||||
batch_syms = all_symbols[batch_start:batch_end]
|
||||
try:
|
||||
for batch_start in range(0, total_syms, SYM_BATCH):
|
||||
batch_end = min(batch_start + SYM_BATCH, total_syms)
|
||||
batch_syms = all_symbols[batch_start:batch_end]
|
||||
|
||||
# 只读取本批 symbol 的数据
|
||||
lf_batch = scan_daily_parquet(daily_glob, cast_options=_cast)
|
||||
lf_batch = lf_batch.filter(pl.col("symbol").is_in(batch_syms))
|
||||
raw = lf_batch.sort(["symbol", "date"]).collect(streaming=True)
|
||||
# 只读取本批 symbol 的数据
|
||||
lf_batch = scan_daily_parquet(daily_files, cast_options=_cast)
|
||||
lf_batch = lf_batch.filter(pl.col("symbol").is_in(batch_syms))
|
||||
raw = lf_batch.sort(["symbol", "date"]).collect(streaming=True)
|
||||
|
||||
if raw.is_empty():
|
||||
continue
|
||||
if raw.is_empty():
|
||||
continue
|
||||
|
||||
# 本批的 factors / instruments
|
||||
batch_factors = (
|
||||
factors.filter(pl.col("symbol").is_in(batch_syms))
|
||||
if not factors.is_empty() else factors
|
||||
)
|
||||
batch_inst = (
|
||||
inst_use.filter(pl.col("symbol").is_in(batch_syms))
|
||||
if not inst_use.is_empty() else inst_use
|
||||
)
|
||||
batch_shares = (
|
||||
historical_shares.filter(pl.col("symbol").is_in(batch_syms))
|
||||
if not historical_shares.is_empty() else historical_shares
|
||||
)
|
||||
# 本批的 factors / instruments
|
||||
batch_factors = (
|
||||
factors.filter(pl.col("symbol").is_in(batch_syms))
|
||||
if not factors.is_empty() else factors
|
||||
)
|
||||
batch_inst = (
|
||||
inst_use.filter(pl.col("symbol").is_in(batch_syms))
|
||||
if not inst_use.is_empty() else inst_use
|
||||
)
|
||||
batch_shares = (
|
||||
historical_shares.filter(pl.col("symbol").is_in(batch_syms))
|
||||
if not historical_shares.is_empty() else historical_shares
|
||||
)
|
||||
|
||||
# 计算
|
||||
enriched = compute_enriched(
|
||||
raw,
|
||||
factors=batch_factors,
|
||||
instruments=batch_inst,
|
||||
historical_shares=batch_shares,
|
||||
)
|
||||
# 计算
|
||||
enriched = compute_enriched(
|
||||
raw,
|
||||
factors=batch_factors,
|
||||
instruments=batch_inst,
|
||||
historical_shares=batch_shares,
|
||||
)
|
||||
|
||||
if not enriched.is_empty():
|
||||
if symbols:
|
||||
# 局部模式: 直接按日期合并写入
|
||||
for date_df in enriched.partition_by("date"):
|
||||
dt = date_df["date"][0]
|
||||
ds = dt.isoformat() if hasattr(dt, "isoformat") else str(dt)
|
||||
out = base / f"date={ds}" / "part.parquet"
|
||||
if not enriched.is_empty():
|
||||
if symbols:
|
||||
# 局部模式: 直接按日期合并写入
|
||||
for date_df in _select_storage_cols(enriched).partition_by("date"):
|
||||
dt = date_df["date"][0]
|
||||
ds = dt.isoformat() if hasattr(dt, "isoformat") else str(dt)
|
||||
out = base / f"date={ds}" / "part.parquet"
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
date_df_storage = _select_storage_cols(date_df)
|
||||
if out.exists():
|
||||
existing = pl.read_parquet(out)
|
||||
existing = existing.filter(~pl.col("symbol").is_in(batch_syms))
|
||||
date_df_storage = pl.concat([existing, date_df_storage], how="diagonal_relaxed")
|
||||
date_df_storage = date_df_storage.sort(["symbol"])
|
||||
publication.write_parquet(date_df_storage, out)
|
||||
written += date_df_storage.height
|
||||
else:
|
||||
# 全量模式: 写单批暂存文件 (按 date,symbol 排序 →
|
||||
# 合并期 parquet 行组统计可按日期裁剪), 随即释放本批内存
|
||||
out = staging_dir / f"batch-{batch_start // SYM_BATCH:04d}.parquet"
|
||||
_select_storage_cols(enriched).sort(["date", "symbol"]).write_parquet(out)
|
||||
staging_files.append(str(out))
|
||||
written += enriched.height
|
||||
|
||||
del raw, enriched, batch_factors, batch_inst, batch_shares
|
||||
gc.collect()
|
||||
|
||||
logger.info("symbol 批次 %d/%d (%s ~ %s), 已处理 %d 行",
|
||||
batch_start // SYM_BATCH + 1,
|
||||
total_batches,
|
||||
batch_syms[0], batch_syms[-1], written)
|
||||
|
||||
# 通知进度
|
||||
if on_batch_done:
|
||||
on_batch_done(batch_start // SYM_BATCH + 1, total_batches)
|
||||
|
||||
# 全量模式: 日期覆盖校验 → 按日期分块流式合并 → 逐分区原子替换
|
||||
if not symbols and staging_files:
|
||||
existing_dates = {
|
||||
p.name.removeprefix("date=")
|
||||
for p in base.glob("date=*")
|
||||
if p.is_dir()
|
||||
}
|
||||
unique_dates = sorted(
|
||||
scan_enriched_parquet(staging_files).select("date").unique()
|
||||
.collect()["date"].to_list()
|
||||
)
|
||||
rebuilt_dates = {
|
||||
ds.isoformat() if hasattr(ds, "isoformat") else str(ds)
|
||||
for ds in unique_dates
|
||||
}
|
||||
missing_dates = existing_dates - rebuilt_dates
|
||||
if missing_dates:
|
||||
sample = ", ".join(sorted(missing_dates)[:5])
|
||||
raise RuntimeError(f"全量重建结果缺少已有日期分区,拒绝覆盖: {sample}")
|
||||
|
||||
base.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
chunk = max(1, -(-len(unique_dates) // _MERGE_DATE_CHUNKS))
|
||||
for ci in range(0, len(unique_dates), chunk):
|
||||
lo = unique_dates[ci]
|
||||
hi = unique_dates[min(ci + chunk, len(unique_dates)) - 1]
|
||||
block = (
|
||||
scan_enriched_parquet(staging_files)
|
||||
.filter((pl.col("date") >= lo) & (pl.col("date") <= hi))
|
||||
.sort(["date", "symbol"])
|
||||
.collect(streaming=True)
|
||||
)
|
||||
for date_df in block.partition_by("date"):
|
||||
ds = date_df["date"][0]
|
||||
ds_str = ds.isoformat() if hasattr(ds, "isoformat") else str(ds)
|
||||
out = base / f"date={ds_str}" / "part.parquet"
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
date_df_storage = _select_storage_cols(date_df)
|
||||
if out.exists():
|
||||
existing = pl.read_parquet(out)
|
||||
existing = existing.filter(~pl.col("symbol").is_in(batch_syms))
|
||||
date_df_storage = pl.concat([existing, date_df_storage], how="diagonal_relaxed")
|
||||
date_df_storage = date_df_storage.sort(["symbol"])
|
||||
publication.write_parquet(date_df_storage, out)
|
||||
written += date_df_storage.height
|
||||
else:
|
||||
# 全量模式: 缓冲到 date_buffers, 最后一次性写入
|
||||
for date_df in enriched.partition_by("date"):
|
||||
dt = date_df["date"][0]
|
||||
ds = dt.isoformat() if hasattr(dt, "isoformat") else str(dt)
|
||||
date_buffers[ds].append(_select_storage_cols(date_df).sort(["symbol"]))
|
||||
written += date_df.height
|
||||
|
||||
del raw, enriched, batch_factors, batch_inst, batch_shares
|
||||
gc.collect()
|
||||
|
||||
logger.info("symbol 批次 %d/%d (%s ~ %s), 已处理 %d 行",
|
||||
batch_start // SYM_BATCH + 1,
|
||||
total_batches,
|
||||
batch_syms[0], batch_syms[-1], written)
|
||||
|
||||
# 通知进度
|
||||
if on_batch_done:
|
||||
on_batch_done(batch_start // SYM_BATCH + 1, total_batches)
|
||||
|
||||
# 全量模式: 按日期分区写入
|
||||
if not symbols and date_buffers:
|
||||
existing_dates = {
|
||||
p.name.removeprefix("date=")
|
||||
for p in base.glob("date=*")
|
||||
if p.is_dir()
|
||||
}
|
||||
rebuilt_dates = set(date_buffers)
|
||||
missing_dates = existing_dates - rebuilt_dates
|
||||
if missing_dates:
|
||||
sample = ", ".join(sorted(missing_dates)[:5])
|
||||
raise RuntimeError(f"全量重建结果缺少已有日期分区,拒绝覆盖: {sample}")
|
||||
|
||||
base.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for ds, dfs in date_buffers.items():
|
||||
out = base / f"date={ds}" / "part.parquet"
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
merged = pl.concat(dfs, how="diagonal_relaxed").sort(["symbol"])
|
||||
publication.write_parquet(merged, out)
|
||||
|
||||
date_buffers.clear()
|
||||
gc.collect()
|
||||
publication.write_parquet(date_df.sort(["symbol"]), out)
|
||||
gc.collect()
|
||||
logger.info("全量暂存合并完成: %d 个日期分区", len(unique_dates))
|
||||
finally:
|
||||
# 无论成功/失败/取消都清掉本次暂存 (历史残留由 _sweep_stale_staging 兜底)
|
||||
if staging_dir is not None:
|
||||
shutil.rmtree(staging_dir, ignore_errors=True)
|
||||
|
||||
publication.commit()
|
||||
t_done = _t.perf_counter()
|
||||
@@ -1949,6 +2185,12 @@ def compute_enriched_today(
|
||||
]
|
||||
df = df.drop([c for c in drop_cols if c in df.columns])
|
||||
|
||||
# 扩展表数值列注入: 当日单日帧, 时序按当日分区对齐 + 快照最新值
|
||||
# (include_snapshot 仅此处为 True —— 单日帧不存在"回看历史"的未来函数问题)。
|
||||
# 帧缓存由 ext_factors 按分区/文件签名管理, 写入端变更自动失效。
|
||||
from app.factors import ext_factors
|
||||
df = ext_factors.attach_ext_columns(df, include_snapshot=True)
|
||||
|
||||
# 自定义信号(日级实时路径同样注入, 但不支持日期偏移条件 → allow_shift=False)
|
||||
# 复用模块级缓存 _custom_signal_exprs_today: 增量热路径每秒级执行,
|
||||
# 不缓存则每轮 glob + 读所有 JSON + 重编译表达式。失效由 invalidate_custom_signals 统一管理。
|
||||
|
||||
@@ -2,7 +2,8 @@
|
||||
|
||||
调度:
|
||||
09:10 盘前 — 同步个股维表 instruments (全量覆盖)
|
||||
15:30 盘后 — 日K同步 + 增量除权因子 + enriched 计算 + 刷新视图
|
||||
15:35 盘后 — 日K同步 + 增量除权因子 + enriched 计算 + 刷新视图
|
||||
(默认 15:35: 盘后固定价 15:30 终止 + 供应商日线定稿缓冲, 见 preferences)
|
||||
|
||||
盘后同步策略:
|
||||
日 K: QuoteService 交易时段已实时落盘 → 有数据时跳过 batch,首次拉 1 年区间
|
||||
@@ -19,9 +20,10 @@ from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from apscheduler.triggers.cron import CronTrigger
|
||||
from apscheduler.triggers.interval import IntervalTrigger
|
||||
|
||||
from app.indicators.pipeline import run_pipeline
|
||||
from app.config import settings
|
||||
from app.services import index_sync, instrument_sync, kline_sync, preferences as _prefs
|
||||
from app.indicators.pipeline import filter_halt_days, run_pipeline
|
||||
from app.services import index_sync, instrument_sync, kline_sync
|
||||
from app.services import preferences as _prefs
|
||||
from app.tickflow.capabilities import Cap, CapabilitySet
|
||||
from app.tickflow.pools import DEMO_SYMBOLS, get_pool
|
||||
from app.tickflow.repository import KlineRepository
|
||||
@@ -31,6 +33,87 @@ logger = logging.getLogger(__name__)
|
||||
ProgressCb = Callable[..., None]
|
||||
|
||||
|
||||
def _prune_partial_enriched_partitions(daily_dir: Path, enriched_dir: Path) -> list[str]:
|
||||
"""删除 symbol 覆盖不完整的 enriched 日期分区, 返回被删的日期 (#223)。
|
||||
|
||||
自选实时路径会在全市场 enriched 生成前提前创建当日分区 (只有几只自选),
|
||||
仅按日期目录计数比较会把它误判为完整分区而跳过计算, 造成日K缺失与
|
||||
均线错误。按与加工相同的停牌过滤口径检查 symbol 覆盖, 不能直接比较
|
||||
行数: 正常剔除停牌记录会让 enriched 少行, 导致每次管道都删除重算。
|
||||
删除后 run_pipeline(new_dates_only=True) 会把它们当"新日期"全市场补齐。
|
||||
daily 同日分区不存在 (今日日K尚未同步) 时不处理, 留给当日正常流程。
|
||||
"""
|
||||
import shutil
|
||||
|
||||
pruned: list[str] = []
|
||||
for part in enriched_dir.glob("date=*"):
|
||||
daily_part = daily_dir / part.stem
|
||||
if not daily_part.exists():
|
||||
continue
|
||||
try:
|
||||
expected: set[str] = set()
|
||||
# 每次只读单文件的停牌判定列, 不加载全历史或指标宽表。
|
||||
for path in daily_part.glob("*.parquet"):
|
||||
schema = pl.read_parquet_schema(path)
|
||||
if not {"symbol", "open", "high"}.issubset(schema):
|
||||
raise ValueError("daily 缺少 symbol/open/high, 无法判断有效标的覆盖")
|
||||
columns = [c for c in ("symbol", "open", "high", "volume", "amount") if c in schema]
|
||||
daily = pl.read_parquet(path, columns=columns)
|
||||
expected.update(filter_halt_days(daily)["symbol"].drop_nulls().to_list())
|
||||
actual: set[str] = set()
|
||||
for path in part.glob("*.parquet"):
|
||||
actual.update(pl.read_parquet(path, columns=["symbol"])["symbol"].drop_nulls().to_list())
|
||||
except Exception as e:
|
||||
logger.warning("enriched 覆盖检查跳过 %s, 保留分区: %s", part.name, e)
|
||||
continue
|
||||
if expected - actual:
|
||||
shutil.rmtree(part, ignore_errors=True)
|
||||
pruned.append(part.stem.split("=")[1])
|
||||
return pruned
|
||||
|
||||
|
||||
def _prune_stale_price_partitions(
|
||||
daily_dir: Path, enriched_dir: Path, max_dates: int = 5
|
||||
) -> list[str]:
|
||||
"""删除收盘价与官方日线不一致的 enriched 日期分区。
|
||||
|
||||
实时 flush 写入的当日分区行数与 daily 相同, 但收盘价可能停留在收盘集合
|
||||
竞价前的快照 (实测: TickFlow 实时端点收盘后仍长期返回旧价, 3392/5554 只
|
||||
股票当日收盘价与官方日线不符), #223 的行数校验识别不到。对最近若干交易日
|
||||
做值级比对: enriched.raw_close 与 daily.close 任一标的差超过半个最小报价
|
||||
单位即删分区, 由后续增量重算按官方日线全市场重建。
|
||||
"""
|
||||
import shutil
|
||||
|
||||
common = sorted(
|
||||
(
|
||||
p.stem.split("=", 1)[1]
|
||||
for p in enriched_dir.glob("date=*")
|
||||
if (daily_dir / p.stem).exists()
|
||||
),
|
||||
reverse=True,
|
||||
)[:max_dates]
|
||||
pruned: list[str] = []
|
||||
for ds in common:
|
||||
try:
|
||||
daily = pl.read_parquet(
|
||||
daily_dir / f"date={ds}" / "*.parquet", columns=["symbol", "close"]
|
||||
)
|
||||
enr = pl.read_parquet(
|
||||
enriched_dir / f"date={ds}" / "*.parquet", columns=["symbol", "raw_close"]
|
||||
)
|
||||
except Exception:
|
||||
continue # 列缺失/不可读 → 交给既有完整性检查兜底
|
||||
joined = enr.join(daily, on="symbol", how="inner").drop_nulls()
|
||||
if joined.is_empty():
|
||||
continue
|
||||
bad = joined.filter((pl.col("raw_close") - pl.col("close")).abs() > 0.005)
|
||||
if not bad.is_empty():
|
||||
shutil.rmtree(enriched_dir / f"date={ds}", ignore_errors=True)
|
||||
pruned.append(ds)
|
||||
return pruned
|
||||
|
||||
|
||||
class PipelineStageError(RuntimeError):
|
||||
"""管道有阶段软失败(数据可能陈旧)时抛出, 让上层 job_store 把任务标记为 failed。
|
||||
|
||||
@@ -343,7 +426,7 @@ def run_now(
|
||||
# - 首次 (enriched 目录不存在) → 全量
|
||||
# - 往前扩展历史 (新日期 < enriched 已有最早日期) → 全量
|
||||
# 前面的除权因子会改变累积因子链,影响后面所有日期的复权价格
|
||||
# - 往后新增日期 (新日期 > enriched 已有最晚日期)
|
||||
# - 往后新增日期或已有历史区间内的缺口
|
||||
# → 增量补新区块(所有标的) + 受除权影响个股全日期重算
|
||||
# - 无新日期 + 有新除权因子 → 增量: 只重算受影响个股的全部日期
|
||||
# - 无新日期 + 无变化 → 跳过
|
||||
@@ -353,6 +436,23 @@ def run_now(
|
||||
daily_days = len(list(daily_dir.glob("date=*"))) if daily_dir.exists() else 0
|
||||
prev_enriched_days = len(list(enriched_dir.glob("date=*"))) if enriched_exists else 0
|
||||
|
||||
# 部分分区修复 (#223) + 收盘价过期分区修复: 删除被实时合并提前创建、覆盖不全
|
||||
# 或收盘价停留在竞价前快照的 enriched 分区, 让下方计数比较与增量计算把它们
|
||||
# 重新当新日期处理 (值级比对以官方日线为准, 实时源不纠错也能自愈)
|
||||
if enriched_exists:
|
||||
partial_pruned = _prune_partial_enriched_partitions(daily_dir, enriched_dir)
|
||||
stale_pruned = _prune_stale_price_partitions(daily_dir, enriched_dir)
|
||||
pruned_dates = sorted(set(partial_pruned) | set(stale_pruned))
|
||||
if pruned_dates:
|
||||
logger.warning(
|
||||
"compute_enriched: 发现 %d 个异常 enriched 分区 (覆盖不全 %d / 收盘价过期 %d), "
|
||||
"已删除待重算: %s",
|
||||
len(pruned_dates), len(partial_pruned), len(stale_pruned),
|
||||
", ".join(pruned_dates[:10]),
|
||||
)
|
||||
enriched_exists = enriched_dir.exists() and any(enriched_dir.glob("date=*"))
|
||||
prev_enriched_days = len(list(enriched_dir.glob("date=*"))) if enriched_exists else 0
|
||||
|
||||
# 判断新日期方向: 找 daily 和 enriched 的日期集合做比较
|
||||
forward_incremental = False
|
||||
backward_extension = False
|
||||
@@ -361,15 +461,14 @@ def run_now(
|
||||
daily_dates = sorted(d.stem.split("=")[1] for d in daily_dir.glob("date=*"))
|
||||
enriched_dates = sorted(d.stem.split("=")[1] for d in enriched_dir.glob("date=*"))
|
||||
earliest_enriched = enriched_dates[0]
|
||||
latest_enriched = enriched_dates[-1]
|
||||
new_dates = set(daily_dates) - set(enriched_dates)
|
||||
if new_dates:
|
||||
# 有新日期早于 enriched 最早日期 → 往前扩展
|
||||
if any(d < earliest_enriched for d in new_dates):
|
||||
backward_extension = True
|
||||
# 有新日期晚于 enriched 最晚日期 → 往后新增
|
||||
if any(d > latest_enriched for d in new_dates):
|
||||
forward_incremental = True
|
||||
# 包含中间被删的异常分区; 没有新增末日也必须补算。
|
||||
# 往前扩展仍由下方优先走全量分支。
|
||||
forward_incremental = True
|
||||
|
||||
def _enriched_batch_progress(cur: int, tot: int) -> None:
|
||||
emit("compute_enriched", 65 + int(23 * cur / tot),
|
||||
@@ -747,7 +846,7 @@ def _run_tracked(fn, job_label: str) -> bool:
|
||||
重任务执行槽: 再挡一层僵尸并发(reap 后线程仍活时不得并行写 parquet)。
|
||||
返回 True 仅表示任务已成功并且执行槽已释放。
|
||||
"""
|
||||
from app.services.pipeline_jobs import JobCancelledError, job_store, release_run_slot, try_acquire_run_slot
|
||||
from app.services.pipeline_jobs import JobCancelledError, job_store, release_run_slot, run_with_capacity, try_acquire_run_slot
|
||||
|
||||
job_id, is_new = job_store.create()
|
||||
if not is_new:
|
||||
@@ -764,8 +863,7 @@ def _run_tracked(fn, job_label: str) -> bool:
|
||||
|
||||
succeeded = False
|
||||
try:
|
||||
job_store.start(job_id)
|
||||
result = fn(on_progress=progress)
|
||||
result = run_with_capacity(job_id, lambda: fn(on_progress=progress))
|
||||
job_store.succeed(job_id, result)
|
||||
succeeded = True
|
||||
logger.info("scheduled %s completed: job_id=%s", job_label, job_id)
|
||||
@@ -850,9 +948,11 @@ async def _run_scheduled_review(repo) -> None:
|
||||
quote_service.push_review_event(json.dumps(
|
||||
{"type": "done", "archived": True}, ensure_ascii=False))
|
||||
|
||||
# 推送到飞书(可选): 运行时读取配置, 用户改设置下次触发即生效。
|
||||
# 推送门控: review_push_mode=manual 时定时复盘只归档不推送,
|
||||
# 由用户对当日报告显式确认后才推; auto 时保持既有自动推送行为。
|
||||
# 失败静默降级, 不影响已归档的报告。
|
||||
_maybe_push_review(content, meta)
|
||||
if _prefs.get_review_push_mode() == "auto":
|
||||
_maybe_push_review(content, meta)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.exception("scheduled review failed: %s", e)
|
||||
# 兜底: 异常时通知前端停止「生成中」状态, 避免页面卡在 streaming
|
||||
@@ -932,11 +1032,12 @@ def _maybe_push_review(content: str, meta: dict) -> None:
|
||||
"""复盘报告归档后, 按 review_push_channels 选定的外部工具逐个推送完整报告。
|
||||
|
||||
定时生成与手动生成共用本函数 (手动归档端点 POST /api/market-recap/reports 也会调用)。
|
||||
channels 为空则不推送; 'feishu' 复用监控中心的全局飞书 Webhook 通道。
|
||||
channels 为空则不推送; 复用监控中心的全局外部渠道配置。
|
||||
推送失败静默降级 (Webhook 是辅助通道), 不影响已归档的报告。
|
||||
"""
|
||||
try:
|
||||
from app.services import preferences, webhook_adapter
|
||||
from app import secrets_store
|
||||
from app.services import email_adapter, preferences, webhook_adapter
|
||||
|
||||
channels = preferences.get_review_push_channels()
|
||||
if not channels:
|
||||
@@ -968,6 +1069,33 @@ def _maybe_push_review(content: str, meta: dict) -> None:
|
||||
url, "每日复盘", full_body
|
||||
)
|
||||
logger.info("review push(wecom) %s", "sent" if ok else "failed")
|
||||
elif ch == "custom":
|
||||
url = preferences.get_custom_webhook_url()
|
||||
if not url:
|
||||
logger.info("review push(custom) skipped: webhook not configured")
|
||||
continue
|
||||
ok = webhook_adapter.send_custom(
|
||||
url,
|
||||
"每日复盘",
|
||||
content,
|
||||
"market_review",
|
||||
meta,
|
||||
secrets_store.get_custom_webhook_secret(),
|
||||
)
|
||||
logger.info("review push(custom) %s", "sent" if ok else "failed")
|
||||
elif ch == "email":
|
||||
config = preferences.get_email_smtp_config()
|
||||
if not email_adapter.is_configured(config):
|
||||
logger.info("review push(email) skipped: SMTP not configured")
|
||||
continue
|
||||
email_body = (f"{subtitle}\n\n{content}" if subtitle else content)
|
||||
ok = email_adapter.send_email(
|
||||
config,
|
||||
secrets_store.get_email_smtp_password(),
|
||||
"每日复盘",
|
||||
email_body,
|
||||
)
|
||||
logger.info("review push(email) %s", "sent" if ok else "failed")
|
||||
# 未来更多渠道在此追加分支
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("review push error: %s", e)
|
||||
@@ -999,7 +1127,7 @@ def start_scheduler(repo: KlineRepository, capset: CapabilitySet) -> AsyncIOSche
|
||||
"""启动调度器。
|
||||
|
||||
工作日 09:10 — 同步个股维表
|
||||
工作日 HH:MM — 盘后管道(时间由用户偏好决定,默认 15:30)
|
||||
工作日 HH:MM — 盘后管道(时间由用户偏好决定,默认 15:35)
|
||||
"""
|
||||
from app.services import preferences
|
||||
sched = preferences.get_pipeline_schedule()
|
||||
|
||||
+26
-6
@@ -19,10 +19,12 @@ from app.api import (
|
||||
backtest,
|
||||
data,
|
||||
ext_data,
|
||||
factors,
|
||||
financials,
|
||||
indices,
|
||||
intraday,
|
||||
kline,
|
||||
lots,
|
||||
market_recap,
|
||||
mining,
|
||||
monitor_rules,
|
||||
@@ -103,6 +105,15 @@ async def _application_lifespan(app: FastAPI):
|
||||
repo = KlineRepository(store)
|
||||
app.state.datastore = store
|
||||
app.state.repo = repo
|
||||
# 自定义/复合因子载入注册表 (P3); 单个失败只跳过该因子 (fail-隔离)
|
||||
from app.factors.store import load_into_registry
|
||||
|
||||
try:
|
||||
loaded_factors = load_into_registry(store.data_dir)
|
||||
if loaded_factors:
|
||||
logger.info("custom factors loaded: %s", len(loaded_factors))
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("custom factors load failed: %s", exc)
|
||||
from app.services.mining_manager import MiningJobManager
|
||||
|
||||
mining_manager = MiningJobManager(store.data_dir)
|
||||
@@ -124,11 +135,6 @@ async def _application_lifespan(app: FastAPI):
|
||||
# instruments/index/ETF 仍同步 (毫秒级)。应用立即 ready, 指标算完后自动替换。
|
||||
repo.refresh_cache(background=True)
|
||||
|
||||
# 能力探测
|
||||
capset = detect_capabilities()
|
||||
app.state.capabilities = capset
|
||||
logger.info("ready; %d capabilities active", len(capset.all()))
|
||||
|
||||
# 自定义数据源配置(可选): 失败只记录错误, 不影响 TickFlow 基准路径。
|
||||
try:
|
||||
from app.data_providers import custom as custom_sources
|
||||
@@ -137,6 +143,11 @@ async def _application_lifespan(app: FastAPI):
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("custom data sources init failed: %s", e)
|
||||
|
||||
# 自定义源必须先注册,能力探测才能补充其数据集能力。
|
||||
capset = detect_capabilities()
|
||||
app.state.capabilities = capset
|
||||
logger.info("ready; %d capabilities active", len(capset.all()))
|
||||
|
||||
# 全局行情服务
|
||||
qs = QuoteService()
|
||||
app.state.quote_service = qs
|
||||
@@ -227,6 +238,10 @@ async def _application_lifespan(app: FastAPI):
|
||||
financial_scheduler.start(store.data_dir, capset)
|
||||
app.state.financial_scheduler = financial_scheduler
|
||||
|
||||
# 自愈看门狗: 探测 polars 闸与写锁, 僵死时退出交由 supervisor 拉起 (兜底层)。
|
||||
from app.watchdog import start_watchdog
|
||||
app.state.watchdog = start_watchdog(app.state, repo)
|
||||
|
||||
# 策略引擎
|
||||
from app.strategy.engine import StrategyEngine
|
||||
from app.strategy import config as strategy_config
|
||||
@@ -273,7 +288,7 @@ async def _application_lifespan(app: FastAPI):
|
||||
return
|
||||
|
||||
with shared_heavy_job_limiter.slot(
|
||||
"normal",
|
||||
"exclusive",
|
||||
cancel_event=matrix_prewarm_owner.cancel_event,
|
||||
):
|
||||
result = prewarm_matrix_cache(
|
||||
@@ -344,6 +359,9 @@ async def _application_lifespan(app: FastAPI):
|
||||
yield
|
||||
finally:
|
||||
repo._on_refresh_done = None # noqa: SLF001
|
||||
wd = getattr(app.state, "watchdog", None)
|
||||
if wd:
|
||||
await wd.stop()
|
||||
if not matrix_prewarm_owner.shutdown(timeout=5.0):
|
||||
logger.warning("matrix cache prewarm did not stop within 5 seconds")
|
||||
mmanager = getattr(app.state, "mining_manager", None)
|
||||
@@ -454,6 +472,7 @@ app.include_router(kline.router)
|
||||
app.include_router(watchlist.router)
|
||||
app.include_router(screener.router)
|
||||
app.include_router(backtest.router)
|
||||
app.include_router(factors.router)
|
||||
app.include_router(mining.router)
|
||||
app.include_router(intraday.router)
|
||||
app.include_router(indices.router)
|
||||
@@ -471,6 +490,7 @@ app.include_router(settings_api.router)
|
||||
app.include_router(strategy.router)
|
||||
app.include_router(signals.router)
|
||||
app.include_router(monitor_rules.router)
|
||||
app.include_router(lots.router)
|
||||
app.include_router(alerts.router)
|
||||
app.include_router(rps.router)
|
||||
|
||||
|
||||
@@ -44,6 +44,8 @@ def trading_minutes_elapsed_from_dt(dt: datetime) -> float:
|
||||
- 开盘前 = 0; 午休(11:30-13:00) = 120(保持上午累计); 收盘后 = 240。
|
||||
- 非交易日(周末) = 240 (视作全天, 避免量比被折算成 0)。
|
||||
"""
|
||||
if dt.weekday() >= 5:
|
||||
return float(_TRADING_TOTAL_MINUTES)
|
||||
t = dt.time()
|
||||
if t < _MORNING_START:
|
||||
return 0.0
|
||||
|
||||
@@ -32,12 +32,13 @@ import logging
|
||||
import math
|
||||
import re
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Iterator
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, date, datetime, timedelta
|
||||
from pathlib import Path
|
||||
|
||||
import polars as pl
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
from app.data_providers.normalizer import DAILY_COLS, normalize_daily
|
||||
from app.indicators.pipeline import filter_halt_days
|
||||
@@ -69,6 +70,12 @@ _DAILY10_DUMP_KIND = "daily-k-10d"
|
||||
_DAILY_DUMP_KIND = "daily-k" # 10 年全量日K dump(约 172MB), 深窗口一次下载覆盖全市场
|
||||
_RECENT_DUMP_DAYS = 12 # 窗口跨度 ≤ 此天数时优先走 10d dump(覆盖 ≈10 个交易日)
|
||||
_PREV_CLOSE_BACKDAYS = 30 # 推导因子时向前找"除权日前收盘"的回看天数(容忍长期停牌)
|
||||
_DAILY_DUMP_BATCH_ROWS = 100_000
|
||||
_HIST_SYMBOL_BATCH = 50
|
||||
_DAILY_DUMP_COLUMNS = [
|
||||
"thscode", "adjusted", "date_ms", "open_price", "high_price", "low_price",
|
||||
"close_price", "volume", "turnover",
|
||||
]
|
||||
|
||||
|
||||
def get_api_key() -> str:
|
||||
@@ -401,11 +408,12 @@ class FuyaoProvider:
|
||||
logger.info("扶摇实时行情拉取完成: %d 条(丢弃 %d 行)", len(records), dropped)
|
||||
return records
|
||||
|
||||
def get_realtime_indices(self, symbols: list[str]) -> list[dict]:
|
||||
def get_realtime_indices(self, symbols: list[str]) -> list[dict] | None:
|
||||
"""指数实时快照 → 内部 realtime record (可选插件协议, quote_service 鸭子类型调用)。
|
||||
|
||||
A 股快照不含指数, 指数在扶摇是独立端点; 覆盖沪深交易所指数 + 同花顺板块,
|
||||
无北交所 (未知代码会整批 1002 连坐, .BJ 直接跳过)。失败软返回空列表。
|
||||
无北交所 (未知代码会整批 1002 连坐, .BJ 直接跳过)。失败返回 None,
|
||||
让上层与“成功但无数据”的空列表区分, 保留上轮有效指数缓存。
|
||||
"""
|
||||
wanted = [s for s in symbols if s and not s.upper().endswith(".BJ")]
|
||||
if not wanted:
|
||||
@@ -414,7 +422,7 @@ class FuyaoProvider:
|
||||
rows, server_ts = self._get_client().index_snapshot(wanted)
|
||||
except FuyaoError as e:
|
||||
logger.warning("扶摇指数行情拉取失败: %s", e)
|
||||
return []
|
||||
return None
|
||||
|
||||
fetched_ms = server_ts or int(time.time() * 1000)
|
||||
records = []
|
||||
@@ -443,34 +451,147 @@ class FuyaoProvider:
|
||||
- 兜底: 单标的 historical 接口(窗口早于 dump 覆盖 / dump 不可用; 10 年自动分片,
|
||||
逐标的节流 + 进度回调)。
|
||||
"""
|
||||
chunks = [
|
||||
df
|
||||
for df in self.iter_daily(
|
||||
symbols,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
asset_type=asset_type,
|
||||
on_chunk_done=on_chunk_done,
|
||||
)
|
||||
if not df.is_empty()
|
||||
]
|
||||
return pl.concat(chunks, how="diagonal_relaxed") if chunks else pl.DataFrame()
|
||||
|
||||
def iter_daily(
|
||||
self,
|
||||
symbols: list[str],
|
||||
start_time: datetime | None,
|
||||
end_time: datetime | None,
|
||||
asset_type: str = "stock",
|
||||
on_chunk_done: Callable[[int, int], None] | None = None,
|
||||
) -> Iterator[pl.DataFrame]:
|
||||
"""分批产出日K,供历史同步逐批落盘,避免全市场结果累积在内存。"""
|
||||
if not symbols or asset_type != "stock":
|
||||
return pl.DataFrame()
|
||||
return
|
||||
end_dt = end_time or datetime.now()
|
||||
start_dt = start_time or (end_dt - timedelta(days=365))
|
||||
start_d, end_d = start_dt.date(), end_dt.date()
|
||||
symset = set(symbols)
|
||||
|
||||
if (end_d - start_d).days <= _RECENT_DUMP_DAYS:
|
||||
try:
|
||||
dump = self._ensure_dump(_DAILY10_DUMP_KIND, "daily_k_10d")
|
||||
if _dump_covers(dump, start_d, end_d):
|
||||
df = self._daily_from_dump(dump, set(symbols), start_d, end_d)
|
||||
df = self._daily_from_dump(dump, symset, start_d, end_d)
|
||||
if on_chunk_done:
|
||||
on_chunk_done(1, 1)
|
||||
logger.info("扶摇日K(10d dump)完成: %d 行 [%s ~ %s]", df.height, start_d, end_d)
|
||||
return df
|
||||
logger.info("扶摇 10d dump 未覆盖窗口 [%s ~ %s], 尝试 10 年 dump", start_d, end_d)
|
||||
if not df.is_empty():
|
||||
yield df
|
||||
return
|
||||
except FuyaoError as e:
|
||||
logger.warning("扶摇日K 10d dump 不可用, 尝试 10 年 dump: %s", e)
|
||||
logger.warning("扶摇 10d dump 不可用: %s", e)
|
||||
|
||||
df = self._daily_from_big_dump(set(symbols), start_d, end_d)
|
||||
if df is not None:
|
||||
if on_chunk_done:
|
||||
on_chunk_done(1, 1)
|
||||
logger.info("扶摇日K(10 年 dump)完成: %d 行 [%s ~ %s]", df.height, start_d, end_d)
|
||||
return df
|
||||
df = self._daily_from_api(symbols, start_d, end_d, on_chunk_done)
|
||||
logger.info("扶摇日K(单标的接口)完成: %d 行 [%s ~ %s]", df.height, start_d, end_d)
|
||||
return df
|
||||
dump_info = self._daily_dump_info()
|
||||
sources: list[tuple[str, date, date]] = []
|
||||
if dump_info:
|
||||
_, dump_min, dump_max = dump_info
|
||||
if start_d < dump_min:
|
||||
sources.append(("api", start_d, min(end_d, dump_min - timedelta(days=1))))
|
||||
overlap_start, overlap_end = max(start_d, dump_min), min(end_d, dump_max)
|
||||
if overlap_start <= overlap_end:
|
||||
sources.append(("dump", overlap_start, overlap_end))
|
||||
tail_start = max(start_d, dump_max + timedelta(days=1))
|
||||
if tail_start <= end_d and not _tail_ok(end_d, dump_max):
|
||||
try:
|
||||
ten = self._ensure_dump(_DAILY10_DUMP_KIND, "daily_k_10d")
|
||||
ten_dates = pl.from_epoch(
|
||||
ten["date_ms"].cast(pl.Int64) + _SH_MS, time_unit="ms"
|
||||
).dt.date()
|
||||
ten_min, ten_max = ten_dates.min(), ten_dates.max()
|
||||
except (FuyaoError, KeyError):
|
||||
ten_min = ten_max = None
|
||||
if (
|
||||
ten_min is not None
|
||||
and ten_min <= dump_max + timedelta(days=1)
|
||||
and _tail_ok(end_d, ten_max)
|
||||
):
|
||||
sources.append(("10d", tail_start, end_d))
|
||||
else:
|
||||
# 多年 dump 与请求终点之间存在不可验证的缺口,不能返回半段数据。
|
||||
sources = [("api", start_d, end_d)]
|
||||
else:
|
||||
sources.append(("api", start_d, end_d))
|
||||
|
||||
dump_batch_count = 0
|
||||
if dump_info:
|
||||
dump_rows = pq.ParquetFile(dump_info[0]).metadata.num_rows
|
||||
dump_batch_count = max(
|
||||
1, (dump_rows + _DAILY_DUMP_BATCH_ROWS - 1) // _DAILY_DUMP_BATCH_ROWS
|
||||
)
|
||||
api_batch_count = (len(symbols) + _HIST_SYMBOL_BATCH - 1) // _HIST_SYMBOL_BATCH
|
||||
total = sum(
|
||||
dump_batch_count if kind == "dump" else 1 if kind == "10d" else api_batch_count
|
||||
for kind, _, _ in sources
|
||||
)
|
||||
done = 0
|
||||
for kind, source_start, source_end in sources:
|
||||
if source_start > source_end:
|
||||
continue
|
||||
if kind == "dump":
|
||||
path = dump_info[0] # type: ignore[index]
|
||||
for df in self._iter_big_dump(path, symset, source_start, source_end):
|
||||
done += 1
|
||||
if on_chunk_done:
|
||||
on_chunk_done(done, total)
|
||||
if not df.is_empty():
|
||||
yield df
|
||||
elif kind == "10d":
|
||||
ten = self._ensure_dump(_DAILY10_DUMP_KIND, "daily_k_10d")
|
||||
df = self._daily_from_dump(ten, symset, source_start, source_end)
|
||||
done += 1
|
||||
if on_chunk_done:
|
||||
on_chunk_done(done, total)
|
||||
if not df.is_empty():
|
||||
yield df
|
||||
else:
|
||||
batches = [
|
||||
symbols[i:i + _HIST_SYMBOL_BATCH]
|
||||
for i in range(0, len(symbols), _HIST_SYMBOL_BATCH)
|
||||
]
|
||||
for batch in batches:
|
||||
rows: list[dict] = []
|
||||
for symbol in batch:
|
||||
rows.extend(_kline_rows(
|
||||
symbol,
|
||||
self._historical_bars(symbol, source_start, source_end),
|
||||
))
|
||||
time.sleep(_HIST_INTERVAL_S)
|
||||
df = normalize_daily(rows, source=self.name)
|
||||
done += 1
|
||||
if on_chunk_done:
|
||||
on_chunk_done(done, total)
|
||||
if not df.is_empty():
|
||||
yield df
|
||||
|
||||
def _daily_dump_info(self) -> tuple[Path, date, date] | None:
|
||||
"""返回多年 dump 的路径和覆盖范围,不把大文件读入进程内存。"""
|
||||
path = None
|
||||
for candidate in sorted(_cache_dir().glob("daily_k__*.parquet"), reverse=True):
|
||||
try:
|
||||
dmin, dmax = _dump_date_range(candidate)
|
||||
except Exception:
|
||||
continue
|
||||
if dmin is not None and dmax is not None:
|
||||
return candidate, dmin, dmax
|
||||
try:
|
||||
path = self._ensure_dump_path(_DAILY_DUMP_KIND, "daily_k")
|
||||
dmin, dmax = _dump_date_range(path)
|
||||
except FuyaoError as e:
|
||||
logger.warning("扶摇 10 年 dump 不可用, 回退单标的接口: %s", e)
|
||||
return None
|
||||
return (path, dmin, dmax) if dmin is not None and dmax is not None else None
|
||||
|
||||
def _daily_from_dump(
|
||||
self, dump: pl.DataFrame, symset: set[str], start_d: date, end_d: date
|
||||
@@ -480,6 +601,34 @@ class FuyaoProvider:
|
||||
)
|
||||
return self._map_daily_dump(df, symset, start_d, end_d)
|
||||
|
||||
def _iter_big_dump(
|
||||
self, path: Path, symset: set[str], start_d: date, end_d: date
|
||||
) -> Iterator[pl.DataFrame]:
|
||||
"""按固定 record batch 读取多年 dump,不做单次全量 collect。"""
|
||||
parquet = pq.ParquetFile(path)
|
||||
columns = [name for name in _DAILY_DUMP_COLUMNS if name in parquet.schema.names]
|
||||
for batch in parquet.iter_batches(
|
||||
batch_size=_DAILY_DUMP_BATCH_ROWS,
|
||||
columns=columns,
|
||||
):
|
||||
raw = pl.from_arrow(batch)
|
||||
if raw.is_empty() or "date_ms" not in raw.columns or "thscode" not in raw.columns:
|
||||
yield pl.DataFrame()
|
||||
continue
|
||||
start_ms, end_ms = _ms_of_date(start_d), _ms_of_date(end_d)
|
||||
raw = raw.filter(
|
||||
(pl.col("date_ms") >= start_ms)
|
||||
& (pl.col("date_ms") <= end_ms)
|
||||
& pl.col("thscode").is_in(sorted(symset))
|
||||
)
|
||||
if raw.is_empty():
|
||||
yield pl.DataFrame()
|
||||
continue
|
||||
raw = raw.with_columns(
|
||||
pl.from_epoch(pl.col("date_ms") + _SH_MS, time_unit="ms").dt.date().alias("date")
|
||||
)
|
||||
yield self._map_daily_dump(raw, symset, start_d, end_d)
|
||||
|
||||
def _map_daily_dump(
|
||||
self, df: pl.DataFrame, symset: set[str], start_d: date, end_d: date
|
||||
) -> pl.DataFrame:
|
||||
@@ -510,77 +659,6 @@ class FuyaoProvider:
|
||||
cols = [c for c in DAILY_COLS if c in out.columns]
|
||||
return out.select(cols).sort(["symbol", "date"]) if not out.is_empty() else out.select(cols)
|
||||
|
||||
def _daily_from_big_dump(
|
||||
self, symset: set[str], start_d: date, end_d: date
|
||||
) -> pl.DataFrame | None:
|
||||
"""深窗口主路径: 10 年全量 dump(lazy 按需筛) + 必要时 10d dump 补尾。
|
||||
|
||||
覆盖不了(窗口早于 10 年 / dump 拉取失败)返回 None, 由调用方走单标的接口。
|
||||
"""
|
||||
path = self._ensure_daily_big_dump(start_d)
|
||||
if path is None:
|
||||
return None
|
||||
_, dmax = _dump_date_range(path)
|
||||
big_hi = min(end_d, dmax)
|
||||
# 窗口/标的过滤下推到 lazy 计划, 只物化需要的行(全量 10 年 ≈ 13.6M 行)
|
||||
window = (
|
||||
pl.scan_parquet(path)
|
||||
.with_columns(
|
||||
pl.from_epoch(pl.col("date_ms") + _SH_MS, time_unit="ms").dt.date().alias("date")
|
||||
)
|
||||
.filter(
|
||||
(pl.col("date") >= start_d)
|
||||
& (pl.col("date") <= big_hi)
|
||||
& pl.col("thscode").is_in(sorted(symset))
|
||||
)
|
||||
.collect()
|
||||
)
|
||||
parts = [self._map_daily_dump(window, symset, start_d, big_hi)]
|
||||
if not _tail_ok(end_d, dmax):
|
||||
# 末端缺口(如 10 年 dump 是旧 release, end 是最近交易日): 10d dump 补尾
|
||||
try:
|
||||
ten = self._ensure_dump(_DAILY10_DUMP_KIND, "daily_k_10d")
|
||||
ten_dates = pl.from_epoch(
|
||||
ten["date_ms"].cast(pl.Int64) + _SH_MS, time_unit="ms"
|
||||
).dt.date()
|
||||
ten_min, ten_max = ten_dates.min(), ten_dates.max()
|
||||
if (
|
||||
ten_min is not None
|
||||
and ten_min <= dmax + timedelta(days=1)
|
||||
and _tail_ok(end_d, ten_max)
|
||||
):
|
||||
tail_start = max(start_d, dmax + timedelta(days=1))
|
||||
parts.append(self._daily_from_dump(ten, symset, tail_start, end_d))
|
||||
else:
|
||||
return None # 中段或尾部仍有缺口 → 单标的兜底, 不交缺口数据
|
||||
except FuyaoError as e:
|
||||
logger.warning("扶摇 10d dump 补尾失败: %s", e)
|
||||
return None
|
||||
non_empty = [p for p in parts if not p.is_empty()]
|
||||
if not non_empty:
|
||||
return pl.DataFrame()
|
||||
out = pl.concat(non_empty, how="vertical_relaxed")
|
||||
return out.unique(subset=["symbol", "date"], keep="last").sort(["symbol", "date"])
|
||||
|
||||
def _daily_from_api(
|
||||
self,
|
||||
symbols: list[str],
|
||||
start_d: date,
|
||||
end_d: date,
|
||||
on_chunk_done: Callable[[int, int], None] | None,
|
||||
) -> pl.DataFrame:
|
||||
frames: list[pl.DataFrame] = []
|
||||
for i, sym in enumerate(symbols):
|
||||
rows = self._historical_bars(sym, start_d, end_d)
|
||||
time.sleep(_HIST_INTERVAL_S)
|
||||
if rows:
|
||||
df = normalize_daily(_kline_rows(sym, rows), default_symbol=sym, source=self.name)
|
||||
if not df.is_empty():
|
||||
frames.append(df)
|
||||
if on_chunk_done:
|
||||
on_chunk_done(i + 1, len(symbols))
|
||||
return pl.concat(frames, how="diagonal_relaxed") if frames else pl.DataFrame()
|
||||
|
||||
def _historical_bars(self, symbol: str, start_d: date, end_d: date) -> list[dict]:
|
||||
"""按 ≤10 年窗口分片拉取单标的原始日K。中途失败软返回已得行, 不抛出。"""
|
||||
out: list[dict] = []
|
||||
|
||||
@@ -213,6 +213,7 @@ async function opRealtime(sdk, job) {
|
||||
volume: q.volume,
|
||||
amount: q.amount,
|
||||
change_pct: q.changePercent,
|
||||
timestamp: q.timestamp,
|
||||
})
|
||||
}
|
||||
return rows
|
||||
|
||||
@@ -275,7 +275,19 @@ class StockSDKProvider:
|
||||
except bridge.StockSDKBridgeError as e:
|
||||
logger.warning("stock-sdk realtime 拉取失败: %s", e)
|
||||
return []
|
||||
return result.get("rows") or []
|
||||
rows = result.get("rows") or []
|
||||
normalized: list[dict] = []
|
||||
for row in rows:
|
||||
item = dict(row)
|
||||
# stock-sdk 的 changePercent 是百分数值(-1.15 = -1.15%);
|
||||
# provider 入口契约统一使用小数制(-0.0115 = -1.15%)。
|
||||
if item.get("change_pct") is not None:
|
||||
item["change_pct"] = float(item["change_pct"]) / 100
|
||||
# stock-sdk 全量实时行情的 amount 单位为万元;内部日K统一使用元。
|
||||
if item.get("amount") is not None:
|
||||
item["amount"] = float(item["amount"]) * 10_000
|
||||
normalized.append(item)
|
||||
return normalized
|
||||
|
||||
# ---- instruments (标的维表) ----
|
||||
def get_instruments(self, asset_type: str = "stock") -> list[dict]:
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
"""polars collect 并发闸。
|
||||
|
||||
polars 的共享执行器 (rayon 工作池 + 流式引擎异步运行时) 在多线程并发 collect
|
||||
时存在死锁问题: 上游 issue #24448 / #23053 / #25754 等同族案例均为「在飞的
|
||||
collect 超过池内工作位 → 持有工作位的任务等待排不上队的任务 → 0 CPU 永久
|
||||
挂起」。本模块用进程级信号量限制同时在飞的 collect 数量 —— 这是上游 issue
|
||||
区被反复验证有效的缓解手段。
|
||||
|
||||
车道设计: 总闸位 polars_collect_permits 个, 其中 background (预热 / 增量 /
|
||||
维表加载等后台计算) 最多占 polars_collect_background_permits 个, 其余闸位
|
||||
保留给 interactive (页面读接口), 保证后台大计算不会把页面请求饿死。获取顺序
|
||||
恒为 background 车道 → 总闸, 不存在环。
|
||||
|
||||
worker 子进程 (回测/优化/挖掘) 单任务串行执行, 不经过本闸。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from typing import Literal
|
||||
|
||||
import polars as pl
|
||||
|
||||
from app.config import settings
|
||||
|
||||
CollectPriority = Literal["interactive", "background"]
|
||||
|
||||
_TOTAL_GATE = threading.BoundedSemaphore(settings.polars_collect_permits)
|
||||
_BACKGROUND_LANE = threading.BoundedSemaphore(settings.polars_collect_background_permits)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def collect_slot(priority: CollectPriority = "interactive") -> Iterator[None]:
|
||||
"""占用一个 collect 闸位; background 需同时占用车道位与总闸位。"""
|
||||
if priority == "background":
|
||||
with _BACKGROUND_LANE, _TOTAL_GATE:
|
||||
yield
|
||||
return
|
||||
with _TOTAL_GATE:
|
||||
yield
|
||||
|
||||
|
||||
def guarded_collect(
|
||||
lf: pl.LazyFrame,
|
||||
*,
|
||||
priority: CollectPriority = "interactive",
|
||||
**kwargs: object,
|
||||
) -> pl.DataFrame:
|
||||
"""在并发闸内执行 LazyFrame.collect; 语义不变, 仅串行化调度。"""
|
||||
with collect_slot(priority):
|
||||
return lf.collect(**kwargs)
|
||||
@@ -99,6 +99,36 @@ def get_ai_config_int(key: str, default: int) -> int:
|
||||
return int(getattr(settings, key, default) or default)
|
||||
|
||||
|
||||
def get_custom_webhook_secret() -> str:
|
||||
"""Return the optional HMAC secret for the generic outbound webhook."""
|
||||
return str(load().get("custom_webhook_secret") or "")
|
||||
|
||||
|
||||
def set_custom_webhook_secret(secret: str) -> str:
|
||||
"""Persist or clear the generic outbound webhook HMAC secret."""
|
||||
value = (secret or "").strip()
|
||||
if value:
|
||||
save({"custom_webhook_secret": value})
|
||||
else:
|
||||
clear("custom_webhook_secret")
|
||||
return value
|
||||
|
||||
|
||||
def get_email_smtp_password() -> str:
|
||||
"""Return the SMTP password used by the email notification channel."""
|
||||
return str(load().get("email_smtp_password") or "")
|
||||
|
||||
|
||||
def set_email_smtp_password(password: str) -> str:
|
||||
"""Persist or clear the SMTP password used by email notifications."""
|
||||
value = password or ""
|
||||
if value:
|
||||
save({"email_smtp_password": value})
|
||||
else:
|
||||
clear("email_smtp_password")
|
||||
return value
|
||||
|
||||
|
||||
def get_env_backed_secret(field: str, env_name: str) -> str:
|
||||
"""取环境变量后备的密钥(插件 API Key 等):secrets.json 优先,否则环境变量。
|
||||
|
||||
|
||||
@@ -27,7 +27,7 @@ from typing import Any
|
||||
|
||||
import polars as pl
|
||||
|
||||
from app.indicators.pipeline import DEVIATION_WINDOWS
|
||||
from app.indicators.pipeline import BENCH_KEYS, DEVIATION_WINDOWS, bench_rt_pct_for
|
||||
|
||||
# ── 规则表 ────────────────────────────────────────────────
|
||||
|
||||
@@ -57,6 +57,18 @@ RULES_META: list[dict[str, Any]] = [
|
||||
_BENCH_RT_CANDIDATES = ["000002.SH", "000001.SH", "399107.SZ", "399001.SZ", "899050.BJ"]
|
||||
|
||||
|
||||
def _bench_key_of(symbol: str) -> str:
|
||||
"""symbol → 板块基准键, 与 pipeline._bench_key_expr 同口径 (SH/STAR/SZ/GEM/BJ)。"""
|
||||
code = symbol.split(".")[0]
|
||||
if symbol.endswith(".BJ"):
|
||||
return "BJ"
|
||||
if symbol.endswith(".SH"):
|
||||
return "STAR" if code.startswith("68") else "SH"
|
||||
if symbol.endswith(".SZ"):
|
||||
return "GEM" if code.startswith("30") else "SZ"
|
||||
return ""
|
||||
|
||||
|
||||
def board_of(symbol: str) -> str:
|
||||
"""按代码前缀判定板块。"""
|
||||
code = symbol.split(".")[0]
|
||||
@@ -134,7 +146,12 @@ def _hist_snapshot(repo: Any) -> dict[str, Any]:
|
||||
|
||||
|
||||
def _bench_rt_pct(quote_service: Any) -> float:
|
||||
"""基准指数今日实时涨跌 (各候选均值, 缺数据时 0)。"""
|
||||
"""基准指数今日实时涨跌 (各候选均值, 小数制, 缺数据时 0)。
|
||||
|
||||
quote_service.get_index_quotes() 返回指数展示缓存, change_pct/pct/pct_change
|
||||
为百分数口径 (CONTRIBUTING §3.1), 消费前显式 /100, 与 enriched 侧小数制
|
||||
change_pct 对齐 (#232); close/prev_close 兜底路径本身是小数, 不转换。
|
||||
"""
|
||||
try:
|
||||
df = quote_service.get_index_quotes()
|
||||
except Exception:
|
||||
@@ -148,7 +165,7 @@ def _bench_rt_pct(quote_service: Any) -> float:
|
||||
if col in df.columns:
|
||||
vals = df[col].drop_nulls()
|
||||
if vals.len() > 0:
|
||||
return float(vals.mean())
|
||||
return float(vals.mean()) / 100.0
|
||||
if {"close", "prev_close"} <= set(df.columns):
|
||||
sub = df.select(["close", "prev_close"]).drop_nulls()
|
||||
if sub.height > 0:
|
||||
@@ -169,6 +186,15 @@ def build_overview(
|
||||
hist_rows: dict[str, dict[str, Any]] = hist["rows"]
|
||||
|
||||
bench_rt = _bench_rt_pct(quote_service) if quote_service is not None else 0.0
|
||||
# 实时叠加按板块基准: 科创板减科创50、创业板减创业板综指, 不再全市场混均值
|
||||
bench_by_key: dict[str, float] = {}
|
||||
if quote_service is not None:
|
||||
try:
|
||||
index_quotes = quote_service.get_index_quotes()
|
||||
except Exception:
|
||||
index_quotes = None
|
||||
for k in BENCH_KEYS:
|
||||
bench_by_key[k] = bench_rt_pct_for(index_quotes, k)
|
||||
# enriched 已含今日收盘 (盘后已同步) 时, 今日涨跌已计入历史偏离, 不再叠加
|
||||
includes_today = cache_date is not None and cache_date >= date.today().isoformat()
|
||||
|
||||
@@ -176,7 +202,9 @@ def build_overview(
|
||||
for symbol, base in hist_rows.items():
|
||||
rule = rule_for(symbol, base.get("name"))
|
||||
rt_pct = base.get("rt_pct")
|
||||
rt_delta = 0.0 if includes_today else ((rt_pct or 0.0) - bench_rt)
|
||||
rt_delta = 0.0 if includes_today else (
|
||||
(rt_pct or 0.0) - bench_by_key.get(_bench_key_of(symbol), 0.0)
|
||||
)
|
||||
|
||||
windows: dict[str, dict[str, Any]] = {}
|
||||
max_closeness = 0.0
|
||||
|
||||
@@ -72,8 +72,8 @@ _ANSI_RE = re.compile(r"\x1b\[[0-9;?]*[ -/]*[@-~]")
|
||||
|
||||
|
||||
# ----------------------------------------------------------------
|
||||
# 用户 focus 输入净化 — 防止通过"特别关注"绕过红线诱导 AI 给出买卖建议
|
||||
# 命中任一敏感词时,整个 focus 被丢弃(返回空串),由各 analyzer 据此跳过注入。
|
||||
# 用户 focus 输入规范化。交易建议类表达不会被静默丢弃,而是由统一提示词
|
||||
# 转换成客观价位、风险和情景分析,避免历史报告显示了 focus、模型却没有收到。
|
||||
# ----------------------------------------------------------------
|
||||
_FOCUS_BLOCKLIST = re.compile(
|
||||
r"买入|卖出|加仓|减仓|轻仓|重仓|半仓|全仓|仓位|止损|止盈|"
|
||||
@@ -87,19 +87,36 @@ _FOCUS_BLOCKLIST = re.compile(
|
||||
|
||||
|
||||
def sanitize_focus(focus: str) -> str:
|
||||
"""净化用户输入的 focus 文本。
|
||||
|
||||
命中交易指令/投资建议类敏感词时返回空串,阻止其注入 AI 提示词。
|
||||
这是对系统提示词红线的兜底:即便用户试图通过 focus 绕过,也不会生效。
|
||||
"""
|
||||
"""规范化 focus 中的首尾空白与连续换行。"""
|
||||
if not focus:
|
||||
return ""
|
||||
text = focus.strip()
|
||||
text = re.sub(r"\s+", " ", focus).strip()
|
||||
return text
|
||||
|
||||
|
||||
def build_focus_instruction(focus: str, *, report_name: str = "分析报告") -> str:
|
||||
"""构建所有报告共用的关注重点指令。
|
||||
|
||||
有关注点时要求模型在固定报告结构之前先直接回应。若原问题涉及交易
|
||||
建议,保留问题语义但要求转换成中立的数据分析,不再无提示地整段丢弃。
|
||||
"""
|
||||
text = sanitize_focus(focus)
|
||||
if not text:
|
||||
return ""
|
||||
|
||||
lines = [
|
||||
"## 用户关注重点(必须优先回应)",
|
||||
f"用户关注: {text}",
|
||||
f"请在完整{report_name}最前面先输出 `### 0. 🔎 关注重点回应`,"
|
||||
"用 2-4 条带具体数据的结论直接回应;随后继续完成既定报告结构,"
|
||||
"并在相关章节加深分析。不要只复述问题。",
|
||||
]
|
||||
if _FOCUS_BLOCKLIST.search(text):
|
||||
return ""
|
||||
return text
|
||||
lines.append(
|
||||
"该关注点含有买卖、仓位、目标价或预测类表达。不得给出相应操作结论;"
|
||||
"请将其转换为客观的技术/财务状态、关键价位、风险因素和条件情景后回应。"
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def current_ai_provider() -> str:
|
||||
@@ -309,11 +326,13 @@ async def stream_ai_text(
|
||||
temperature: float | None = 0.5,
|
||||
max_tokens: int | None = 4000,
|
||||
timeout: float = 180.0,
|
||||
prefer_final_answer: bool = False,
|
||||
) -> AsyncIterator[str]:
|
||||
"""Yield text deltas from the configured provider.
|
||||
|
||||
Codex CLI only exposes the final assistant message for this use case, so it
|
||||
yields one complete chunk after the command exits.
|
||||
yields one complete chunk after the command exits. ``prefer_final_answer``
|
||||
lets compatible providers prioritize visible content over hidden reasoning.
|
||||
|
||||
max_tokens=None 表示不限制输出(同 generate_ai_text 的说明)。
|
||||
"""
|
||||
@@ -328,6 +347,7 @@ async def stream_ai_text(
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
timeout=timeout,
|
||||
prefer_final_answer=prefer_final_answer,
|
||||
):
|
||||
yield chunk
|
||||
|
||||
@@ -374,6 +394,7 @@ async def _stream_openai(
|
||||
temperature: float | None,
|
||||
max_tokens: int | None,
|
||||
timeout: float,
|
||||
prefer_final_answer: bool,
|
||||
) -> AsyncIterator[str]:
|
||||
ai_key = secrets_store.get_ai_key()
|
||||
if not ai_key:
|
||||
@@ -381,15 +402,16 @@ async def _stream_openai(
|
||||
|
||||
client = _openai_client(ai_key, timeout)
|
||||
model = current_ai_model()
|
||||
base_url = secrets_store.get_ai_config("ai_base_url", settings.ai_base_url)
|
||||
req_messages = list(messages)
|
||||
|
||||
async def _iter(stream):
|
||||
async for chunk in stream:
|
||||
delta = chunk.choices[0].delta if chunk.choices else None
|
||||
if delta and delta.content:
|
||||
yield delta.content
|
||||
|
||||
kwargs = _openai_kwargs(temperature=temperature, max_tokens=max_tokens)
|
||||
kwargs = _openai_kwargs(
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
model=model,
|
||||
base_url=base_url,
|
||||
prefer_final_answer=prefer_final_answer,
|
||||
)
|
||||
while True:
|
||||
try:
|
||||
stream = await client.chat.completions.create(
|
||||
@@ -410,7 +432,7 @@ async def _stream_openai(
|
||||
raise
|
||||
|
||||
try:
|
||||
async for piece in _iter(stream):
|
||||
async for piece in _iter_openai_text(stream):
|
||||
yield piece
|
||||
except Exception as exc:
|
||||
if _is_openai_transport_error(exc):
|
||||
@@ -418,6 +440,53 @@ async def _stream_openai(
|
||||
raise
|
||||
|
||||
|
||||
_LENGTH_FINISH_REASONS = {"length", "max_tokens", "max_output_tokens"}
|
||||
|
||||
|
||||
async def _iter_openai_text(stream) -> AsyncIterator[str]:
|
||||
"""Normalize an OpenAI-compatible stream into complete text deltas.
|
||||
|
||||
Reasoning models may spend the entire completion budget on
|
||||
``reasoning_content`` and finish with HTTP 200 but no user-visible text.
|
||||
Treat that response, and any length-truncated partial response, as a
|
||||
terminal generation error instead of silently reporting success.
|
||||
"""
|
||||
content_seen = False
|
||||
reasoning_seen = False
|
||||
finish_reason = ""
|
||||
|
||||
async for chunk in stream:
|
||||
choices = getattr(chunk, "choices", None) or []
|
||||
if not choices:
|
||||
continue
|
||||
choice = choices[0]
|
||||
reason = getattr(choice, "finish_reason", None)
|
||||
if reason:
|
||||
finish_reason = str(reason)
|
||||
|
||||
delta = getattr(choice, "delta", None)
|
||||
if delta is None:
|
||||
continue
|
||||
if getattr(delta, "reasoning_content", None):
|
||||
reasoning_seen = True
|
||||
content = getattr(delta, "content", None)
|
||||
if content:
|
||||
content_seen = True
|
||||
yield content
|
||||
|
||||
if finish_reason in _LENGTH_FINISH_REASONS:
|
||||
if reasoning_seen and not content_seen:
|
||||
raise RuntimeError(
|
||||
"AI 推理达到输出长度上限, 未生成正文; 请提高输出 Token 上限或改用非推理模型"
|
||||
)
|
||||
raise RuntimeError("AI 输出达到长度上限, 内容不完整; 请提高输出 Token 上限后重试")
|
||||
|
||||
if not content_seen:
|
||||
if reasoning_seen:
|
||||
raise RuntimeError("AI 仅返回推理内容, 未生成正文; 请检查模型配置或改用非推理模型")
|
||||
raise RuntimeError("AI 服务未返回正文内容; 请检查模型配置或稍后重试")
|
||||
|
||||
|
||||
def _openai_client(api_key: str, timeout: float):
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
@@ -435,6 +504,7 @@ def _openai_client(api_key: str, timeout: float):
|
||||
# 只在 400 明确指出对应参数时移除该参数并重试; 每个参数最多移除一次。
|
||||
_TEMP_REJECT_HINTS = ("temperature", "only 1 is allowed")
|
||||
_REASONING_EFFORT_REJECT_HINTS = ("reasoning_effort", "reasoning effort")
|
||||
_THINKING_BODY_REJECT_HINTS = ("thinking",)
|
||||
|
||||
|
||||
def _is_temperature_rejected(exc: Exception) -> bool:
|
||||
@@ -457,6 +527,16 @@ def _is_reasoning_effort_rejected(exc: Exception) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _is_thinking_body_rejected(exc: Exception) -> bool:
|
||||
"""True if the upstream 400 specifically rejects the thinking extra_body."""
|
||||
if getattr(exc, "status_code", None) != 400:
|
||||
return False
|
||||
text = _openai_error_detail(exc) or str(exc)
|
||||
return _openai_error_param(exc) == "thinking" or any(
|
||||
h in text.lower() for h in _THINKING_BODY_REJECT_HINTS
|
||||
)
|
||||
|
||||
|
||||
def _openai_error_param(exc: Exception) -> str:
|
||||
body = getattr(exc, "body", None)
|
||||
if not isinstance(body, dict):
|
||||
@@ -476,11 +556,26 @@ def _openai_retry_kwargs(exc: Exception, kwargs: dict) -> dict | None:
|
||||
if "reasoning_effort" in retry_kwargs and _is_reasoning_effort_rejected(exc):
|
||||
retry_kwargs.pop("reasoning_effort")
|
||||
return retry_kwargs
|
||||
if "extra_body" in retry_kwargs and _is_thinking_body_rejected(exc):
|
||||
# DeepSeek thinking 禁用参数被拒 (模型/API 版本差异): 回退默认思考模式
|
||||
# 重试; 报告若因此被推理挤占正文, 由 _iter_openai_text 显式报错。
|
||||
retry_kwargs.pop("extra_body")
|
||||
return retry_kwargs
|
||||
return None
|
||||
|
||||
|
||||
def _openai_kwargs(*, temperature: float | None, max_tokens: int | None) -> dict:
|
||||
"""Build OpenAI create() kwargs; optional parameters are omitted when empty.
|
||||
_DEEPSEEK_V4_MODELS = {"deepseek-v4-flash", "deepseek-v4-pro"}
|
||||
|
||||
|
||||
def _openai_kwargs(
|
||||
*,
|
||||
temperature: float | None,
|
||||
max_tokens: int | None,
|
||||
model: str = "",
|
||||
base_url: str = "",
|
||||
prefer_final_answer: bool = False,
|
||||
) -> dict:
|
||||
"""Build OpenAI create() kwargs and map supported provider capabilities.
|
||||
|
||||
max_tokens=None 时不传 — 由服务端默认上限管理(推理模型的思考 token 也
|
||||
计入该参数预算, 限制会挤占正文, 见 stream_ai_text 文档)。
|
||||
@@ -494,6 +589,15 @@ def _openai_kwargs(*, temperature: float | None, max_tokens: int | None) -> dict
|
||||
reasoning_effort = current_openai_reasoning_effort()
|
||||
if reasoning_effort:
|
||||
kwargs["reasoning_effort"] = reasoning_effort
|
||||
if (
|
||||
prefer_final_answer
|
||||
and model.strip().lower() in _DEEPSEEK_V4_MODELS
|
||||
and urlsplit(base_url.strip()).hostname == "api.deepseek.com"
|
||||
):
|
||||
# DeepSeek V4 defaults to thinking mode. For report-style tasks the
|
||||
# hidden reasoning shares max_tokens with the final answer and can
|
||||
# exhaust the budget before any visible content is emitted.
|
||||
kwargs["extra_body"] = {"thinking": {"type": "disabled"}}
|
||||
return kwargs
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
"""自动挖掘 L1 编排: 全量因子统计筛选 → 达标池。
|
||||
|
||||
流程定位 (对应方案「四层漏斗」):
|
||||
- L1 本模块: 注册表全量因子批量检验, 按置信档门槛筛出达标因子 (近期窗口,
|
||||
仅作"有信号"的先验过滤; 最终达标由挖掘引擎的逐折训练选择与嵌套样本外
|
||||
验证裁定)。
|
||||
- L2/L3/L4 由现有挖掘引擎完成: 相关性剪枝 (prune_correlated_factors)、
|
||||
束搜索组合 (beam_search_factor_combinations)、嵌套样本外验证与达标
|
||||
门槛 (evaluate_candidate_gate), 本模块不重复实现。
|
||||
|
||||
达标判据与检验页服务端判读同源 (|t_NW| / BH q / |IC| / |IR|), 按档放宽或收紧;
|
||||
q 值缺失时按"通过"处理 (探索档小样本下 BH 校正保守)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, timedelta
|
||||
from typing import Any, Literal
|
||||
|
||||
from app.backtest.factor import FactorBacktestService, FactorBatchConfig
|
||||
from app.factors.registry import factor_columns_view
|
||||
|
||||
Profile = Literal["exploratory", "balanced", "strict"]
|
||||
|
||||
# 挖掘请求的因子池上限 (与 MiningStartRequest.factor_names max_length 对齐)
|
||||
MAX_AUTO_POOL = 48
|
||||
|
||||
# L1 筛选窗口: 近一年 (与挖掘窗口解耦, 只筛"近期有信号", 长窗口验证交给引擎)
|
||||
SCREEN_WINDOW_DAYS = 365
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ScreenGate:
|
||||
min_abs_ic: float
|
||||
min_abs_ir: float
|
||||
min_abs_t: float
|
||||
max_q: float
|
||||
|
||||
def to_dict(self) -> dict[str, float]:
|
||||
return {
|
||||
"min_abs_ic": self.min_abs_ic,
|
||||
"min_abs_ir": self.min_abs_ir,
|
||||
"min_abs_t": self.min_abs_t,
|
||||
"max_q": self.max_q,
|
||||
}
|
||||
|
||||
|
||||
SCREEN_GATES: dict[str, ScreenGate] = {
|
||||
"exploratory": ScreenGate(min_abs_ic=0.02, min_abs_ir=0.15, min_abs_t=1.5, max_q=0.20),
|
||||
"balanced": ScreenGate(min_abs_ic=0.02, min_abs_ir=0.30, min_abs_t=2.0, max_q=0.10),
|
||||
"strict": ScreenGate(min_abs_ic=0.03, min_abs_ir=0.50, min_abs_t=2.5, max_q=0.05),
|
||||
}
|
||||
|
||||
|
||||
def classify_factor(item: dict[str, Any], gate: ScreenGate) -> str | None:
|
||||
"""返回 None 表示达标; 否则返回首个未过的门槛, 格式统一为「类别 (细节)」。"""
|
||||
if item.get("error"):
|
||||
return f"计算失败 ({str(item['error'])[:40]})"
|
||||
ic = item.get("ic_mean")
|
||||
ir = item.get("ir")
|
||||
t = item.get("t_newey_west")
|
||||
q = item.get("q_value")
|
||||
if ic is None or ir is None:
|
||||
return "样本不足 (无有效 IC/IR)"
|
||||
if abs(ic) < gate.min_abs_ic:
|
||||
return f"预测力弱 (|IC|<{gate.min_abs_ic:.2f})"
|
||||
if abs(ir) < gate.min_abs_ir:
|
||||
return f"稳定度低 (|IR|<{gate.min_abs_ir:.2f})"
|
||||
if t is None:
|
||||
return "样本不足 (无 NW t 值)"
|
||||
if abs(t) < gate.min_abs_t:
|
||||
return f"不显著 (|t|<{gate.min_abs_t:.1f})"
|
||||
if q is not None and q > gate.max_q:
|
||||
return f"多重检验未过 (q>{gate.max_q:.2f})"
|
||||
return None
|
||||
|
||||
|
||||
def _short_reason(reason: str) -> str:
|
||||
"""失败原因归并到短类别 (「类别 (细节)」的前半段), 供原因分布统计。"""
|
||||
return reason.split(" (", 1)[0].strip()
|
||||
|
||||
|
||||
def _finite_or_none(value: Any) -> float | None:
|
||||
"""NaN/Inf 一律归 None, 避免写入任务存储时产生非法 JSON。"""
|
||||
if isinstance(value, (int, float)) and math.isfinite(value):
|
||||
return float(value)
|
||||
return None
|
||||
|
||||
|
||||
def _metric_row(item: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"factor_name": item.get("factor_name"),
|
||||
"label": item.get("label") or item.get("factor_name"),
|
||||
"group": item.get("group") or "",
|
||||
"ic": _finite_or_none(item.get("ic_mean")),
|
||||
"ir": _finite_or_none(item.get("ir")),
|
||||
"t": _finite_or_none(item.get("t_newey_west")),
|
||||
"q": _finite_or_none(item.get("q_value")),
|
||||
"direction": 1 if (item.get("ic_mean") or 0) >= 0 else -1,
|
||||
}
|
||||
|
||||
|
||||
def screen_all_factors(
|
||||
engine: Any,
|
||||
*,
|
||||
asset_type: str,
|
||||
start: date | None,
|
||||
end: date,
|
||||
profile: str,
|
||||
max_factors: int = MAX_AUTO_POOL,
|
||||
) -> dict[str, Any]:
|
||||
"""L1 全量筛选: 注册表全部适用因子批量检验 → 达标池 + 失败原因分布。
|
||||
|
||||
start=None 时取近 SCREEN_WINDOW_DAYS 天; 显式 start 只会收紧 (不放宽) 筛选窗口。
|
||||
"""
|
||||
gate = SCREEN_GATES.get(profile)
|
||||
if gate is None:
|
||||
raise ValueError(f"unknown mining profile: {profile}")
|
||||
|
||||
candidates = [
|
||||
str(item["id"])
|
||||
for item in factor_columns_view()
|
||||
if asset_type in item.get("asset_types", ["stock"])
|
||||
]
|
||||
screen_start = max(start or date.min, end - timedelta(days=SCREEN_WINDOW_DAYS))
|
||||
began = time.perf_counter()
|
||||
service = FactorBacktestService(engine)
|
||||
batch = service.run_batch(FactorBatchConfig(
|
||||
factor_names=candidates,
|
||||
symbols=None,
|
||||
start=screen_start,
|
||||
end=end,
|
||||
rebalance="daily",
|
||||
asset_type=asset_type,
|
||||
))
|
||||
elapsed_ms = round((time.perf_counter() - began) * 1000, 1)
|
||||
|
||||
qualified: list[dict[str, Any]] = []
|
||||
failed: list[dict[str, Any]] = []
|
||||
by_name = {str(getattr(item, "factor_name", None)): item for item in batch.results}
|
||||
for name in candidates:
|
||||
item = by_name.get(name)
|
||||
if item is None:
|
||||
failed.append({"factor_name": name, "label": name, "group": "",
|
||||
"ic": None, "ir": None, "t": None, "q": None,
|
||||
"reason": "未返回结果"})
|
||||
continue
|
||||
# 非有限值先清洗 (NaN 与任何比较均为 False, 会绕过门槛误判达标)
|
||||
reason = classify_factor({
|
||||
"error": getattr(item, "error", None),
|
||||
"ic_mean": _finite_or_none(getattr(item, "ic_mean", None)),
|
||||
"ir": _finite_or_none(getattr(item, "ir", None)),
|
||||
"t_newey_west": _finite_or_none(getattr(item, "t_newey_west", None)),
|
||||
"q_value": _finite_or_none(getattr(item, "q_value", None)),
|
||||
}, gate)
|
||||
row = _metric_row({
|
||||
"factor_name": getattr(item, "factor_name", None),
|
||||
"label": getattr(item, "label", None),
|
||||
"group": getattr(item, "group", None),
|
||||
"ic_mean": getattr(item, "ic_mean", None),
|
||||
"ir": getattr(item, "ir", None),
|
||||
"t_newey_west": getattr(item, "t_newey_west", None),
|
||||
"q_value": getattr(item, "q_value", None),
|
||||
})
|
||||
if reason is None:
|
||||
qualified.append(row)
|
||||
else:
|
||||
failed.append({**row, "reason": reason})
|
||||
|
||||
# 池按 |IC|*|IR| 降序 (截面信噪比口径), 截断到挖掘上限
|
||||
qualified.sort(key=lambda row: abs(row["ic"] or 0.0) * abs(row["ir"] or 0.0), reverse=True)
|
||||
pool = [row["factor_name"] for row in qualified[:max_factors]]
|
||||
|
||||
reason_counts: dict[str, int] = {}
|
||||
for row in failed:
|
||||
category = _short_reason(row["reason"])
|
||||
reason_counts[category] = reason_counts.get(category, 0) + 1
|
||||
|
||||
return {
|
||||
"profile": profile,
|
||||
"gate": gate.to_dict(),
|
||||
"screen_window": {"start": screen_start.isoformat(), "end": end.isoformat()},
|
||||
"n_total": len(candidates),
|
||||
"n_qualified": len(qualified),
|
||||
"pool": pool,
|
||||
"pool_truncated": len(qualified) > len(pool),
|
||||
"qualified": qualified,
|
||||
"failed": failed,
|
||||
"reason_counts": dict(sorted(reason_counts.items(), key=lambda kv: -kv[1])),
|
||||
"elapsed_ms": elapsed_ms,
|
||||
}
|
||||
@@ -5,9 +5,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import date
|
||||
from datetime import date, timedelta
|
||||
from typing import Literal
|
||||
|
||||
import numpy as np
|
||||
@@ -20,6 +21,10 @@ from app.tickflow.repository import KlineRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 旧信号回测的指标 warmup 日历窗口 (#201): 与 backtest.factor.FACTOR_WARMUP_DAYS
|
||||
# 同源 (120 交易日 → 保守取日历日), 覆盖 MA60/MACD/BOLL 等最长回看
|
||||
_WARMUP_CALENDAR_DAYS = 120 * 1.6
|
||||
|
||||
# vectorbt 是 optional extras(见 pyproject.toml).未装时只有 backtest 不可用,其他功能正常.
|
||||
_vbt = None
|
||||
_vbt_unavailable_reason: str | None = None
|
||||
@@ -118,6 +123,31 @@ _SIGNAL_COLS: dict[SignalKind, str] = {
|
||||
}
|
||||
|
||||
|
||||
def _build_max_hold_exits(entries: pd.DataFrame, max_hold_days: int) -> pd.DataFrame:
|
||||
"""为每个入场信号在 max_hold_days 个交易日后生成一个强制退出信号。
|
||||
|
||||
返回与 entries 同形状的布尔矩阵, 仅在「入场位之后第 max_hold_days 个交易日」置
|
||||
True(不含入场位本身), 供调用方与用户 exits 做 OR。
|
||||
|
||||
两处易错点(见 #198):
|
||||
- 必须从全 False 起步。若用 `entries.copy()` 起步会把入场位当成退出位,
|
||||
导致入场当日即被强制平仓。
|
||||
- 用单步定位写入 `iloc[row, col_loc]`。链式 `iloc[row][col] = True` 写入的是
|
||||
临时行副本, 在 pandas Copy-on-Write 语义下不会落到原矩阵(pandas 3.x 直接报错),
|
||||
强制退出信号会静默丢失。
|
||||
"""
|
||||
out = pd.DataFrame(False, index=entries.index, columns=entries.columns)
|
||||
n = len(entries)
|
||||
for col in entries.columns:
|
||||
col_loc = out.columns.get_loc(col)
|
||||
entry_rows = np.where(entries[col].to_numpy())[0]
|
||||
for i in entry_rows:
|
||||
end_i = min(int(i) + max_hold_days, n - 1)
|
||||
if end_i > i:
|
||||
out.iloc[end_i, col_loc] = True
|
||||
return out
|
||||
|
||||
|
||||
class BacktestService:
|
||||
def __init__(self, repo: KlineRepository) -> None:
|
||||
self.repo = repo
|
||||
@@ -137,11 +167,17 @@ class BacktestService:
|
||||
try:
|
||||
from app.tickflow.repository import enriched_dirname
|
||||
enriched_glob = str(self.repo.store.data_dir / enriched_dirname(asset_type) / "**" / "*.parquet")
|
||||
# 指标 warmup (#201): MA/MACD/RSI/BOLL 需要区间前的历史窗口,
|
||||
# 直接按 [start,end] 过滤后 compute_all 会让区间头部的指标失真。
|
||||
# 与挖掘侧同款公式 (mining_runtime: warmup = max(120, bars*1.6)),
|
||||
# 此处指标最长回看约 120 交易日, 取保守日历日窗口; 数据不足时
|
||||
# 自然退化 (有多少算多少)。计算完成后裁回 [start,end]。
|
||||
warmup_start = start - timedelta(days=_WARMUP_CALENDAR_DAYS)
|
||||
df = (
|
||||
scan_enriched_parquet(enriched_glob)
|
||||
.filter(
|
||||
(pl.col("symbol").is_in(symbols))
|
||||
& (pl.col("date") >= start)
|
||||
& (pl.col("date") >= warmup_start)
|
||||
& (pl.col("date") <= end)
|
||||
)
|
||||
.sort(["date", "symbol"])
|
||||
@@ -157,6 +193,7 @@ class BacktestService:
|
||||
# 即时计算指标 + 信号
|
||||
from app.indicators.pipeline import compute_all
|
||||
df = compute_all(df)
|
||||
df = df.filter(pl.col("date") >= start)
|
||||
|
||||
# 选择需要的列
|
||||
needed_cols = [
|
||||
@@ -205,6 +242,12 @@ class BacktestService:
|
||||
return result if result is not None else pd.DataFrame()
|
||||
|
||||
def run(self, config: BacktestConfig) -> BacktestResult:
|
||||
from app.services.heavy_job_limiter import shared_heavy_job_limiter
|
||||
|
||||
with shared_heavy_job_limiter.slot("exclusive"):
|
||||
return self._run(config)
|
||||
|
||||
def _run(self, config: BacktestConfig) -> BacktestResult:
|
||||
vbt = _get_vbt()
|
||||
run_id = uuid.uuid4().hex[:10]
|
||||
|
||||
@@ -270,16 +313,10 @@ class BacktestService:
|
||||
if config.stop_loss_pct is not None:
|
||||
pf_kwargs["sl_stop"] = abs(config.stop_loss_pct)
|
||||
if config.max_hold_days is not None:
|
||||
# vectorbt 没有内置 max-hold;用时间退出近似:
|
||||
# 在 max_hold_days 后强制 exit
|
||||
exits_idx = entries.copy()
|
||||
for col in entries.columns:
|
||||
entry_rows = np.where(entries[col].values)[0]
|
||||
for i in entry_rows:
|
||||
end_i = min(i + config.max_hold_days, len(entries) - 1)
|
||||
if end_i > i:
|
||||
exits_idx.iloc[end_i][col] = True
|
||||
pf_kwargs["exits"] = (exits | exits_idx).astype(bool)
|
||||
# vectorbt 没有内置 max-hold;用时间退出近似:入场后第 max_hold_days
|
||||
# 个交易日强制 exit, 与用户 exits 做 OR(保留原有信号退出)。
|
||||
forced_exits = _build_max_hold_exits(entries, config.max_hold_days)
|
||||
pf_kwargs["exits"] = (exits | forced_exits).astype(bool)
|
||||
|
||||
pf = vbt.Portfolio.from_signals(**pf_kwargs)
|
||||
except Exception as e: # noqa: BLE001
|
||||
@@ -388,10 +425,16 @@ def _config_to_dict(c: BacktestConfig) -> dict:
|
||||
|
||||
|
||||
def _json_safe(v):
|
||||
# 非有限浮点 (inf / NaN) 必须先于原生标量分支拦下: Starlette 的 JSONResponse 用
|
||||
# json.dumps(allow_nan=False) 渲染, 漏一个就是整个响应 500。pf.stats() 经
|
||||
# pandas Series.to_dict() 出来时 numpy 标量已被装箱成原生 float (全胜时
|
||||
# Profit Factor = inf, 零波动时 Sharpe = NaN), 两条分支都要覆盖。
|
||||
if isinstance(v, (float, np.floating)) and not math.isfinite(float(v)):
|
||||
return None
|
||||
if isinstance(v, (int, float, str, bool)) or v is None:
|
||||
return v
|
||||
if isinstance(v, (np.floating, np.integer)):
|
||||
return float(v) if not np.isnan(float(v)) else None
|
||||
return float(v)
|
||||
if hasattr(v, "isoformat"):
|
||||
return v.isoformat()
|
||||
return str(v)
|
||||
|
||||
@@ -125,13 +125,11 @@ def _compute_rotation_signals(dates: list[str], columns: dict) -> dict:
|
||||
dates_asc = list(reversed(dates))
|
||||
|
||||
# 收集每个概念在各日期的 (排名, 涨幅)。排名 = 该日在列中的索引 + 1。
|
||||
concept_data: dict[str, list[tuple[int, float]]] = {}
|
||||
concept_data: dict[str, dict[str, tuple[int, float]]] = {}
|
||||
for d in dates_asc:
|
||||
col = columns.get(d) or []
|
||||
for idx, (name, pct) in enumerate(col):
|
||||
concept_data.setdefault(name, []).append((idx + 1, pct))
|
||||
|
||||
n_dates = len(dates_asc)
|
||||
concept_data.setdefault(name, {})[d] = (idx + 1, pct)
|
||||
|
||||
def _stats(ranks_pcts: list[tuple[int, float]]) -> dict:
|
||||
ranks = [r for r, _ in ranks_pcts]
|
||||
@@ -151,10 +149,10 @@ def _compute_rotation_signals(dates: list[str], columns: dict) -> dict:
|
||||
institutional: list[dict] = []
|
||||
hot_money: list[dict] = []
|
||||
|
||||
for concept, rp in concept_data.items():
|
||||
# 缺失日补 (大排名, 0 涨幅) 保持时间轴对齐
|
||||
if len(rp) < n_dates:
|
||||
rp = rp + [(999, 0.0)] * (n_dates - len(rp))
|
||||
for concept, by_date in concept_data.items():
|
||||
# 缺失日按日期归位补 (大排名, 0 涨幅) —— 补位必须落在缺席的那一天,
|
||||
# 一律追加到末尾会把"只在最近几日上榜"的新晋概念读成退潮。
|
||||
rp = [by_date.get(d, (999, 0.0)) for d in dates_asc]
|
||||
s = _stats(rp)
|
||||
s["concept"] = concept
|
||||
|
||||
@@ -208,11 +206,24 @@ def _compute_rotation_signals(dates: list[str], columns: dict) -> dict:
|
||||
# ================================================================
|
||||
|
||||
def _fmt_pct(v) -> str:
|
||||
"""概念/行业涨幅: 小数口径 (0.0522 = +5.22%), 展示前乘 100。"""
|
||||
if v is None:
|
||||
return "—"
|
||||
return f"{v*100:+.2f}%"
|
||||
|
||||
|
||||
def _fmt_index_pct(v) -> str:
|
||||
"""指数涨跌幅: 百分数口径 (CONTRIBUTING §3.1), 直接展示, 不能再乘一次 100。
|
||||
|
||||
build_market_overview 的 indices[].change_pct 在数据边界已转成百分数
|
||||
(quote_service._build_index_quotes 与 _index_quotes 的 DB 兜底都已乘过 100),
|
||||
与 market_recap._build_indices_block 的展示口径一致。
|
||||
"""
|
||||
if v is None:
|
||||
return "—"
|
||||
return f"{v:+.2f}%"
|
||||
|
||||
|
||||
def _build_market_block(overview: dict) -> str:
|
||||
"""大盘背景精简块 (复用 market_overview 已算好的字段)。"""
|
||||
indices = overview.get("indices") or []
|
||||
@@ -224,7 +235,7 @@ def _build_market_block(overview: dict) -> str:
|
||||
for idx in indices[:4]:
|
||||
name = idx.get("name") or idx.get("symbol") or "?"
|
||||
chg = idx.get("change_pct")
|
||||
idx_lines.append(f"{name} {_fmt_pct(chg)}")
|
||||
idx_lines.append(f"{name} {_fmt_index_pct(chg)}")
|
||||
idx_str = " / ".join(idx_lines) or "指数缺失"
|
||||
|
||||
total_amount = (amt.get("total") or 0) / 1e8 # 元 → 亿
|
||||
@@ -278,10 +289,10 @@ def _build_user_prompt(signals: dict, overview: dict, days: int, dates: list[str
|
||||
_build_signal_block("🎰 游资特征 (排名波动大)", signals.get("hot_money", [])),
|
||||
]
|
||||
|
||||
from app.services.ai_provider import sanitize_focus
|
||||
safe_focus = sanitize_focus(focus)
|
||||
if safe_focus:
|
||||
parts.extend(["", f"本次分析请特别关注: {safe_focus}"])
|
||||
from app.services.ai_provider import build_focus_instruction
|
||||
focus_instruction = build_focus_instruction(focus, report_name=f"{dim}轮动分析报告")
|
||||
if focus_instruction:
|
||||
parts.extend(["", focus_instruction])
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
@@ -374,6 +385,7 @@ async def analyze_rotation_stream(
|
||||
temperature=0.5,
|
||||
# 不限制输出(推理模型思考 token 计入预算, 见 ai_provider.stream_ai_text)
|
||||
max_tokens=None,
|
||||
prefer_final_answer=True,
|
||||
):
|
||||
got_content = True
|
||||
yield json.dumps({"type": "delta", "content": delta}, ensure_ascii=False)
|
||||
|
||||
@@ -8,6 +8,7 @@
|
||||
- null → batch 拉取 / 盘后计算写入的权威历史 → 完整
|
||||
- d < 今天 且 时刻 < d 15:00 → 盘中快照 (停机前实时写的) → 坏
|
||||
- d < 今天 且 时刻 ≥ d 15:00 → 尾盘定版 (close_final) → 完整
|
||||
- batch 权威行中仅夹杂少量零成交实时行 → 停牌残留 → 忽略
|
||||
- d == 今天 → 实时更新中, 属正常, 不校验
|
||||
- 分区缺失的工作日 → 缺口 (工作日近似; 节假日误报的代价是一次空范围拉取,
|
||||
merge-upsert 空写, 无害)
|
||||
@@ -108,6 +109,62 @@ def _is_snapshot(day: date, quote_ts_ms: int | None) -> bool:
|
||||
return ts.date() == day and ts.time() < CLOSE_CUTOFF
|
||||
|
||||
|
||||
def _partition_is_snapshot(day: date, part_dir: Path, quote_ts_max_ms: int | None) -> bool:
|
||||
"""判断整个分区是否仍是盘中快照, 而非同步后遗留的停牌实时行。
|
||||
|
||||
batch 行用 null quote_ts 标识权威历史。实时轮询曾把停牌股票的 09:15、
|
||||
零成交记录写入分区; 后续 batch 会过滤停牌日, merge-upsert 因而留下这些
|
||||
孤立行。若分区已有 batch 行, 且当日收盘前的实时行全部零成交, 则它们不应
|
||||
让整个分区反复进入修复。整分区都是实时行时仍按快照处理, 包括盘前零成交。
|
||||
"""
|
||||
if not _is_snapshot(day, quote_ts_max_ms):
|
||||
return False
|
||||
|
||||
start_ms = int(datetime.combine(day, dt_time.min, tzinfo=CN_TZ).timestamp() * 1000)
|
||||
cutoff_ms = int(datetime.combine(day, CLOSE_CUTOFF, tzinfo=CN_TZ).timestamp() * 1000)
|
||||
authoritative_rows = 0
|
||||
suspicious_rows = 0
|
||||
|
||||
for path in sorted(part_dir.glob("*.parquet")):
|
||||
try:
|
||||
schema = pl.read_parquet_schema(path)
|
||||
if "quote_ts" not in schema:
|
||||
continue
|
||||
columns = [
|
||||
name for name in ("quote_ts", "volume", "amount")
|
||||
if name in schema
|
||||
]
|
||||
frame = pl.read_parquet(path, columns=columns).with_columns(
|
||||
pl.col("quote_ts").cast(pl.Int64, strict=False),
|
||||
)
|
||||
authoritative_rows += frame["quote_ts"].null_count()
|
||||
suspicious = frame.filter(
|
||||
pl.col("quote_ts").is_between(start_ms, cutoff_ms, closed="left")
|
||||
)
|
||||
if suspicious.is_empty():
|
||||
continue
|
||||
suspicious_rows += suspicious.height
|
||||
|
||||
activity_columns = [
|
||||
name for name in ("volume", "amount") if name in suspicious.columns
|
||||
]
|
||||
if not activity_columns:
|
||||
return True
|
||||
has_activity = suspicious.select(
|
||||
pl.any_horizontal(
|
||||
pl.col(name).cast(pl.Float64, strict=False).fill_null(0) > 0
|
||||
for name in activity_columns
|
||||
).any()
|
||||
).item()
|
||||
if has_activity:
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.debug("snapshot residue scan skipped %s: %s", path, e)
|
||||
return True
|
||||
|
||||
return suspicious_rows > 0 and authoritative_rows <= suspicious_rows
|
||||
|
||||
|
||||
def _candidate_days(today: date, lookback_days: int) -> list[date]:
|
||||
"""最近 lookback_days 自然日内、严格早于今天的工作日 (节假日近似, 误报无害)。"""
|
||||
days: list[date] = []
|
||||
@@ -157,7 +214,7 @@ def scan_recent_integrity(
|
||||
continue
|
||||
part_dir = base / f"date={day.isoformat()}"
|
||||
quote_ts = _quote_ts_max_ms(part_dir)
|
||||
if _is_snapshot(day, quote_ts):
|
||||
if _partition_is_snapshot(day, part_dir, quote_ts):
|
||||
issues.append(IntegrityIssue(day=day, table=table, kind="snapshot"))
|
||||
|
||||
issues.sort(key=lambda i: (i.day, i.table))
|
||||
@@ -256,6 +313,7 @@ def launch_integrity_repair(app_state, start_date: date, reason: str) -> tuple[s
|
||||
JobCancelledError,
|
||||
job_store,
|
||||
release_run_slot,
|
||||
run_with_capacity,
|
||||
try_acquire_run_slot,
|
||||
)
|
||||
from app.services.repair_daily import run_repair_daily
|
||||
@@ -282,8 +340,7 @@ def launch_integrity_repair(app_state, start_date: date, reason: str) -> tuple[s
|
||||
if not try_acquire_run_slot(job_id):
|
||||
job_store.fail(job_id, "已有数据任务在运行(或上一次任务卡死未结束),请稍后再试")
|
||||
return
|
||||
job_store.start(job_id)
|
||||
result = _run()
|
||||
result = run_with_capacity(job_id, _run)
|
||||
if isinstance(result, dict) and "error" in result:
|
||||
job_store.fail(job_id, str(result["error"]))
|
||||
else:
|
||||
|
||||
@@ -292,9 +292,30 @@ class DepthService:
|
||||
self._persist(enriched_date)
|
||||
|
||||
def _call_depth_batch(self, symbols: list[str]) -> dict:
|
||||
"""调 tf.depth.batch, 按 capset 的 batch 切片 + 节流。返回 {symbol: MarketDepth}。"""
|
||||
from app.tickflow.client import get_client
|
||||
tf = get_client()
|
||||
"""按独立五档路由取数; 所有 provider 共用分片限速且失败不跨源回退。"""
|
||||
from app.services import preferences
|
||||
|
||||
provider_name = preferences.get_depth5_data_provider()
|
||||
if provider_name == "tickflow":
|
||||
from app.data_providers.registry import get_provider
|
||||
|
||||
provider = get_provider("tickflow")
|
||||
else:
|
||||
from app.data_providers import custom as custom_sources
|
||||
|
||||
try:
|
||||
if not custom_sources.provider_has_dataset(provider_name, "depth5"):
|
||||
logger.warning("depth provider %s 未声明 depth5, 跳过本轮", provider_name)
|
||||
return {}
|
||||
provider = custom_sources.get_provider(provider_name)
|
||||
except Exception as e:
|
||||
logger.warning("depth provider %s 解析失败, 跳过本轮: %s", provider_name, e)
|
||||
return {}
|
||||
|
||||
fetch_depth = getattr(provider, "get_depth_batch", None)
|
||||
if not callable(fetch_depth):
|
||||
logger.warning("depth provider %s 未实现 get_depth_batch, 跳过本轮", provider_name)
|
||||
return {}
|
||||
|
||||
capset = self._get_capset()
|
||||
limit = resolve_limit(capset, Cap.DEPTH5_BATCH, default_batch=100, default_rpm=30)
|
||||
@@ -304,12 +325,23 @@ class DepthService:
|
||||
for i, chunk in enumerate(chunks):
|
||||
sleep_between_batches(i, limit.rpm, default_interval=2.0)
|
||||
try:
|
||||
# SDK 的 batch 内部已按 batch_size 切, 这里再切一层防单请求过大
|
||||
data = tf.depth.batch(chunk)
|
||||
data = fetch_depth(chunk)
|
||||
if isinstance(data, dict):
|
||||
result.update(data)
|
||||
else:
|
||||
logger.warning(
|
||||
"depth provider %s 第 %d 批返回非 dict, 已跳过",
|
||||
provider_name,
|
||||
i + 1,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("depth.batch 第 %d 批失败(%d 只): %s", i + 1, len(chunk), e)
|
||||
logger.warning(
|
||||
"depth provider %s 第 %d 批失败(%d 只): %s",
|
||||
provider_name,
|
||||
i + 1,
|
||||
len(chunk),
|
||||
e,
|
||||
)
|
||||
# 单批失败不影响其他批
|
||||
return result
|
||||
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
"""SMTP email notification adapter.
|
||||
|
||||
Transport failures are isolated from alert persistence and SSE delivery. SMTP credentials
|
||||
are supplied by the caller from ``secrets_store`` and never logged here.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import smtplib
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from email.message import EmailMessage
|
||||
from email.utils import parseaddr
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SECURITY_MODES = {"ssl", "starttls", "none"}
|
||||
_MAX_ATTEMPTS = 2
|
||||
|
||||
|
||||
def is_valid_email(address: str) -> bool:
|
||||
"""Small dependency-free mailbox validation suitable for configuration checks."""
|
||||
parsed = parseaddr((address or "").strip())[1]
|
||||
if parsed != (address or "").strip() or parsed.count("@") != 1:
|
||||
return False
|
||||
local, domain = parsed.rsplit("@", 1)
|
||||
return bool(local and domain and "." in domain and " " not in parsed)
|
||||
|
||||
|
||||
def is_configured(config: dict) -> bool:
|
||||
"""Return whether the non-secret fields are sufficient to attempt delivery."""
|
||||
sender = str(config.get("from_address") or config.get("username") or "").strip()
|
||||
recipients = config.get("to_addresses") or []
|
||||
return bool(config.get("host") and sender and recipients)
|
||||
|
||||
|
||||
def send_email(
|
||||
config: dict,
|
||||
password: str,
|
||||
subject: str,
|
||||
body: str,
|
||||
*,
|
||||
max_attempts: int = _MAX_ATTEMPTS,
|
||||
) -> bool:
|
||||
"""Send one UTF-8 plain-text email through SSL, STARTTLS, or plain SMTP."""
|
||||
if not is_configured(config):
|
||||
return False
|
||||
|
||||
host = str(config.get("host") or "").strip()
|
||||
try:
|
||||
port = int(config.get("port", 465))
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
security = str(config.get("security") or "ssl")
|
||||
username = str(config.get("username") or "").strip()
|
||||
sender = str(config.get("from_address") or username).strip()
|
||||
recipients = [str(item).strip() for item in config.get("to_addresses", [])]
|
||||
if (
|
||||
not 1 <= port <= 65535
|
||||
or security not in SECURITY_MODES
|
||||
or not is_valid_email(sender)
|
||||
or not recipients
|
||||
or any(not is_valid_email(item) for item in recipients)
|
||||
):
|
||||
return False
|
||||
|
||||
message = EmailMessage()
|
||||
message["Subject"] = str(subject or "TickFlow 通知")
|
||||
message["From"] = sender
|
||||
message["To"] = ", ".join(recipients)
|
||||
message.set_content(str(body or ""))
|
||||
|
||||
last_err = ""
|
||||
for attempt in range(1, max_attempts + 1):
|
||||
smtp = None
|
||||
try:
|
||||
if security == "ssl":
|
||||
smtp = smtplib.SMTP_SSL(host, port, timeout=10)
|
||||
else:
|
||||
smtp = smtplib.SMTP(host, port, timeout=10)
|
||||
if security == "starttls":
|
||||
smtp.ehlo()
|
||||
smtp.starttls()
|
||||
smtp.ehlo()
|
||||
if username:
|
||||
smtp.login(username, password)
|
||||
smtp.send_message(message)
|
||||
# Delivery already succeeded; a failed QUIT must not retry and duplicate the email.
|
||||
try:
|
||||
smtp.quit()
|
||||
except Exception:
|
||||
with suppress(Exception):
|
||||
smtp.close()
|
||||
return True
|
||||
except Exception as exc: # SMTP/network errors must not escape
|
||||
last_err = str(exc)
|
||||
if smtp is not None:
|
||||
with suppress(Exception):
|
||||
smtp.close()
|
||||
if attempt < max_attempts:
|
||||
time.sleep(1)
|
||||
|
||||
logger.warning("邮件推送最终失败(已尝试 %d 次): %s", max_attempts, last_err)
|
||||
return False
|
||||
@@ -1,6 +1,7 @@
|
||||
"""扩展数据服务 — 配置管理 + 文件解析 + Parquet 存储。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import codecs
|
||||
import copy
|
||||
import json
|
||||
import logging
|
||||
@@ -40,7 +41,8 @@ class PullConfig:
|
||||
"url", "method", "headers", "body", "response_path",
|
||||
"field_map", "schedule_minutes", "enabled",
|
||||
"last_run", "last_status", "last_message", "last_rows",
|
||||
"next_run", "time_window_start", "time_window_end",
|
||||
"next_run", "time_window_start", "time_window_end", "date_param",
|
||||
"auth",
|
||||
)
|
||||
|
||||
def __init__(
|
||||
@@ -60,6 +62,8 @@ class PullConfig:
|
||||
next_run: str | None = None,
|
||||
time_window_start: str | None = None,
|
||||
time_window_end: str | None = None,
|
||||
date_param: str | None = None,
|
||||
auth: dict | None = None,
|
||||
) -> None:
|
||||
self.url = url
|
||||
self.method = method # GET | POST
|
||||
@@ -76,6 +80,12 @@ class PullConfig:
|
||||
self.next_run = next_run # 下次预计运行 (ISO, 调度器写入)
|
||||
self.time_window_start = time_window_start # 每日拉取窗口起始 "HH:MM", None=不限
|
||||
self.time_window_end = time_window_end # 每日拉取窗口结束 "HH:MM", None=不限
|
||||
# 接口按日期查询的参数名 (如 "date"): 非 None 时请求
|
||||
# 带 ?{date_param}=YYYY-MM-DD, 支持历史回补; None = 接口只有当日快照
|
||||
self.date_param = date_param
|
||||
# 拉取接口鉴权方式 {"type": "none|bearer|header|query", "header": ..., "param": ...},
|
||||
# 与自定义行情源 AuthConfig 同口径; Key 本体存 secrets_store, 不落 config.json
|
||||
self.auth = auth
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
@@ -94,6 +104,8 @@ class PullConfig:
|
||||
"next_run": self.next_run,
|
||||
"time_window_start": self.time_window_start,
|
||||
"time_window_end": self.time_window_end,
|
||||
"date_param": self.date_param,
|
||||
"auth": self.auth,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
@@ -116,9 +128,25 @@ class PullConfig:
|
||||
next_run=d.get("next_run"),
|
||||
time_window_start=d.get("time_window_start"),
|
||||
time_window_end=d.get("time_window_end"),
|
||||
date_param=d.get("date_param"),
|
||||
auth=d.get("auth"),
|
||||
)
|
||||
|
||||
|
||||
def ext_api_key_field(config_id: str) -> str:
|
||||
"""扩展数据拉取 API Key 在 secrets.json 中的字段名。"""
|
||||
return f"ext_{config_id}_api_key"
|
||||
|
||||
|
||||
def get_ext_api_key(config_id: str) -> str:
|
||||
"""取扩展数据拉取接口的 API Key: secrets.json 优先, 环境变量 EXT_{ID}_API_KEY 兜底。"""
|
||||
from app import secrets_store
|
||||
|
||||
return secrets_store.get_env_backed_secret(
|
||||
ext_api_key_field(config_id), f"EXT_{config_id.upper()}_API_KEY"
|
||||
)
|
||||
|
||||
|
||||
class ExtConfig:
|
||||
"""一个扩展数据源的完整配置。"""
|
||||
__slots__ = (
|
||||
@@ -265,7 +293,7 @@ class ExtConfigStore:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def upsert(self, config: ExtConfig) -> None:
|
||||
def upsert(self, config: ExtConfig, *, keep_strategy_cache: bool = False) -> None:
|
||||
config.updated_at = datetime.now().isoformat()
|
||||
cp = self._config_path(config.id)
|
||||
cp.parent.mkdir(parents=True, exist_ok=True)
|
||||
@@ -273,6 +301,10 @@ class ExtConfigStore:
|
||||
json.dumps(config.to_dict(), ensure_ascii=False, indent=2),
|
||||
encoding="utf-8",
|
||||
)
|
||||
# 字段集/模式变化会改变扩展列集合: 失效扩展帧缓存与策略结果缓存。
|
||||
# 定时拉取循环的 last_run/next_run 例行回写传 keep_strategy_cache=True,
|
||||
# 否则每轮拉取后策略页缓存被状态回写清空 (数据写入链路已另行放行)。
|
||||
_invalidate_ext_derived(self._base.parent, keep_strategy_cache=keep_strategy_cache)
|
||||
|
||||
def delete(self, config_id: str) -> bool:
|
||||
import shutil
|
||||
@@ -283,6 +315,7 @@ class ExtConfigStore:
|
||||
if not cp.exists():
|
||||
return False
|
||||
shutil.rmtree(cp.parent, ignore_errors=True)
|
||||
_invalidate_ext_derived(self._base.parent)
|
||||
return True
|
||||
|
||||
def _migrate_legacy(self, old_path: Path) -> None:
|
||||
@@ -444,6 +477,44 @@ def apply_config_mapping(df: pl.DataFrame, config: ExtConfig, data_dir: Path) ->
|
||||
return df
|
||||
|
||||
|
||||
# 编码识别与转换的分块大小,与 ext_data 上传写入用的块大小一致。
|
||||
_TRANSCODE_CHUNK_BYTES = 1024 * 1024
|
||||
|
||||
|
||||
def _decodes_as(file_path: Path, encoding: str) -> bool:
|
||||
"""整个文件能否按 encoding 完整解码,逐块判断,不把文件读进内存。"""
|
||||
decoder = codecs.getincrementaldecoder(encoding)()
|
||||
try:
|
||||
with file_path.open("rb") as src:
|
||||
while chunk := src.read(_TRANSCODE_CHUNK_BYTES):
|
||||
decoder.decode(chunk)
|
||||
decoder.decode(b"", True) # 结尾处的半个字符也算解码失败
|
||||
except UnicodeDecodeError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _transcode_to_utf8(file_path: Path, out_path: Path, encoding: str) -> bool:
|
||||
"""按 encoding 逐块转成 UTF-8 写入 out_path;解码失败则删除半成品返回 False。
|
||||
|
||||
增量解码器负责跨块边界的多字节字符:GBK 一个汉字两字节,正好落在块边界
|
||||
上时前半截会被留到下一块,不会被误判成解码失败。
|
||||
"""
|
||||
decoder = codecs.getincrementaldecoder(encoding)()
|
||||
try:
|
||||
with (
|
||||
file_path.open("rb") as src,
|
||||
out_path.open("w", encoding="utf-8", newline="") as dst,
|
||||
):
|
||||
while chunk := src.read(_TRANSCODE_CHUNK_BYTES):
|
||||
dst.write(decoder.decode(chunk))
|
||||
dst.write(decoder.decode(b"", True))
|
||||
except UnicodeDecodeError:
|
||||
out_path.unlink(missing_ok=True)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def ensure_utf8_csv(file_path: Path) -> Path:
|
||||
"""确保 CSV 文件以 UTF-8 编码可读,非 UTF-8(如 GBK/GB18030)则转换。
|
||||
|
||||
@@ -454,21 +525,14 @@ def ensure_utf8_csv(file_path: Path) -> Path:
|
||||
返回值:若已是 UTF-8 则返回原路径;否则在同目录写一个 *.utf8 文件并返回它
|
||||
(调用方用临时目录,随目录一起清理)。
|
||||
"""
|
||||
raw = file_path.read_bytes()
|
||||
# BOM 处理:UTF-8-SIG 等带 BOM 文件直接交给 Polars(它认识 BOM)
|
||||
try:
|
||||
raw.decode("utf-8")
|
||||
if _decodes_as(file_path, "utf-8"):
|
||||
return file_path # 已是合法 UTF-8
|
||||
except UnicodeDecodeError:
|
||||
pass
|
||||
# 依次尝试常见中文编码,第一个能完整解码的即为命中
|
||||
for enc in ("gb18030", "gbk", "gb2312", "big5"):
|
||||
try:
|
||||
text = raw.decode(enc)
|
||||
except UnicodeDecodeError:
|
||||
continue
|
||||
out_path = file_path.with_suffix(file_path.suffix + ".utf8")
|
||||
out_path.write_text(text, encoding="utf-8")
|
||||
if not _transcode_to_utf8(file_path, out_path, enc):
|
||||
continue
|
||||
logger.info("CSV 编码转换 %s → %s (%s)", file_path.name, out_path.name, enc)
|
||||
return out_path
|
||||
# 都无法解码:返回原路径,让 Polars 抛出更精确的原始错误
|
||||
@@ -531,6 +595,8 @@ def write_ext_parquet(
|
||||
config: ExtConfig,
|
||||
data_dir: Path,
|
||||
snapshot_date: date | None = None,
|
||||
*,
|
||||
keep_strategy_cache: bool = False,
|
||||
) -> int:
|
||||
"""将 DataFrame 写入扩展数据 Parquet。
|
||||
|
||||
@@ -582,9 +648,26 @@ def write_ext_parquet(
|
||||
df = cast_df_to_schema(df, config.fields)
|
||||
df.write_parquet(out_path)
|
||||
logger.info("扩展表写入: %s → %s (%d 行)", config.id, out_path, len(df))
|
||||
# 扩展列已接入 enriched 帧/因子注册表: 写入后必须失效相关缓存
|
||||
_invalidate_ext_derived(data_dir, keep_strategy_cache=keep_strategy_cache)
|
||||
return len(df)
|
||||
|
||||
|
||||
def _invalidate_ext_derived(data_dir: Path, *, keep_strategy_cache: bool = False) -> None:
|
||||
"""扩展数据/配置变更 → 扩展帧缓存 + 因子同步状态 + 策略结果缓存。
|
||||
|
||||
惰性导入避免与 ext_factors (反向惰性引用本模块) 构成模块级环。
|
||||
repo 内存 enriched 缓存由 API 层 repo.clear_cache() 补充清理。
|
||||
keep_strategy_cache 语义见 ext_factors.invalidate_ext_caches。
|
||||
"""
|
||||
try:
|
||||
from app.factors.ext_factors import invalidate_ext_caches
|
||||
|
||||
invalidate_ext_caches(data_dir, keep_strategy_cache=keep_strategy_cache)
|
||||
except Exception as e:
|
||||
logger.warning("扩展数据缓存失效失败: %s", e)
|
||||
|
||||
|
||||
def delete_ext_parquet(config_id: str, data_dir: Path) -> None:
|
||||
"""删除扩展数据源关联的所有 Parquet 数据(保留 config.json)。
|
||||
|
||||
@@ -601,6 +684,7 @@ def delete_ext_parquet(config_id: str, data_dir: Path) -> None:
|
||||
if ts_dir.exists():
|
||||
import shutil
|
||||
shutil.rmtree(ts_dir, ignore_errors=True)
|
||||
_invalidate_ext_derived(data_dir)
|
||||
|
||||
|
||||
def fix_symbol_format(config: ExtConfig, data_dir: Path) -> int:
|
||||
@@ -657,6 +741,8 @@ def rows_to_parquet(
|
||||
config: ExtConfig,
|
||||
data_dir: Path,
|
||||
snapshot_date: date | None = None,
|
||||
*,
|
||||
keep_strategy_cache: bool = False,
|
||||
) -> int:
|
||||
"""将 JSON 行列表转为 DataFrame 写入 Parquet,复用 write_ext_parquet 的存储逻辑。
|
||||
|
||||
@@ -667,4 +753,7 @@ def rows_to_parquet(
|
||||
df = apply_config_mapping(df, config, data_dir)
|
||||
if "symbol" in df.columns:
|
||||
df = df.with_columns(pl.col("symbol").cast(pl.Utf8))
|
||||
return write_ext_parquet(df, config, data_dir, snapshot_date=snapshot_date)
|
||||
return write_ext_parquet(
|
||||
df, config, data_dir, snapshot_date=snapshot_date,
|
||||
keep_strategy_cache=keep_strategy_cache,
|
||||
)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""内置扩展数据预设 — 概念/行业首次启动自动拉取。
|
||||
"""内置扩展数据预设 — 概念/行业启动时只创建配置, 等待用户手动获取 (#199)。
|
||||
|
||||
设计原则:
|
||||
- 扩展数据通用逻辑零改动 (ExtConfig / fetch_and_ingest / API / 前端均不动)
|
||||
@@ -55,14 +55,16 @@ def _concept_preset() -> ExtConfig:
|
||||
ExtField("股票简称", "string", "股票简称"),
|
||||
ExtField("所属概念", "string", "所属概念"),
|
||||
],
|
||||
description="同花顺概念分类 (首次启动自动拉取, 可在扩展数据页手动更新)",
|
||||
description="同花顺概念分类 (启动仅创建配置, 在概念/行业页手动获取)",
|
||||
symbol_map={"type": "mapped", "col": "股票代码"},
|
||||
code_map={"type": "computed", "from": "symbol", "method": "strip_exchange"},
|
||||
pull=PullConfig(
|
||||
url=_CONCEPT_DATA_URL,
|
||||
method="GET",
|
||||
schedule_minutes=1440,
|
||||
enabled=True,
|
||||
# enabled=False: ensure_builtin_presets 承诺启动不拉取, PullScheduler
|
||||
# 只调度 enabled 配置; 手动获取走 fetch_preset 独立路径不受影响 (#199)
|
||||
enabled=False,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -84,14 +86,15 @@ def _industry_preset() -> ExtConfig:
|
||||
ExtField("股票简称", "string", "股票简称"),
|
||||
ExtField("所属同花顺行业", "string", "所属同花顺行业"),
|
||||
],
|
||||
description="同花顺行业分类 (首次启动自动拉取, 可在扩展数据页手动更新)",
|
||||
description="同花顺行业分类 (启动仅创建配置, 在概念/行业页手动获取)",
|
||||
symbol_map={"type": "mapped", "col": "股票代码"},
|
||||
code_map={"type": "computed", "from": "symbol", "method": "strip_exchange"},
|
||||
pull=PullConfig(
|
||||
url=_INDUSTRY_DATA_URL,
|
||||
method="GET",
|
||||
schedule_minutes=1440,
|
||||
enabled=True,
|
||||
# 同概念 preset: 出厂禁用, 避免启动即网络拉取 (#199)
|
||||
enabled=False,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -179,8 +182,11 @@ async def _fetch_json(url: str) -> list[dict]:
|
||||
"""
|
||||
import httpx
|
||||
|
||||
# 延迟导入避免与 ext_pull 循环依赖; 出站请求带 tsp 标识头
|
||||
from app.services.ext_pull import outbound_headers
|
||||
|
||||
async with httpx.AsyncClient(timeout=30) as client:
|
||||
resp = await client.get(url)
|
||||
resp = await client.get(url, headers=outbound_headers())
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
|
||||
@@ -5,30 +5,58 @@ import asyncio
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
from datetime import date, datetime, timezone
|
||||
from datetime import UTC, date, datetime, timezone
|
||||
from functools import reduce
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from app.market_time import cn_now, cn_today
|
||||
from app.services.ext_data import (
|
||||
ExtConfig,
|
||||
ExtConfigStore,
|
||||
PullConfig,
|
||||
rows_to_parquet,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def outbound_headers(user_headers: dict[str, str] | None = None) -> dict[str, str]:
|
||||
"""扩展数据出站请求的默认标识头。
|
||||
|
||||
默认携带 User-Agent: tsp/<版本> 与 X-TSP-Client: tick-stock-panel,
|
||||
供服务端 (如 tickflow-hub) 识别本项目的请求。用户在拉取配置里显式
|
||||
设置的同名头优先 (大小写不敏感), 不被标识头覆盖。
|
||||
"""
|
||||
from app import __version__
|
||||
|
||||
defaults = {
|
||||
"User-Agent": f"tsp/{__version__}",
|
||||
"X-TSP-Client": "tick-stock-panel",
|
||||
}
|
||||
override = {k.lower() for k in (user_headers or {})}
|
||||
return {
|
||||
**{k: v for k, v in defaults.items() if k.lower() not in override},
|
||||
**(user_headers or {}),
|
||||
}
|
||||
|
||||
|
||||
def _in_time_window(start: str | None, end: str | None) -> bool:
|
||||
"""检查当前本地时间是否在每日时间窗口内。
|
||||
"""检查当前北京时间是否在每日时间窗口内。
|
||||
|
||||
start/end 为 "HH:MM" 格式。两者都为 None 时不限制(返回 True)。
|
||||
支持跨午夜窗口(如 22:00-02:00)。
|
||||
|
||||
用北京时间而不是本地时间: 这个窗口是照着 A 股交易时段设的, 而
|
||||
market_time 模块开篇就写明「服务器/容器本地时区不可靠 (python:slim
|
||||
镜像默认 UTC)」。UTC 容器里 9:30-15:00 的窗口实际落在北京 17:30-23:00,
|
||||
每天都在收盘之后。
|
||||
"""
|
||||
if not start or not end:
|
||||
return True
|
||||
now = datetime.now().strftime("%H:%M")
|
||||
now = cn_now().strftime("%H:%M")
|
||||
if start <= end:
|
||||
return start <= now < end
|
||||
# 跨午夜: 如 22:00-02:00
|
||||
@@ -102,21 +130,76 @@ def _apply_preset_flatten(config_id: str, rows: list[dict]) -> list[dict]:
|
||||
return flatten(rows)
|
||||
|
||||
|
||||
async def fetch_and_ingest(
|
||||
config: ExtConfig,
|
||||
data_dir,
|
||||
) -> tuple[int, str]:
|
||||
"""执行一次拉取: 请求外部 API → 解析响应 → 写入 Parquet。
|
||||
def _with_date_param(url: str, date_param: str | None, day: date) -> str:
|
||||
"""接口按日查询参数: ?{date_param}=YYYY-MM-DD (已有 query 用 &)。"""
|
||||
if not date_param:
|
||||
return url
|
||||
sep = "&" if "?" in url else "?"
|
||||
return f"{url}{sep}{date_param}={day.isoformat()}"
|
||||
|
||||
Returns:
|
||||
(rows_written, date_str)
|
||||
|
||||
def _apply_auth(config_id: str, auth: dict | None, url: str, headers: dict[str, str]) -> str:
|
||||
"""把 secrets_store 里的 API Key 注入出站请求。
|
||||
|
||||
鉴权三型与自定义行情源 AuthConfig 同口径: bearer → {header: "Bearer <key>"},
|
||||
header → {header: <key>}, query → ?{param}=<key>。Key 只存 secrets.json,
|
||||
不落 config.json; 配置了鉴权但未设置 Key 时 fail-closed 直接报错,
|
||||
避免不带凭据请求被服务端记成无效调用。返回 (可能追加了参数的) url。
|
||||
"""
|
||||
pull = config.pull
|
||||
if not pull or not pull.url:
|
||||
raise ValueError("拉取未配置或 URL 为空")
|
||||
from urllib.parse import quote
|
||||
|
||||
from app.services.ext_data import get_ext_api_key
|
||||
|
||||
auth_type = str((auth or {}).get("type") or "none").lower()
|
||||
if auth_type == "none":
|
||||
return url
|
||||
key = get_ext_api_key(config_id)
|
||||
if not key:
|
||||
raise ValueError(f"已配置 {auth_type} 鉴权但未设置 API Key, 请在拉取设置中填写")
|
||||
if auth_type == "bearer":
|
||||
headers[str(auth.get("header") or "Authorization")] = f"Bearer {key}"
|
||||
elif auth_type == "header":
|
||||
headers[str(auth.get("header") or "Authorization")] = key
|
||||
elif auth_type == "query":
|
||||
name = str(auth.get("param") or "token")
|
||||
sep = "&" if "?" in url else "?"
|
||||
url = f"{url}{sep}{name}={quote(key, safe='')}"
|
||||
else:
|
||||
raise ValueError(f"未知鉴权类型: {auth_type!r} (可选 none/bearer/header/query)")
|
||||
return url
|
||||
|
||||
|
||||
def _assert_rows_date(rows: list[dict], day: date) -> None:
|
||||
"""金融契约: 响应行的 date 字段 (若提供) 必须与请求日期一致。
|
||||
|
||||
服务端忽略日期参数返回当日数据时会静默把当日值写进历史分区,
|
||||
造成整个时序口径错乱 —— 此处 fail-closed 拒绝 (实测确实有忽略
|
||||
?date= 的接口)。date 字段缺省的接口不做校验。
|
||||
"""
|
||||
want = day.isoformat()
|
||||
for r in rows[:20]:
|
||||
if not isinstance(r, dict):
|
||||
continue
|
||||
raw = r.get("date")
|
||||
if raw is None:
|
||||
continue
|
||||
if str(raw)[:10] != want:
|
||||
raise ValueError(
|
||||
f"接口返回的日期 {str(raw)[:10]!r} 与请求日期 {want} 不一致 "
|
||||
"(接口可能不支持日期参数), 已拒绝写入该分区"
|
||||
)
|
||||
|
||||
|
||||
async def _request_json(pull: PullConfig, config_id: str, day: date | None = None) -> Any:
|
||||
"""发起一次拉取请求并返回解析后的 JSON。
|
||||
|
||||
正式拉取 (带日期参数) 与设置页"测试" (不带) 共用同一实现,
|
||||
保证 UA 标识头与 API Key 鉴权注入只有一套口径。
|
||||
"""
|
||||
url = _with_date_param(pull.url, pull.date_param, day) if day else pull.url
|
||||
async with httpx.AsyncClient(timeout=30) as client:
|
||||
headers = pull.headers or {}
|
||||
headers = outbound_headers(pull.headers)
|
||||
url = _apply_auth(config_id, pull.auth, url, headers)
|
||||
kwargs: dict[str, Any] = {"headers": headers}
|
||||
|
||||
if pull.method.upper() == "POST" and pull.body:
|
||||
@@ -124,19 +207,28 @@ async def fetch_and_ingest(
|
||||
if "content-type" not in {k.lower() for k in headers}:
|
||||
kwargs["headers"]["Content-Type"] = "application/json"
|
||||
|
||||
resp = await client.request(pull.method.upper(), pull.url, **kwargs)
|
||||
resp = await client.request(pull.method.upper(), url, **kwargs)
|
||||
resp.raise_for_status()
|
||||
try:
|
||||
return resp.json()
|
||||
except Exception as e:
|
||||
raise ValueError(f"响应不是有效 JSON: {e}") from e
|
||||
|
||||
# 解析 JSON
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception as e:
|
||||
raise ValueError(f"响应不是有效 JSON: {e}") from e
|
||||
|
||||
async def fetch_rows_for_date(config: ExtConfig, target_date: date) -> list[dict]:
|
||||
"""按日期请求外部 API 并解析为行 (不写盘)。空数据返回 []。
|
||||
|
||||
与 fetch_and_ingest 共用同一解析链 (response_path/预设转换/字段映射/
|
||||
关联字段校验), 历史回补与当日拉取不产生第二套口径。
|
||||
"""
|
||||
pull = config.pull
|
||||
if not pull or not pull.url:
|
||||
raise ValueError("拉取未配置或 URL 为空")
|
||||
|
||||
data = await _request_json(pull, config.id, day=target_date)
|
||||
|
||||
# 提取行
|
||||
rows = _extract_rows(data, pull.response_path)
|
||||
if not rows:
|
||||
raise ValueError("提取到的行数为 0")
|
||||
|
||||
# 内置预设 (概念/行业): 应用结构转换, 让产出 schema 与分析页一致。
|
||||
# 否则 raw 接口列 (concepts/industries 数组、name) 会直接覆盖正确的 part.parquet,
|
||||
@@ -157,10 +249,143 @@ async def fetch_and_ingest(
|
||||
if rows and not ({"symbol", "code"} & row_keys or mapped_cols & row_keys):
|
||||
raise ValueError("数据行中缺少 symbol/code 字段,请配置字段映射或标的映射")
|
||||
|
||||
# 写入
|
||||
snap = date.today()
|
||||
n = rows_to_parquet(rows, config, data_dir, snapshot_date=snap)
|
||||
return n, snap.isoformat()
|
||||
_assert_rows_date(rows, target_date)
|
||||
return rows
|
||||
|
||||
|
||||
async def fetch_and_ingest(
|
||||
config: ExtConfig,
|
||||
data_dir,
|
||||
target_date: date | None = None,
|
||||
*,
|
||||
keep_strategy_cache: bool = False,
|
||||
) -> tuple[int, str]:
|
||||
"""执行一次拉取: 请求外部 API → 解析响应 → 写入 Parquet。
|
||||
|
||||
target_date 默认当日; 历史回补传入目标日期 (写入对应分区)。
|
||||
keep_strategy_cache=True 由定时拉取循环传入: 例行刷新不清策略结果缓存。
|
||||
Returns:
|
||||
(rows_written, date_str)
|
||||
"""
|
||||
# 同上: 落盘分区按北京日期, 否则 UTC 容器在北京时间 08:00 之前写的是前一天。
|
||||
day = target_date or cn_today()
|
||||
rows = await fetch_rows_for_date(config, day)
|
||||
if not rows:
|
||||
raise ValueError("提取到的行数为 0")
|
||||
n = rows_to_parquet(
|
||||
rows, config, data_dir, snapshot_date=day,
|
||||
keep_strategy_cache=keep_strategy_cache,
|
||||
)
|
||||
return n, day.isoformat()
|
||||
|
||||
|
||||
MAX_BACKFILL_DAYS = 120 # 单次回补上限: 同步端点, 控制请求时长
|
||||
_BACKFILL_DAY_INTERVAL_S = 0.3 # 相邻请求间隔 (对数据源限速)
|
||||
_BACKFILL_429_WAIT_S = 30.0 # 429 限流退避时长 (服务端按分钟配额)
|
||||
_BACKFILL_MAX_CONSECUTIVE_429 = 3 # 连续 429 天数达到阈值 → 中止本次回补
|
||||
_RATE_LIMIT_ABORT_REASON = "限流中止 (429), 稍后重跑回补可自动续补剩余日期"
|
||||
|
||||
|
||||
def _status_code(e: BaseException) -> int | None:
|
||||
"""从 httpx.HTTPStatusError 提取状态码; 非该类异常返回 None。"""
|
||||
resp = getattr(e, "response", None)
|
||||
return getattr(resp, "status_code", None)
|
||||
|
||||
|
||||
def _day_partition(data_dir, config_id: str, day: date) -> Path:
|
||||
return Path(data_dir) / "ext_data" / config_id / "timeseries" / f"date={day.isoformat()}" / "part.parquet"
|
||||
|
||||
|
||||
async def backfill_history(
|
||||
config: ExtConfig,
|
||||
data_dir,
|
||||
start: date,
|
||||
end: date,
|
||||
) -> dict:
|
||||
"""按本地交易日逐日回补 timeseries 历史分区 (幂等, 已有分区跳过)。
|
||||
|
||||
前提: 接口支持按日期查询 (pull.date_param 已配置)。交易日取本地日K
|
||||
分区日期 —— 非交易日无人气数据, 也避免无谓请求。单日失败不中断,
|
||||
汇总进 failed 清单返回; 该日无数据 (空响应或 404) 计入 empty 跳过。
|
||||
"""
|
||||
if config.mode != "timeseries":
|
||||
raise ValueError("仅 timeseries 模式支持历史回补 (snapshot 无历史概念)")
|
||||
pull = config.pull
|
||||
if not pull or not pull.url:
|
||||
raise ValueError("拉取未配置或 URL 为空")
|
||||
if not pull.date_param:
|
||||
raise ValueError("接口未配置日期参数 (date_param) —— 需接口支持 ?日期参数= 历史查询")
|
||||
if start > end:
|
||||
raise ValueError("开始日期不能晚于结束日期")
|
||||
if (end - start).days + 1 > MAX_BACKFILL_DAYS:
|
||||
raise ValueError(f"单次回补上限 {MAX_BACKFILL_DAYS} 天, 请分段执行")
|
||||
|
||||
from app.services.dragon_tiger import _local_trading_days
|
||||
|
||||
days = [d for d in _local_trading_days(data_dir) if start <= d <= end]
|
||||
if not days:
|
||||
raise ValueError("范围内无本地交易日 (需先同步日K以确定交易日历)")
|
||||
|
||||
fetched = skipped = empty = 0
|
||||
rows_written = 0
|
||||
failed: list[dict] = []
|
||||
consecutive_429 = 0 # 连续限流天数 (重试成功即清零); 达到阈值中止本次回补
|
||||
for i, d in enumerate(days):
|
||||
part = _day_partition(data_dir, config.id, d)
|
||||
if part.exists():
|
||||
skipped += 1
|
||||
continue
|
||||
try:
|
||||
rows = await fetch_rows_for_date(config, d)
|
||||
consecutive_429 = 0
|
||||
if not rows:
|
||||
empty += 1 # 该日无数据 (服务端未归档), 不是错误
|
||||
else:
|
||||
rows_written += rows_to_parquet(rows, config, data_dir, snapshot_date=d)
|
||||
fetched += 1
|
||||
except httpx.HTTPStatusError as e:
|
||||
if _status_code(e) == 404:
|
||||
# 接口契约 (tickflow-hub /exports、/fuyao-rank): 该日无快照
|
||||
# 返回 404 —— 视为该日无数据跳过, 不计入失败
|
||||
empty += 1
|
||||
elif _status_code(e) != 429:
|
||||
failed.append({"date": d.isoformat(), "reason": str(e)[:200]})
|
||||
else:
|
||||
# 服务端按分钟配额限流: 退避后原地重试一次; 连续多日 429
|
||||
# 说明配额窗口已耗尽, 中止剩余天数 (幂等, 重跑即可续补)。
|
||||
consecutive_429 += 1
|
||||
if consecutive_429 >= _BACKFILL_MAX_CONSECUTIVE_429:
|
||||
remaining = [dd for dd in days[i:] if not _day_partition(data_dir, config.id, dd).exists()]
|
||||
failed.extend({"date": dd.isoformat(), "reason": _RATE_LIMIT_ABORT_REASON}
|
||||
for dd in remaining)
|
||||
break
|
||||
await asyncio.sleep(_BACKFILL_429_WAIT_S)
|
||||
try:
|
||||
rows = await fetch_rows_for_date(config, d)
|
||||
consecutive_429 = 0
|
||||
if not rows:
|
||||
empty += 1
|
||||
else:
|
||||
rows_written += rows_to_parquet(rows, config, data_dir, snapshot_date=d)
|
||||
fetched += 1
|
||||
except Exception as e2:
|
||||
if _status_code(e2) == 404: # 退避重试后无该日快照 → 同样视为无数据
|
||||
empty += 1
|
||||
else:
|
||||
failed.append({"date": d.isoformat(), "reason": str(e2)[:200]})
|
||||
except Exception as e:
|
||||
failed.append({"date": d.isoformat(), "reason": str(e)[:200]})
|
||||
if i + 1 < len(days):
|
||||
await asyncio.sleep(_BACKFILL_DAY_INTERVAL_S) # 限速, 对数据源礼貌
|
||||
return {
|
||||
"total_days": len(days),
|
||||
"fetched": fetched,
|
||||
"skipped_existing": skipped,
|
||||
"empty": empty,
|
||||
"failed": failed,
|
||||
"rows_written": rows_written,
|
||||
}
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -221,22 +446,21 @@ class PullScheduler:
|
||||
configs = store.load_all()
|
||||
|
||||
active_ids: set[str] = set()
|
||||
new_configs: list[ExtConfig] = []
|
||||
enabled_configs: list[ExtConfig] = []
|
||||
|
||||
for config in configs:
|
||||
if not config.pull or not config.pull.enabled or not config.pull.url:
|
||||
continue
|
||||
active_ids.add(config.id)
|
||||
if config.id not in self._tasks:
|
||||
new_configs.append(config)
|
||||
enabled_configs.append(config)
|
||||
|
||||
# 需要移除的 id (快照当前 task 字典的键, 避免遍历时改字典)
|
||||
remove_ids = [cid for cid in list(self._tasks) if cid not in active_ids]
|
||||
|
||||
# 所有对 _tasks 的修改都提交到主循环里执行, 保证线程安全
|
||||
# 对 _tasks 的一切读判断 (含增删 diff) 都放进主循环闭包里执行:
|
||||
# refresh 可能从工作线程调用, 若在调用方线程读 _tasks 再把决策
|
||||
# 提交回主循环, 两步之间主循环可能已改动字典 (TOCTOU, #203)。
|
||||
# 此处只携带与 _tasks 无关的 config 数据跨线程。
|
||||
def _apply() -> None:
|
||||
for config in new_configs:
|
||||
if config.id not in self._tasks: # 二次校验, 防重复
|
||||
for config in enabled_configs:
|
||||
if config.id not in self._tasks:
|
||||
self._tasks[config.id] = self._loop.create_task(
|
||||
self._run_loop(config)
|
||||
)
|
||||
@@ -244,11 +468,10 @@ class PullScheduler:
|
||||
"PullScheduler: scheduled %s (every %d min)",
|
||||
config.id, config.pull.schedule_minutes,
|
||||
)
|
||||
for cid in remove_ids:
|
||||
task = self._tasks.pop(cid, None)
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
logger.info("PullScheduler: removed %s", cid)
|
||||
for cid in [c for c in self._tasks if c not in active_ids]:
|
||||
task = self._tasks.pop(cid)
|
||||
task.cancel()
|
||||
logger.info("PullScheduler: removed %s", cid)
|
||||
|
||||
self._submit(_apply)
|
||||
|
||||
@@ -273,7 +496,7 @@ class PullScheduler:
|
||||
fresh.pull.last_run = datetime.now(timezone.utc).isoformat()
|
||||
fresh.pull.last_status = "skipped"
|
||||
fresh.pull.last_message = "不在拉取时间窗口内"
|
||||
store.upsert(fresh)
|
||||
store.upsert(fresh, keep_strategy_cache=True)
|
||||
logger.info("PullScheduler: %s skipped (outside time window)", config.id)
|
||||
interval = max(pull.schedule_minutes * 60, 60)
|
||||
await asyncio.sleep(interval)
|
||||
@@ -281,12 +504,16 @@ class PullScheduler:
|
||||
|
||||
# 先执行一次 (启用即拉取, 让用户立刻看到生效)
|
||||
try:
|
||||
n, d = await fetch_and_ingest(fresh, self._data_dir)
|
||||
# 例行定时刷新: 不清策略结果缓存 (见 invalidate_ext_caches),
|
||||
# 否则策略页每轮拉取后整页空白, 直到下次全量重算完成。
|
||||
n, d = await fetch_and_ingest(
|
||||
fresh, self._data_dir, keep_strategy_cache=True
|
||||
)
|
||||
fresh.pull.last_run = datetime.now(timezone.utc).isoformat()
|
||||
fresh.pull.last_status = "success"
|
||||
fresh.pull.last_message = f"{n} rows @ {d}"
|
||||
fresh.pull.last_rows = n
|
||||
store.upsert(fresh)
|
||||
store.upsert(fresh, keep_strategy_cache=True)
|
||||
logger.info("PullScheduler: %s success, %d rows", config.id, n)
|
||||
except Exception as e:
|
||||
fresh2 = store.get(config.id)
|
||||
@@ -294,19 +521,19 @@ class PullScheduler:
|
||||
fresh2.pull.last_run = datetime.now(timezone.utc).isoformat()
|
||||
fresh2.pull.last_status = "error"
|
||||
fresh2.pull.last_message = str(e)[:200]
|
||||
store.upsert(fresh2)
|
||||
store.upsert(fresh2, keep_strategy_cache=True)
|
||||
logger.warning("PullScheduler: %s error: %s", config.id, e)
|
||||
|
||||
# 间隔取自最新配置 (每次重新读取, 修复改间隔不生效)
|
||||
interval = max(pull.schedule_minutes * 60, 60) # 至少 60s
|
||||
# 预告下次运行时间, 供前端展示
|
||||
next_dt = datetime.now(timezone.utc).timestamp() + interval
|
||||
next_dt = datetime.now(UTC).timestamp() + interval
|
||||
latest = store.get(config.id)
|
||||
if latest and latest.pull:
|
||||
latest.pull.next_run = datetime.fromtimestamp(
|
||||
next_dt, tz=timezone.utc
|
||||
next_dt, tz=UTC
|
||||
).isoformat()
|
||||
store.upsert(latest)
|
||||
store.upsert(latest, keep_strategy_cache=True)
|
||||
|
||||
await asyncio.sleep(interval)
|
||||
if not self._running:
|
||||
|
||||
@@ -136,13 +136,10 @@ def _build_user_prompt(fins: dict[str, list[dict]], symbol: str, focus: str) ->
|
||||
data_json,
|
||||
"```",
|
||||
]
|
||||
from app.services.ai_provider import sanitize_focus
|
||||
safe_focus = sanitize_focus(focus)
|
||||
if safe_focus:
|
||||
lines.extend([
|
||||
"",
|
||||
f"本次分析请特别关注: {safe_focus}",
|
||||
])
|
||||
from app.services.ai_provider import build_focus_instruction
|
||||
focus_instruction = build_focus_instruction(focus, report_name="财务分析报告")
|
||||
if focus_instruction:
|
||||
lines.extend(["", focus_instruction])
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
@@ -187,6 +184,7 @@ async def analyze_financials_stream(
|
||||
temperature=0.4,
|
||||
# 不限制输出(推理模型思考 token 计入预算, 见 ai_provider.stream_ai_text)
|
||||
max_tokens=None,
|
||||
prefer_final_answer=True,
|
||||
):
|
||||
got_content = True
|
||||
yield json.dumps({"type": "delta", "content": delta}, ensure_ascii=False)
|
||||
|
||||
@@ -160,7 +160,8 @@ def _merge_report_history(*frames: pl.DataFrame) -> pl.DataFrame:
|
||||
语义(区分"覆盖"与"填空"): 每列独立取 announce_date 最新的非空值 —
|
||||
新同步行有值则覆盖旧值, 新行缺的列(如 fuyao 不提供的字段)由旧行补齐,
|
||||
实现多数据源并集共存。历史报告期不可变, 合并不会引入过期数据。
|
||||
无 announce_date 的帧按输入顺序, 后写优先(与旧行为 keep="last" 一致)。
|
||||
无 announce_date 的帧按输入顺序, 后写优先(与旧行为 keep="last" 一致);
|
||||
公告日为空视为最旧, 不得压过带公告日的行。
|
||||
"""
|
||||
valid = [
|
||||
frame
|
||||
@@ -176,7 +177,10 @@ def _merge_report_history(*frames: pl.DataFrame) -> pl.DataFrame:
|
||||
sort_keys = ["symbol", "period_end"] + (
|
||||
["announce_date"] if "announce_date" in merged.columns else []
|
||||
)
|
||||
merged = merged.sort(sort_keys, nulls_last=True)
|
||||
# 公告日为空排在最前: 排到最后会让"公告日未知"的旧行在逐列 last() 时胜出,
|
||||
# 产出 announce_date 是新公告、数值却是旧值的自相矛盾行。symbol/period_end
|
||||
# 已在上面过滤掉空值, 不受该参数影响。
|
||||
merged = merged.sort(sort_keys, nulls_last=False)
|
||||
value_cols = [c for c in merged.columns if c not in ("symbol", "period_end")]
|
||||
return (
|
||||
merged.group_by("symbol", "period_end")
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
"""文件系统小工具 — 原子写等。
|
||||
|
||||
历史遗留: json_report_store / strategy_cache / kline_sync 等模块里各有一份内联的
|
||||
同款原子写。新代码统一用本模块的 atomic_write_text, 一处实现一处维护。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def atomic_write_text(path: Path, text: str) -> None:
|
||||
"""临时文件 + os.replace 原子替换, 避免读侧读到半截 JSON。"""
|
||||
tmp = path.with_suffix(path.suffix + ".tmp")
|
||||
tmp.write_text(text, encoding="utf-8")
|
||||
os.replace(tmp, path)
|
||||
@@ -4,11 +4,12 @@ from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from collections import deque
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from typing import ClassVar, Literal
|
||||
|
||||
HeavyJobKind = Literal["normal", "mining"]
|
||||
HeavyJobKind = Literal["normal", "mining", "exclusive"]
|
||||
|
||||
|
||||
class HeavyJobLimitTimeoutError(TimeoutError):
|
||||
@@ -20,7 +21,7 @@ class HeavyJobCancelledError(RuntimeError):
|
||||
|
||||
|
||||
class HeavyJobLimiter:
|
||||
"""A weighted limiter where normal jobs cost one slot and mining costs two."""
|
||||
"""FIFO weighted capacity; exclusive jobs reserve the entire process budget."""
|
||||
|
||||
_WEIGHTS: ClassVar[dict[HeavyJobKind, int]] = {"normal": 1, "mining": 2}
|
||||
|
||||
@@ -32,8 +33,10 @@ class HeavyJobLimiter:
|
||||
self.capacity = capacity
|
||||
self._cancel_poll_interval = cancel_poll_interval
|
||||
self._used = 0
|
||||
self._acquired = {"normal": 0, "mining": 0}
|
||||
self._acquired = {"normal": 0, "mining": 0, "exclusive": 0}
|
||||
self._condition = threading.Condition()
|
||||
self._waiters: deque[object] = deque()
|
||||
self._local = threading.local()
|
||||
|
||||
@property
|
||||
def in_use(self) -> int:
|
||||
@@ -61,23 +64,29 @@ class HeavyJobLimiter:
|
||||
|
||||
deadline = None if timeout is None else time.monotonic() + timeout
|
||||
with self._condition:
|
||||
while True:
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
return False
|
||||
if self._used + weight <= self.capacity:
|
||||
self._used += weight
|
||||
self._acquired[kind] += 1
|
||||
return True
|
||||
ticket = object()
|
||||
self._waiters.append(ticket)
|
||||
try:
|
||||
while True:
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
return False
|
||||
if self._waiters[0] is ticket and self._used + weight <= self.capacity:
|
||||
self._used += weight
|
||||
self._acquired[kind] += 1
|
||||
return True
|
||||
|
||||
remaining = None if deadline is None else deadline - time.monotonic()
|
||||
if remaining is not None and remaining <= 0:
|
||||
return False
|
||||
wait_for = remaining
|
||||
if cancel_event is not None:
|
||||
wait_for = self._cancel_poll_interval
|
||||
if remaining is not None:
|
||||
wait_for = min(wait_for, remaining)
|
||||
self._condition.wait(wait_for)
|
||||
remaining = None if deadline is None else deadline - time.monotonic()
|
||||
if remaining is not None and remaining <= 0:
|
||||
return False
|
||||
wait_for = remaining
|
||||
if cancel_event is not None:
|
||||
wait_for = self._cancel_poll_interval
|
||||
if remaining is not None:
|
||||
wait_for = min(wait_for, remaining)
|
||||
self._condition.wait(wait_for)
|
||||
finally:
|
||||
self._waiters.remove(ticket)
|
||||
self._condition.notify_all()
|
||||
|
||||
def release(self, kind: HeavyJobKind = "normal") -> None:
|
||||
"""Return capacity previously acquired for ``kind``."""
|
||||
@@ -97,21 +106,33 @@ class HeavyJobLimiter:
|
||||
timeout: float | None = None,
|
||||
cancel_event: threading.Event | None = None,
|
||||
) -> Iterator[HeavyJobLimiter]:
|
||||
"""Acquire weighted capacity for the duration of a ``with`` block."""
|
||||
"""Reserve capacity in the executing thread, reusing an outer reservation."""
|
||||
weight = self._weight(kind)
|
||||
held = getattr(self._local, "weight", 0)
|
||||
if held:
|
||||
if weight > held:
|
||||
raise RuntimeError("cannot upgrade a held heavy-job reservation")
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
raise HeavyJobCancelledError(f"{kind} job was cancelled")
|
||||
yield self
|
||||
return
|
||||
acquired = self.acquire(kind, timeout=timeout, cancel_event=cancel_event)
|
||||
if not acquired:
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
raise HeavyJobCancelledError(f"{kind} job was cancelled while waiting")
|
||||
raise HeavyJobLimitTimeoutError(f"timed out waiting for {kind} job capacity")
|
||||
try:
|
||||
self._local.weight = weight
|
||||
yield self
|
||||
finally:
|
||||
self._local.weight = 0
|
||||
self.release(kind)
|
||||
|
||||
@classmethod
|
||||
def _weight(cls, kind: HeavyJobKind) -> int:
|
||||
def _weight(self, kind: HeavyJobKind) -> int:
|
||||
if kind == "exclusive":
|
||||
return self.capacity
|
||||
try:
|
||||
return cls._WEIGHTS[kind]
|
||||
return self._WEIGHTS[kind]
|
||||
except KeyError as exc:
|
||||
raise ValueError(f"unsupported heavy job kind: {kind!r}") from exc
|
||||
|
||||
|
||||
@@ -21,6 +21,8 @@ import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
from app.market_time import cn_now
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -121,6 +123,9 @@ class JsonReportStore:
|
||||
|
||||
@staticmethod
|
||||
def _now_iso() -> str:
|
||||
"""当前本地时间 ISO 字符串(带秒精度,前端 toLocaleString 友好)。"""
|
||||
from datetime import datetime
|
||||
return datetime.now().isoformat(timespec="seconds")
|
||||
"""当前北京时间 ISO 字符串(带秒精度,前端 toLocaleString 友好)。
|
||||
|
||||
用北京墙钟而非宿主机时钟: 容器默认 UTC 时, 前端把这串 naive 时间按浏览器
|
||||
本地时区解析, 刚生成的报告会显示成「8 小时前」。
|
||||
"""
|
||||
return cn_now().replace(tzinfo=None).isoformat(timespec="seconds")
|
||||
|
||||
@@ -7,7 +7,11 @@
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
import shutil
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from datetime import date, datetime, timedelta
|
||||
|
||||
@@ -177,6 +181,17 @@ def sync_and_persist_daily_batch(
|
||||
end_time = end_date or datetime.now()
|
||||
days = count or 365
|
||||
start_time = start_date or (end_time - timedelta(days=days))
|
||||
iter_daily = getattr(provider, "iter_daily", None)
|
||||
if callable(iter_daily):
|
||||
return _persist_daily_chunks(
|
||||
iter_daily(
|
||||
symbols,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
on_chunk_done=on_chunk_done,
|
||||
),
|
||||
repo,
|
||||
)
|
||||
df = provider.get_daily(
|
||||
symbols,
|
||||
start_time=start_time,
|
||||
@@ -228,6 +243,49 @@ def sync_and_persist_daily_batch(
|
||||
return df.height
|
||||
|
||||
|
||||
def _persist_daily_chunks(chunks, repo: KlineRepository) -> int:
|
||||
"""先把流式 provider 结果写入私有 staging,完整取数后再提交正式分区。"""
|
||||
staging_base = repo.store.data_dir / ".daily_sync_staging"
|
||||
_sweep_stale_daily_staging(staging_base)
|
||||
root = staging_base / uuid.uuid4().hex
|
||||
written = 0
|
||||
try:
|
||||
for index, df in enumerate(chunks):
|
||||
if df.is_empty():
|
||||
continue
|
||||
for date_df in df.partition_by("date"):
|
||||
dt = date_df["date"][0]
|
||||
ds = dt.isoformat() if hasattr(dt, "isoformat") else str(dt)
|
||||
out = root / f"date={ds}" / f"part-{index}.parquet"
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
date_df.write_parquet(out)
|
||||
written += date_df.height
|
||||
|
||||
for date_dir in sorted(root.glob("date=*")):
|
||||
files = sorted(date_dir.glob("*.parquet"))
|
||||
if files:
|
||||
repo.append_daily(pl.scan_parquet(files).collect(engine="streaming"))
|
||||
finally:
|
||||
shutil.rmtree(root, ignore_errors=True)
|
||||
with contextlib.suppress(OSError):
|
||||
root.parent.rmdir()
|
||||
|
||||
return written
|
||||
|
||||
|
||||
def _sweep_stale_daily_staging(staging_base, max_age_s: int = 24 * 60 * 60) -> None:
|
||||
"""清理崩溃遗留的旧同步目录,不碰仍可能活跃的新目录。"""
|
||||
if not staging_base.exists():
|
||||
return
|
||||
cutoff = time.time() - max_age_s
|
||||
for run_dir in staging_base.iterdir():
|
||||
try:
|
||||
if run_dir.is_dir() and run_dir.stat().st_mtime < cutoff:
|
||||
shutil.rmtree(run_dir)
|
||||
except OSError:
|
||||
logger.warning("failed to clean stale daily staging: %s", run_dir)
|
||||
|
||||
|
||||
def sync_daily_by_quotes(repo: KlineRepository) -> int:
|
||||
"""用实时行情接口拉全市场当日数据,覆写 kline_daily 今天分区。
|
||||
|
||||
@@ -260,11 +318,15 @@ def sync_daily_by_quotes(repo: KlineRepository) -> int:
|
||||
"close": q.get("last_price"),
|
||||
"volume": q.get("volume"),
|
||||
"amount": q.get("amount"),
|
||||
# 快照时刻标记: data_integrity 靠 quote_ts 区分盘中快照与盘后权威历史,
|
||||
# 缺失会让盘中覆写的分区在停机后被当成完整历史, 永远不进修复。
|
||||
"quote_ts": q.get("timestamp"),
|
||||
})
|
||||
|
||||
df = pl.DataFrame(records)
|
||||
if df.is_empty():
|
||||
return 0
|
||||
df = df.with_columns(pl.col("quote_ts").cast(pl.Int64, strict=False))
|
||||
|
||||
# 分区日期用北京交易日 (与 quote_service._build_daily 的 cn_today 一致),
|
||||
# 避免 UTC 服务器在盘中把日分区写成服务器本地日期。
|
||||
@@ -703,10 +765,8 @@ def _try_custom_minute(
|
||||
(None, True) → 未配自定义源 / 未配 minute dataset / 自定义源异常 → 走 TickFlow
|
||||
(df, False) → 自定义源成功(含空 df) → 直接用, 不回退
|
||||
|
||||
降级策略 (C): 自定义源异常时无条件 fall through 到 TickFlow,
|
||||
由 TickFlow 路径自身 try/except 兜底。Pro+ 用户 TickFlow 成功返回数据,
|
||||
None 档用户 TickFlow 失败返回空。不显式判断 tier, 避免 #126 augmented
|
||||
capability 逻辑干扰。
|
||||
自定义源异常时返回 fallback=True。单股拉取调用方另行检查 TickFlow 原生
|
||||
能力, 避免自定义源增广能力误放行无权限请求。
|
||||
|
||||
resolver 异常边界由 _resolve_minute_provider 统一兜底; 业务调用
|
||||
(provider.get_minute) 仍在本函数 try 块内, 与 resolver 异常分离
|
||||
@@ -1162,8 +1222,14 @@ def fetch_minute_single(
|
||||
symbol: str,
|
||||
trade_date: date,
|
||||
asset_type: AssetType = "stock",
|
||||
*,
|
||||
capset: CapabilitySet,
|
||||
) -> pl.DataFrame:
|
||||
"""实时拉取单股单日分钟 K(不写入本地)。优先自定义分钟源, 回退 TickFlow。"""
|
||||
"""实时拉取单股单日分钟 K(不写入本地)。
|
||||
|
||||
优先使用当前自定义分钟源。仅当 TickFlow 原生单股分钟能力存在时才允许
|
||||
回退 TickFlow; 自定义源增广只授予 batch 能力, 不会误放行该回退路径。
|
||||
"""
|
||||
from datetime import datetime
|
||||
# 北京时间窗口必须带时区: naive datetime 会被 .timestamp() 按服务器本地时区解释,
|
||||
# UTC 容器上窗口整体偏移 8 小时, 分时补拉必然为空。
|
||||
@@ -1180,6 +1246,9 @@ def fetch_minute_single(
|
||||
# 见 sync_minute_batch 同分支注释: df 在此必非 None。
|
||||
return df if df is not None else pl.DataFrame()
|
||||
|
||||
if not capset.has(Cap.KLINE_MINUTE_BY_SYMBOL):
|
||||
return pl.DataFrame()
|
||||
|
||||
tf = get_client()
|
||||
try:
|
||||
raw = tf.klines.batch(
|
||||
@@ -1212,29 +1281,38 @@ def fetch_adj_factor_single(symbol: str) -> pl.DataFrame:
|
||||
return _normalize_adj_factor(raw)
|
||||
|
||||
|
||||
def _as_beijing(d: datetime) -> datetime:
|
||||
"""落盘的分钟 datetime 是北京墙钟 naive, 带上北京时区再交给取数窗口。
|
||||
|
||||
naive 值经 _datetime_to_ms 会被 .timestamp() 按服务器本地时区解释, 与同
|
||||
窗口另一端的服务器本地时间混用后整体错位 (UTC 容器上错 8 小时)。
|
||||
"""
|
||||
return d if d.tzinfo is not None else d.replace(tzinfo=CN_TZ)
|
||||
|
||||
|
||||
def _latest_minute_datetime(repo: KlineRepository) -> datetime | None:
|
||||
"""本地分钟 K 数据的最新时间。"""
|
||||
"""本地分钟 K 数据的最新时间 (北京时区)。"""
|
||||
try:
|
||||
res = repo.execute_one("SELECT max(datetime) FROM kline_minute")
|
||||
if res and res[0]:
|
||||
d = res[0]
|
||||
if isinstance(d, datetime):
|
||||
return d
|
||||
return datetime.fromisoformat(str(d))
|
||||
return _as_beijing(d)
|
||||
return _as_beijing(datetime.fromisoformat(str(d)))
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def _earliest_minute_datetime(repo: KlineRepository) -> datetime | None:
|
||||
"""本地分钟 K 数据的最早时间 (用于向前扩展的起点)。"""
|
||||
"""本地分钟 K 数据的最早时间 (北京时区, 用于向前扩展的起点)。"""
|
||||
try:
|
||||
res = repo.execute_one("SELECT min(datetime) FROM kline_minute")
|
||||
if res and res[0]:
|
||||
d = res[0]
|
||||
if isinstance(d, datetime):
|
||||
return d
|
||||
return datetime.fromisoformat(str(d))
|
||||
return _as_beijing(d)
|
||||
return _as_beijing(datetime.fromisoformat(str(d)))
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
return None
|
||||
@@ -1356,7 +1434,9 @@ def sync_and_persist_minute(
|
||||
# 迁移:旧版按 symbol= 分区转为 date= 分区
|
||||
_migrate_symbol_to_date_partition(repo)
|
||||
|
||||
now = datetime.now()
|
||||
# 窗口两端统一为北京时区: 起止点会与本地分钟 K 的北京墙钟混用, 用服务器
|
||||
# 本地时间会让窗口整体错位 (UTC 容器上起点晚于终点, 增量补拉一个请求都发不出)。
|
||||
now = cn_now()
|
||||
|
||||
if extend_backward:
|
||||
# 向前扩展模式: 从本地最早数据往前补, 叠加已有数据避免缺口。
|
||||
|
||||
@@ -141,11 +141,8 @@ def compute_mainline_range(repo, data_dir: Path, start: date, end: date,
|
||||
if not enriched_dir.exists():
|
||||
return pl.DataFrame()
|
||||
|
||||
# 兼容返回裸 DataFrame 的实现: 元组解包会把两列 DataFrame 拆成两个 Series,
|
||||
# Series.is_empty() 能通过但后续 group_by 报 'Series' object has no attribute
|
||||
# 'group_by'(用户反馈的重算偶发报错), 故按实际形态取值而不盲目解包
|
||||
loaded = _load_concept_map_df(repo, kind)
|
||||
map_df = loaded[0] if isinstance(loaded, tuple) else loaded
|
||||
# _load_concept_map_df 恒返回 (map_df, count), 命中缓存不再返回裸 DataFrame (#186)
|
||||
map_df, _ = _load_concept_map_df(repo, kind)
|
||||
if map_df.is_empty():
|
||||
return pl.DataFrame()
|
||||
|
||||
|
||||
@@ -237,6 +237,16 @@ def _symbol_keys(row: dict, config: ExtConfig) -> list[str]:
|
||||
return keys
|
||||
|
||||
|
||||
def _leader_sort_key(row: dict) -> float:
|
||||
"""领涨股排序键: 缺涨跌幅的成分股排最后。
|
||||
|
||||
0.00% 是有效涨跌幅, 不能与"无行情"合并成同一个哨兵值 —— 板块整体下跌时
|
||||
平盘股就是领涨股。
|
||||
"""
|
||||
value = _finite(row.get("change_pct"))
|
||||
return value if value is not None else float("-inf")
|
||||
|
||||
|
||||
def _dimension_rank(rows: list[dict], repo, kind: str, limit: int = 5, level: int | None = None) -> dict:
|
||||
if not rows:
|
||||
return {"leading": [], "lagging": []}
|
||||
@@ -281,7 +291,7 @@ def _dimension_rank(rows: list[dict], repo, kind: str, limit: int = 5, level: in
|
||||
changes = [v for v in changes if v is not None]
|
||||
if not changes:
|
||||
continue
|
||||
leader = max(stocks, key=lambda s: _finite(s.get("change_pct")) or -999)
|
||||
leader = max(stocks, key=_leader_sort_key)
|
||||
items.append({
|
||||
"name": name,
|
||||
"count": len(stocks),
|
||||
|
||||
@@ -237,10 +237,10 @@ def _build_user_prompt(overview: dict, news: list[dict], focus: str, lhb_context
|
||||
"消息催化一节请直接从量价异动给出可能的催化逻辑结论,不要编造具体消息,也不要复述本说明。)",
|
||||
])
|
||||
|
||||
from app.services.ai_provider import sanitize_focus
|
||||
safe_focus = sanitize_focus(focus)
|
||||
if safe_focus:
|
||||
parts.extend(["", f"本次复盘请特别关注: {safe_focus}"])
|
||||
from app.services.ai_provider import build_focus_instruction
|
||||
focus_instruction = build_focus_instruction(focus, report_name="大盘复盘报告")
|
||||
if focus_instruction:
|
||||
parts.extend(["", focus_instruction])
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
@@ -336,6 +336,7 @@ async def recap_market_stream(
|
||||
temperature=0.5,
|
||||
# 不限制输出(推理模型思考 token 计入预算, 见 ai_provider.stream_ai_text)
|
||||
max_tokens=None,
|
||||
prefer_final_answer=True,
|
||||
):
|
||||
got_content = True
|
||||
yield json.dumps({"type": "delta", "content": delta}, ensure_ascii=False)
|
||||
|
||||
@@ -19,9 +19,10 @@ import logging
|
||||
import os
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Literal, TypeVar
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -479,6 +480,28 @@ def _duration_s(j: dict[str, Any]) -> float | None:
|
||||
# 进程内单例
|
||||
job_store = JobStore()
|
||||
|
||||
_Result = TypeVar("_Result")
|
||||
|
||||
|
||||
def run_with_capacity(job_id: str, fn: Callable[[], _Result]) -> _Result:
|
||||
"""Wait in the worker, keeping the reservation until its real execution ends."""
|
||||
from app.services.heavy_job_limiter import (
|
||||
HeavyJobCancelledError,
|
||||
shared_heavy_job_limiter,
|
||||
)
|
||||
|
||||
job_store.progress(job_id, "init", 0, "等待其他计算任务完成…")
|
||||
with _CANCEL_FLAGS_LOCK:
|
||||
cancel_event = _CANCEL_FLAGS.get(job_id)
|
||||
try:
|
||||
with shared_heavy_job_limiter.slot("exclusive", cancel_event=cancel_event):
|
||||
if is_cancelled(job_id):
|
||||
raise JobCancelledError(job_id)
|
||||
job_store.start(job_id)
|
||||
return fn()
|
||||
except HeavyJobCancelledError as exc:
|
||||
raise JobCancelledError(job_id) from exc
|
||||
|
||||
|
||||
# ================================================================
|
||||
# 重任务互斥执行槽 — 防「僵尸并发」, 带所有权 token
|
||||
|
||||
@@ -88,13 +88,12 @@ def get_realtime_quote_interval() -> float:
|
||||
|
||||
|
||||
def set_realtime_quote_interval(interval: float) -> float:
|
||||
"""保存行情轮询间隔(不在此做 min/max 校验,由调用方按档位限制)。"""
|
||||
current = load()
|
||||
current["realtime_quote_interval"] = interval
|
||||
_path().write_text(
|
||||
json.dumps(current, indent=2, ensure_ascii=False), encoding="utf-8",
|
||||
)
|
||||
_invalidate_cache()
|
||||
"""保存行情轮询间隔(不在此做 min/max 校验,由调用方按档位限制)。
|
||||
|
||||
走 save() 而不是自己 load + write_text: 锁外的 read-modify-write 会用旧快照
|
||||
整体覆盖文件, 把并发写入的另一个偏好丢掉 (见 save 的 docstring)。
|
||||
"""
|
||||
save({"realtime_quote_interval": interval})
|
||||
return interval
|
||||
|
||||
|
||||
@@ -248,7 +247,7 @@ def get_data_source_long_job_timeout_s() -> int:
|
||||
|
||||
|
||||
def get_minute_batch_compress() -> bool:
|
||||
"""分时批量响应是否启用 gzip 传输压缩。默认开启 (公网部署传输是大头);
|
||||
"""分时详情与批量响应是否启用 gzip 传输压缩。默认开启 (公网部署传输是大头);
|
||||
本机/内网可关闭省服务端 CPU。每次请求即时读取, 开关保存后立即生效。
|
||||
"""
|
||||
raw = load().get("minute_batch_compress", True)
|
||||
@@ -256,7 +255,7 @@ def get_minute_batch_compress() -> bool:
|
||||
|
||||
|
||||
def get_daily_batch_compress() -> bool:
|
||||
"""日K批量响应是否启用 gzip 传输压缩 (与分时各自独立配置)。默认开启。"""
|
||||
"""日K详情与批量响应是否启用 gzip 传输压缩 (与分时各自独立配置)。默认开启。"""
|
||||
raw = load().get("daily_batch_compress", True)
|
||||
return bool(raw)
|
||||
|
||||
@@ -308,8 +307,8 @@ def get_financial_provider() -> str:
|
||||
# ===== 盘后管道拉取内容开关 (A股 / ETF / 指数 独立控制) =====
|
||||
|
||||
def get_pipeline_pull_a_share() -> bool:
|
||||
"""A 股日K固定拉取。"""
|
||||
return True
|
||||
"""是否拉取 A 股日K。默认 True。"""
|
||||
return load().get("pipeline_pull_a_share", True)
|
||||
|
||||
|
||||
def get_pipeline_pull_etf() -> bool:
|
||||
@@ -453,7 +452,7 @@ def set_mainline_filter_config(cfg: dict) -> dict:
|
||||
return get_mainline_filter_config()
|
||||
|
||||
|
||||
_PIPELINE_PULL_KEYS = ("pipeline_pull_etf", "pipeline_pull_index")
|
||||
_PIPELINE_PULL_KEYS = ("pipeline_pull_a_share", "pipeline_pull_etf", "pipeline_pull_index")
|
||||
|
||||
|
||||
def get_pipeline_pull_types() -> dict:
|
||||
@@ -487,17 +486,23 @@ def set_pipeline_index_symbols(symbols: str) -> str:
|
||||
|
||||
|
||||
def get_pipeline_schedule() -> dict:
|
||||
"""返回盘后管道调度时间 {"hour": 15, "minute": 30}。"""
|
||||
d = load().get("pipeline_schedule", {"hour": 15, "minute": 30})
|
||||
return {"hour": d.get("hour", 15), "minute": d.get("minute", 30)}
|
||||
"""返回盘后管道调度时间 {"hour": 15, "minute": 35}。
|
||||
|
||||
默认 15:35 而非 15:30 整: 盘后固定价交易 15:30 才彻底结束, 且供应商
|
||||
聚合含盘后量的官方日K需要时间 —— 整点即拉可能写入不含盘后成交的
|
||||
日线, 也与 quote 定版重试窗口终点 (15:30) 精确重合。留 5 分钟缓冲。
|
||||
"""
|
||||
d = load().get("pipeline_schedule", {"hour": 15, "minute": 35})
|
||||
return {"hour": d.get("hour", 15), "minute": d.get("minute", 35)}
|
||||
|
||||
|
||||
def set_pipeline_schedule(hour: int, minute: int) -> dict:
|
||||
h = max(0, min(23, hour))
|
||||
m = max(0, min(59, minute))
|
||||
# 盘后不早于 15:00
|
||||
if h * 60 + m < 15 * 60:
|
||||
h, m = 15, 0
|
||||
# 盘后管道不早于 15:35: 15:30 盘后固定价才终止 (量/额此前仍会变),
|
||||
# 且供应商官方日线定稿需要缓冲 —— 更早启动可能固化不含盘后量的当日分区
|
||||
if h * 60 + m < 15 * 60 + 35:
|
||||
h, m = 15, 35
|
||||
save({"pipeline_schedule": {"hour": h, "minute": m}})
|
||||
return {"hour": h, "minute": m}
|
||||
|
||||
@@ -580,21 +585,23 @@ def set_depth_finalize_time(hour: int, minute: int) -> dict:
|
||||
return {"hour": h, "minute": m}
|
||||
|
||||
|
||||
# 复盘推送可选渠道白名单 (企业微信已实现, 与飞书并列)
|
||||
# 监控与复盘共用的外部推送渠道白名单。
|
||||
# 多选: 不推送 = 空数组, 而非 'none'
|
||||
REVIEW_PUSH_CHANNELS = {"feishu", "wecom"}
|
||||
PUSH_CHANNELS = {"feishu", "wecom", "custom", "email"}
|
||||
|
||||
|
||||
def get_review_schedule() -> dict:
|
||||
"""定时复盘调度 {"enabled": False, "hour": 15, "minute": 10}。默认关闭。
|
||||
"""定时复盘调度 {"enabled": False, "hour": 15, "minute": 40}。默认关闭。
|
||||
|
||||
A股 15:00 收盘, 默认时间设为 15:10(收盘后即时复盘), 强制下限 15:00。
|
||||
默认 15:40: 盘后管道默认 15:35 启动, 留 5 分钟缓冲, 复盘使用管道
|
||||
产出的最终口径数据 (含盘后量校正的日K/enriched)。强制下限 15:00 —
|
||||
偏好收盘后即时复盘 (走实时快照缓存, 不等管道) 的用户可自行调早。
|
||||
"""
|
||||
d = load().get("review_schedule", {"enabled": False, "hour": 15, "minute": 10})
|
||||
d = load().get("review_schedule", {"enabled": False, "hour": 15, "minute": 40})
|
||||
return {
|
||||
"enabled": bool(d.get("enabled", False)),
|
||||
"hour": d.get("hour", 15),
|
||||
"minute": d.get("minute", 10),
|
||||
"minute": d.get("minute", 40),
|
||||
}
|
||||
|
||||
|
||||
@@ -662,7 +669,7 @@ def get_review_push_channels() -> list[str]:
|
||||
d = load()
|
||||
raw = d.get("review_push_channels")
|
||||
if isinstance(raw, list):
|
||||
return [c for c in raw if c in REVIEW_PUSH_CHANNELS]
|
||||
return [c for c in raw if c in PUSH_CHANNELS]
|
||||
# 兼容老单选字符串
|
||||
if d.get("review_push_channel") == "feishu":
|
||||
return ["feishu"]
|
||||
@@ -677,13 +684,33 @@ def set_review_push_channels(channels: list[str]) -> list[str]:
|
||||
seen: set[str] = set()
|
||||
cleaned: list[str] = []
|
||||
for c in channels or []:
|
||||
if c in REVIEW_PUSH_CHANNELS and c not in seen:
|
||||
if c in PUSH_CHANNELS and c not in seen:
|
||||
seen.add(c)
|
||||
cleaned.append(c)
|
||||
save({"review_push_channels": cleaned})
|
||||
return cleaned
|
||||
|
||||
|
||||
REVIEW_PUSH_MODES = frozenset({"auto", "manual"})
|
||||
|
||||
|
||||
def get_review_push_mode() -> str:
|
||||
"""复盘推送触发方式: auto=归档后自动推; manual=仅显式 push。默认 manual。
|
||||
|
||||
定时复盘与手动保存复盘共用此开关。manual 时定时路径只归档不推送,
|
||||
手动路径需 save_report 显式传 push=True 才推。
|
||||
"""
|
||||
mode = load().get("review_push_mode", "manual")
|
||||
return mode if mode in REVIEW_PUSH_MODES else "manual"
|
||||
|
||||
|
||||
def set_review_push_mode(mode: str) -> str:
|
||||
"""保存复盘推送触发方式, 白名单外的值回退 manual。"""
|
||||
mode = mode if mode in REVIEW_PUSH_MODES else "manual"
|
||||
save({"review_push_mode": mode})
|
||||
return mode
|
||||
|
||||
|
||||
|
||||
# ===== 实时监控 =====
|
||||
|
||||
@@ -797,6 +824,71 @@ def set_wecom_webhook_url(url: str) -> str:
|
||||
return get_wecom_webhook_url()
|
||||
|
||||
|
||||
def get_custom_webhook_url() -> str:
|
||||
"""Generic third-party JSON Webhook URL shared by enabled rules and reviews."""
|
||||
return str(load().get("custom_webhook_url") or "")
|
||||
|
||||
|
||||
def set_custom_webhook_url(url: str) -> str:
|
||||
"""Persist or clear the generic third-party JSON Webhook URL."""
|
||||
value = str(url or "").strip()
|
||||
save({"custom_webhook_url": value})
|
||||
return value
|
||||
|
||||
|
||||
_EMAIL_SMTP_DEFAULTS = {
|
||||
"host": "",
|
||||
"port": 465,
|
||||
"security": "ssl",
|
||||
"username": "",
|
||||
"from_address": "",
|
||||
"to_addresses": [],
|
||||
}
|
||||
|
||||
|
||||
def get_email_smtp_config() -> dict:
|
||||
"""Return non-secret SMTP settings for the email notification channel."""
|
||||
raw = load().get("email_smtp_config")
|
||||
if not isinstance(raw, dict):
|
||||
raw = {}
|
||||
security = raw.get("security", _EMAIL_SMTP_DEFAULTS["security"])
|
||||
if security not in {"ssl", "starttls", "none"}:
|
||||
security = _EMAIL_SMTP_DEFAULTS["security"]
|
||||
try:
|
||||
port = int(raw.get("port", _EMAIL_SMTP_DEFAULTS["port"]))
|
||||
except (TypeError, ValueError):
|
||||
port = _EMAIL_SMTP_DEFAULTS["port"]
|
||||
if not 1 <= port <= 65535:
|
||||
port = _EMAIL_SMTP_DEFAULTS["port"]
|
||||
recipients = raw.get("to_addresses")
|
||||
if not isinstance(recipients, list):
|
||||
recipients = []
|
||||
return {
|
||||
"host": str(raw.get("host") or "").strip(),
|
||||
"port": port,
|
||||
"security": security,
|
||||
"username": str(raw.get("username") or "").strip(),
|
||||
"from_address": str(raw.get("from_address") or "").strip(),
|
||||
"to_addresses": [str(item).strip() for item in recipients if str(item).strip()],
|
||||
}
|
||||
|
||||
|
||||
def set_email_smtp_config(config: dict) -> dict:
|
||||
"""Atomically persist the non-secret SMTP configuration group."""
|
||||
normalized = {
|
||||
"host": str(config.get("host") or "").strip(),
|
||||
"port": int(config.get("port", 465)),
|
||||
"security": str(config.get("security") or "ssl"),
|
||||
"username": str(config.get("username") or "").strip(),
|
||||
"from_address": str(config.get("from_address") or "").strip(),
|
||||
"to_addresses": [
|
||||
str(item).strip() for item in config.get("to_addresses", []) if str(item).strip()
|
||||
],
|
||||
}
|
||||
save({"email_smtp_config": normalized})
|
||||
return get_email_smtp_config()
|
||||
|
||||
|
||||
# ===== 企业微信智能机器人 (API 模式 / 长连接) =====
|
||||
|
||||
|
||||
@@ -863,7 +955,7 @@ def get_webhook_default_channels() -> list[str]:
|
||||
d = load()
|
||||
raw = d.get("webhook_default_channels")
|
||||
if isinstance(raw, list):
|
||||
return [c for c in raw if c in REVIEW_PUSH_CHANNELS]
|
||||
return [c for c in raw if c in PUSH_CHANNELS]
|
||||
# 兼容老布尔开关 (勾选即双推)
|
||||
if d.get("webhook_enabled_default") is True:
|
||||
return ["feishu", "wecom"]
|
||||
@@ -875,7 +967,7 @@ def set_webhook_default_channels(channels: list[str]) -> list[str]:
|
||||
seen: set[str] = set()
|
||||
cleaned: list[str] = []
|
||||
for c in channels or []:
|
||||
if c in REVIEW_PUSH_CHANNELS and c not in seen:
|
||||
if c in PUSH_CHANNELS and c not in seen:
|
||||
seen.add(c)
|
||||
cleaned.append(c)
|
||||
save({"webhook_default_channels": cleaned})
|
||||
|
||||
@@ -33,10 +33,40 @@ from datetime import date, datetime, time as dt_time
|
||||
|
||||
import polars as pl
|
||||
|
||||
from app.market_time import cn_now, cn_today
|
||||
from app.market_time import CN_TZ, cn_now, cn_today
|
||||
from app.parquet import scan_daily_parquet
|
||||
from app.polars_guard import guarded_collect
|
||||
from app.services.index_const import CORE_INDEX_SYMBOLS
|
||||
from app.strategy.intraday_signals import IntradaySignalEvaluator
|
||||
from app.strategy.monitor import format_alert_quote
|
||||
|
||||
# 告警来源 → 中文标签 (webhook 标题 / 系统通知标题共用)
|
||||
SOURCE_LABELS = {
|
||||
"strategy": "策略", "signal": "信号", "price": "价格",
|
||||
"market": "异动", "ladder": "连板梯队", "sector": "板块",
|
||||
"volume_delta": "放量", "abnormal": "异动", "date": "日期提醒",
|
||||
}
|
||||
|
||||
# final 定版确认容差: 快照时间戳允许早于边界 5s 内 (供应商时间戳精度不一)
|
||||
_FINAL_CONFIRM_SLACK_MS = 5_000
|
||||
|
||||
# final 定版边界与重试窗口终点 (北京时间)。收盘窗口终点 15:30, 恰与盘后管道
|
||||
# 启动同时: 管道运行期间轮询本就被暂停, 此后未确认的定版不再写盘, 当日分区
|
||||
# 由管道按官方日线值级校正 —— 避免定版重试与权威重建互相覆盖。
|
||||
_FINAL_BOUNDARY = {"morning_final": dt_time(11, 30), "close_final": dt_time(15, 0)}
|
||||
_FINAL_DEADLINE = {"morning_final": dt_time(12, 10), "close_final": dt_time(15, 30)}
|
||||
|
||||
|
||||
def _body_with_quote(body: str, ev: dict) -> str:
|
||||
"""推送正文尾部补上触发时的现价/涨跌幅 (日期提醒无行情, 自然为空)。
|
||||
|
||||
默认告警的 message 已由引擎拼过引语 (monitor._default_message), 这里仅在正文
|
||||
尚未带引语时追加, 避免「现价」出现两遍 (自定义 message 的规则则补上这一句)。
|
||||
"""
|
||||
quote_tail = format_alert_quote(ev.get("price"), ev.get("change_pct"))
|
||||
if not quote_tail or body.endswith(quote_tail):
|
||||
return body
|
||||
return f"{body} · {quote_tail}"
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -220,6 +250,8 @@ class QuoteService:
|
||||
# 午休/收盘最终同步状态: 到边界后必须成功拉取一版行情, 再进入休盘态。
|
||||
self._final_sync_done: set[tuple[date, str]] = set()
|
||||
self._final_sync_failed: dict[tuple[date, str], str] = {}
|
||||
# 最近一次 final 定版拉取是否取得边界后快照 (None=非 final 拉取)
|
||||
self._last_final_confirmed: bool | None = None
|
||||
self._holiday_active = False # 交易日探针当前是否判休市 (日志去重)
|
||||
# 轮询放量 (volume_delta 规则): 上一轮全市场股票快照的 (累计成交量[手], 累计成交额[元])。
|
||||
# 每轮全量快照后更新 (含非连续竞价时段, 保证 13:00 恢复时 prev 是 12:59
|
||||
@@ -542,8 +574,17 @@ class QuoteService:
|
||||
}
|
||||
|
||||
def refresh(self) -> dict:
|
||||
"""手动触发一次行情拉取。"""
|
||||
self._fetch_quotes()
|
||||
"""手动触发一次行情拉取。
|
||||
|
||||
午休/收盘定版阶段同样走边界确认: 避免盘后手动刷新把竞价前的陈旧收盘价
|
||||
重新写回当日分区, 覆盖盘后管道按官方日线重建的结果。
|
||||
"""
|
||||
phase = self._market_phase()
|
||||
is_final = phase in {"morning_final", "close_final"}
|
||||
self._fetch_quotes(
|
||||
final=is_final,
|
||||
final_boundary_ms=self._final_boundary_ms(phase) if is_final else None,
|
||||
)
|
||||
return self.status()
|
||||
|
||||
# ================================================================
|
||||
@@ -559,16 +600,36 @@ class QuoteService:
|
||||
phase = self._market_phase()
|
||||
if self._should_fetch_for_phase(phase):
|
||||
is_final = phase in {"morning_final", "close_final"}
|
||||
ok = self._fetch_quotes(final=is_final)
|
||||
ok = self._fetch_quotes(
|
||||
final=is_final,
|
||||
final_boundary_ms=self._final_boundary_ms(phase),
|
||||
)
|
||||
if is_final:
|
||||
key = self._final_sync_key(phase)
|
||||
if key and ok:
|
||||
label = "午休" if phase == "morning_final" else "收盘"
|
||||
if key and ok and self._last_final_confirmed:
|
||||
self._final_sync_done.add(key)
|
||||
self._final_sync_failed.pop(key, None)
|
||||
logger.info("%s 最终行情同步完成, 进入休盘态", "午休" if phase == "morning_final" else "收盘")
|
||||
logger.info("%s 最终行情同步完成 (快照时间戳已达边界), 进入休盘态", label)
|
||||
elif key and self._past_final_deadline(phase):
|
||||
# 重试窗口结束仍未取得边界后快照: 接受现状停止轮询。
|
||||
# 实测有实时源收盘后长期返回竞价前旧价 (快照时间戳可信但价格不更新),
|
||||
# 此时盲目落盘只会固化旧价 —— 交由 15:30 盘后管道按官方日线校正。
|
||||
self._final_sync_done.add(key)
|
||||
self._final_sync_failed[key] = (
|
||||
"fetch_failed" if not ok else "unconfirmed_snapshot"
|
||||
)
|
||||
logger.warning(
|
||||
"%s 定版窗口结束仍未取得边界后快照 (%s), 停止轮询; "
|
||||
"当日分区由盘后管道按官方日线值级校正",
|
||||
label, "拉取失败" if not ok else "快照未确认",
|
||||
)
|
||||
elif key:
|
||||
self._final_sync_failed[key] = "fetch_failed"
|
||||
logger.warning("%s 最终行情同步失败, 将继续重试", "午休" if phase == "morning_final" else "收盘")
|
||||
self._final_sync_failed[key] = (
|
||||
"fetch_failed" if not ok else "unconfirmed_snapshot"
|
||||
)
|
||||
if not ok:
|
||||
logger.warning("%s 最终行情同步失败, 将继续重试", label)
|
||||
else:
|
||||
logger.debug("非轮询阶段(%s), 跳过行情轮询", phase)
|
||||
except Exception as e: # noqa: BLE001
|
||||
@@ -579,16 +640,20 @@ class QuoteService:
|
||||
time.sleep(0.5)
|
||||
waited += 0.5
|
||||
|
||||
def _fetch_quotes(self, *, final: bool = False) -> bool:
|
||||
"""拉取行情。加锁串行化 (后台轮询 vs 手动 refresh)。返回本轮是否成功更新。"""
|
||||
def _fetch_quotes(self, *, final: bool = False, final_boundary_ms: int | None = None) -> bool:
|
||||
"""拉取行情。加锁串行化 (后台轮询 vs 手动 refresh)。返回本轮是否成功更新。
|
||||
|
||||
final_boundary_ms: final 定版的边界时间戳 (ms)。传入时快照时间戳未达边界
|
||||
的本轮不落盘 (见 _process_full_market_records)。
|
||||
"""
|
||||
with self._fetch_lock:
|
||||
before = self._fetched_at
|
||||
if final:
|
||||
logger.info("最终行情同步开始")
|
||||
self._fetch_full_market_quotes()
|
||||
self._fetch_full_market_quotes(final_boundary_ms=final_boundary_ms)
|
||||
return self._fetched_at > before
|
||||
|
||||
def _fetch_full_market_quotes(self) -> None:
|
||||
def _fetch_full_market_quotes(self, final_boundary_ms: int | None = None) -> None:
|
||||
"""拉取全市场行情 → 写 daily + 计算 enriched + 更新缓存。"""
|
||||
from app.services import preferences
|
||||
|
||||
@@ -604,17 +669,37 @@ class QuoteService:
|
||||
# 指数补充: A 股快照通常不含指数。插件可选实现
|
||||
# get_realtime_indices(symbols) 用独立端点补拉 (如 fuyao 指数快照);
|
||||
# 未实现的源指数缓存为空, 由日K兜底接管。
|
||||
replace_index_cache = True
|
||||
fetch_indices = getattr(provider, "get_realtime_indices", None)
|
||||
if callable(fetch_indices):
|
||||
wanted = sorted(set(CORE_INDEX_SYMBOLS) | self._collect_monitor_index_symbols())
|
||||
# 偏离值基准指数 (科创50/创业板综指等) 一并拉取, 供盘中
|
||||
# attach_deviation_columns_today 实时外推; 展示层仍按核心
|
||||
# 四只过滤, 多拉的指数不进侧栏。
|
||||
from app.indicators.pipeline import BENCHMARK_INDEX_SYMBOLS
|
||||
wanted = sorted(
|
||||
set(CORE_INDEX_SYMBOLS)
|
||||
| BENCHMARK_INDEX_SYMBOLS
|
||||
| self._collect_monitor_index_symbols()
|
||||
)
|
||||
try:
|
||||
records = records + (fetch_indices(wanted) or [])
|
||||
fetched_indices = fetch_indices(wanted)
|
||||
if fetched_indices is None:
|
||||
replace_index_cache = False
|
||||
else:
|
||||
records = records + fetched_indices
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("自定义源指数行情拉取失败: %s", e)
|
||||
replace_index_cache = False
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("自定义实时行情拉取失败: %s", e)
|
||||
return
|
||||
self._process_full_market_records(records, t0=t0, now_ts=now_ts)
|
||||
self._process_full_market_records(
|
||||
records,
|
||||
t0=t0,
|
||||
now_ts=now_ts,
|
||||
replace_index_cache=replace_index_cache,
|
||||
final_boundary_ms=final_boundary_ms,
|
||||
)
|
||||
return
|
||||
# 自定义源未配置 realtime → 回退 TickFlow
|
||||
|
||||
@@ -653,8 +738,11 @@ class QuoteService:
|
||||
logger.info("拉取全市场行情 (universes=%s, SDK超时=30s×重试3)", universes)
|
||||
resp.extend(tf.quotes.get_by_universes(universes=universes) or [])
|
||||
logger.info("全市场行情拉取完成: %d 条 (%.2fs)", len(resp), time.perf_counter() - _u0)
|
||||
# 指数: 固定核心四只 + 监控规则标的, 按码显式拉取
|
||||
_core_syms = sorted(core_index_symbols | monitor_index_symbols)
|
||||
# 指数: 固定核心四只 + 偏离值基准指数 + 监控规则标的, 按码显式拉取
|
||||
from app.indicators.pipeline import BENCHMARK_INDEX_SYMBOLS
|
||||
_core_syms = sorted(
|
||||
core_index_symbols | BENCHMARK_INDEX_SYMBOLS | monitor_index_symbols
|
||||
)
|
||||
if _core_syms:
|
||||
_i0 = time.perf_counter()
|
||||
resp.extend(tf.quotes.get(symbols=_core_syms) or [])
|
||||
@@ -699,10 +787,25 @@ class QuoteService:
|
||||
"session": q.get("session"),
|
||||
})
|
||||
|
||||
self._process_full_market_records(records, t0=t0, now_ts=now_ts)
|
||||
self._process_full_market_records(
|
||||
records, t0=t0, now_ts=now_ts, final_boundary_ms=final_boundary_ms
|
||||
)
|
||||
|
||||
def _process_full_market_records(self, records: list[dict], *, t0: float, now_ts: float) -> None:
|
||||
"""把全市场 records 写盘并增量计算 enriched。"""
|
||||
def _process_full_market_records(
|
||||
self,
|
||||
records: list[dict],
|
||||
*,
|
||||
t0: float,
|
||||
now_ts: float,
|
||||
replace_index_cache: bool = True,
|
||||
final_boundary_ms: int | None = None,
|
||||
) -> None:
|
||||
"""把全市场 records 写盘并增量计算 enriched。
|
||||
|
||||
final_boundary_ms (final 定版边界) 传入时, 快照最大时间戳未达边界的本轮
|
||||
只更新展示缓存, 不写 daily/enriched、不评估监控 —— 防止收盘后数据源仍
|
||||
返回竞价前旧价时把陈旧收盘价固化到当日分区。
|
||||
"""
|
||||
from app.services import preferences
|
||||
all_index_symbols = set(self._repo.get_index_symbol_set()) if self._repo else set()
|
||||
core_index_symbols = set(CORE_INDEX_SYMBOLS)
|
||||
@@ -717,6 +820,16 @@ class QuoteService:
|
||||
logger.warning("行情数据为空")
|
||||
return
|
||||
|
||||
# ---- final 定版确认: 快照最大时间戳达到边界 (含容差) 才允许落盘 ----
|
||||
confirmed_final: bool | None = None
|
||||
if final_boundary_ms is not None:
|
||||
ts_vals = [t for t in (r.get("timestamp") for r in records) if t]
|
||||
max_ts = max(ts_vals) if ts_vals else None
|
||||
confirmed_final = bool(
|
||||
max_ts is not None and max_ts >= final_boundary_ms - _FINAL_CONFIRM_SLACK_MS
|
||||
)
|
||||
self._last_final_confirmed = confirmed_final
|
||||
|
||||
index_records = [r for r in records if r.get("symbol") in all_index_symbols]
|
||||
etf_records = [r for r in records if r.get("symbol") in all_etf_symbols]
|
||||
stock_records = [
|
||||
@@ -733,13 +846,26 @@ class QuoteService:
|
||||
self._fetch_ms = fetch_ms
|
||||
self._fetched_at = fetched_at
|
||||
self._symbol_count = len(stock_records)
|
||||
self._index_symbol_count = len(index_records)
|
||||
self._etf_symbol_count = len(etf_records)
|
||||
self._index_quotes_cache = self._build_index_quotes(index_records)
|
||||
if replace_index_cache:
|
||||
self._index_symbol_count = len(index_records)
|
||||
self._index_quotes_cache = self._build_index_quotes(index_records)
|
||||
else:
|
||||
logger.info("指数本轮获取失败,沿用上轮缓存: %d 只", self._index_symbol_count)
|
||||
|
||||
_persist_last_fetch(fetched_at)
|
||||
logger.info("行情刷新: %d 只股票, %d 只ETF, %d 只指数, 耗时 %.0fms", len(stock_records), len(etf_records), len(index_records), fetch_ms)
|
||||
|
||||
if confirmed_final is False:
|
||||
# 边界前的陈旧快照: 展示缓存已更新, 落盘与监控评估留待边界后快照。
|
||||
# 轮询线程会在定版窗口内持续重试, 窗口结束由 _poll_loop 放弃并告警。
|
||||
logger.info(
|
||||
"final 快照未达定版边界 (max quote_ts=%s, 边界=%s), 本轮跳过落盘",
|
||||
max_ts, final_boundary_ms,
|
||||
)
|
||||
self._broadcast_quote_updated()
|
||||
return
|
||||
|
||||
# 轮询放量状态更新 (volume_delta 规则的差值来源)
|
||||
self._update_volume_delta(stock_records, fetched_at)
|
||||
|
||||
@@ -827,6 +953,23 @@ class QuoteService:
|
||||
result = df.select(select_exprs).with_columns(
|
||||
pl.lit(cn_today()).cast(pl.Date).alias("date"),
|
||||
)
|
||||
# 停牌股回归: 实时源对停牌标的返回停牌前最后一份快照 — OHLCV 全为旧日
|
||||
# 真实值, 仅 timestamp 停在旧日。这类记录不属于当日, 不过滤会把旧日 K 线
|
||||
# 原样复制成当日假蜡烛 (如 301266.SZ 2026-09-04)。按 quote_ts 的北京
|
||||
# 日期归属过滤; 时间戳缺失/为空的源无法判断, 维持原行为保留。
|
||||
if "quote_ts" in result.columns:
|
||||
day_start_ms = int(
|
||||
datetime.combine(cn_today(), dt_time(0, 0), tzinfo=CN_TZ).timestamp() * 1000
|
||||
)
|
||||
result = result.filter(
|
||||
pl.col("quote_ts").is_null()
|
||||
| pl.col("quote_ts").is_between(day_start_ms, day_start_ms + 86_400_000, closed="left")
|
||||
)
|
||||
# 停牌/尚无集合竞价的记录 open/high 均为 0。必须在下方用 close 填充前
|
||||
# 过滤, 否则零成交行会被伪装成有效日K, 并在 batch 同步后作为实时残留
|
||||
# 反复触发历史完整性修复。
|
||||
from app.indicators.pipeline import filter_halt_days
|
||||
result = filter_halt_days(result)
|
||||
# 修复: API 在非交易时段可能返回 open/high/low=0 或 null,
|
||||
# 导致蜡烛从 0 开始。用 close 填充这些异常值。
|
||||
for col in ("open", "high", "low"):
|
||||
@@ -933,6 +1076,20 @@ class QuoteService:
|
||||
return (cn_today(), "close")
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def _final_boundary_ms(cls, phase: str) -> int | None:
|
||||
"""final 阶段定版边界的 epoch ms (按北京时间当日换算, 不依赖服务器时区)。"""
|
||||
b = _FINAL_BOUNDARY.get(phase)
|
||||
if b is None:
|
||||
return None
|
||||
return int(datetime.combine(cn_today(), b, tzinfo=CN_TZ).timestamp() * 1000)
|
||||
|
||||
@classmethod
|
||||
def _past_final_deadline(cls, phase: str) -> bool:
|
||||
"""是否已过 final 重试窗口终点 (用于放弃未确认的定版重试)。"""
|
||||
dl = _FINAL_DEADLINE.get(phase)
|
||||
return dl is not None and cn_now().time() >= dl
|
||||
|
||||
def _holiday_gate(self) -> bool:
|
||||
"""交易日探针门控: 确定休市 → False (停止轮询, 含 final 定版)。
|
||||
|
||||
@@ -1058,6 +1215,12 @@ class QuoteService:
|
||||
rule_events += engine.evaluate_abnormal(_overview.get("rows") or [])
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("异动监控规则评估失败 (不影响其他告警): %s", e)
|
||||
# 日期提醒轮: 纯日历、无行情, 已在盘中; 引擎内按天 cooldown 保证每天一次
|
||||
if engine.has_rule_type("date"):
|
||||
try:
|
||||
rule_events = rule_events + engine.evaluate_date_rules()
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("日期提醒评估失败 (不影响其他告警): %s", e)
|
||||
# ETF 规则轮: 股票快照不含 ETF, 用 ETF enriched 快照单独评估。
|
||||
# 独立 try —— ETF 轮任何异常都不得丢弃本轮已算出的股票告警。
|
||||
# refresh=False —— 不在轮询线程上触发 ETF 冷缓存的同步重算 (缓存由 ETF 实时
|
||||
@@ -1278,9 +1441,20 @@ class QuoteService:
|
||||
prev_close=prev_close,
|
||||
asset_type=asset_type,
|
||||
now=now,
|
||||
signals=self._load_intraday_signal_defs(),
|
||||
)
|
||||
return self._intraday_signal_evaluator.inject(enriched, signals)
|
||||
|
||||
def _load_intraday_signal_defs(self) -> list[dict]:
|
||||
"""加载自定义盘中信号定义(带指纹缓存); 失败时退化为仅内置 4 信号。"""
|
||||
try:
|
||||
from app.strategy import custom_signals
|
||||
|
||||
return custom_signals.load_intraday_all(self._repo.store.data_dir)
|
||||
except Exception as e:
|
||||
logger.warning("load intraday signal defs failed: %s", e)
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _continuous_session_start_ms() -> float:
|
||||
"""当前连续竞价时段的起点 (北京时间 9:30 或 13:00) 的 epoch 毫秒。"""
|
||||
@@ -1402,7 +1576,7 @@ class QuoteService:
|
||||
def _maybe_send_webhook(self, rule_events: list[dict], engine) -> None:
|
||||
"""把告警通过 Webhook 推送到外部 IM (由规则 webhook_channels 指定渠道)。
|
||||
|
||||
- 飞书 / 企业微信任一已配置即生效 (两个都没配才跳过)
|
||||
- 飞书 / 企业微信 / 第三方 Webhook / 邮件均按规则独立选择
|
||||
- 仅推送 webhook_channels 非空的规则触发的告警, 且只投递被勾选的渠道
|
||||
- 失败静默, 不阻断主流程
|
||||
- 去重: 复用 MonitorRuleEngine 的 cooldown, 此处不重复去重
|
||||
@@ -1411,40 +1585,40 @@ class QuoteService:
|
||||
以便反查引擎规则判断是否启用推送。
|
||||
"""
|
||||
try:
|
||||
from app.services import preferences
|
||||
from app.services import webhook_adapter
|
||||
from app import secrets_store
|
||||
from app.services import email_adapter, preferences, webhook_adapter
|
||||
|
||||
feishu_url = preferences.get_feishu_webhook_url()
|
||||
feishu_secret = preferences.get_feishu_webhook_secret()
|
||||
wecom_url = preferences.get_wecom_webhook_url()
|
||||
# 两个通道都没配置才跳过
|
||||
if not feishu_url and not wecom_url:
|
||||
custom_url = preferences.get_custom_webhook_url()
|
||||
custom_secret = secrets_store.get_custom_webhook_secret()
|
||||
email_config = preferences.get_email_smtp_config()
|
||||
email_password = secrets_store.get_email_smtp_password()
|
||||
if not any((feishu_url, wecom_url, custom_url, email_adapter.is_configured(email_config))):
|
||||
return
|
||||
|
||||
# 反查规则, 过滤出启用推送的事件
|
||||
source_labels = {
|
||||
"strategy": "策略", "signal": "信号",
|
||||
"price": "价格", "market": "异动", "ladder": "连板梯队",
|
||||
"sector": "板块", "volume_delta": "放量",
|
||||
}
|
||||
rules = engine.rules if engine is not None else {}
|
||||
enqueued = 0
|
||||
for ev in rule_events:
|
||||
rule = rules.get(ev.get("rule_id"))
|
||||
# webhook_channels 指定命中的渠道 (['feishu'] / ['wecom'] / ['feishu','wecom'] / []).
|
||||
# webhook_channels 指定本规则需要投递的外部渠道。
|
||||
# 空列表 = 该规则不推送。仅推送「渠道已选 + 对应地址已配置」的组合。
|
||||
channels = rule.get("webhook_channels") if rule else None
|
||||
if not channels:
|
||||
continue
|
||||
source = ev.get("source", "")
|
||||
source_label = source_labels.get(source, source or "通知")
|
||||
source_label = SOURCE_LABELS.get(source, source or "通知")
|
||||
symbol = ev.get("symbol") or ""
|
||||
name = ev.get("name") or ""
|
||||
message = ev.get("message") or ""
|
||||
title = source_label
|
||||
body = f"{symbol} {name} {message}".strip() if symbol else (message or name)
|
||||
# 补上触发时的现价/涨跌幅, 让推送可执行 (止损到底触发在哪个价位)
|
||||
body = _body_with_quote(body, ev)
|
||||
# 提交到独立线程池, 不阻塞行情轮询线程 (webhook 慢/重试不拖累实时行情+告警)。
|
||||
# 按渠道独立投递: 飞书 / 企业微信谁被勾选且已配置就推谁。
|
||||
# 按渠道独立投递: 只投递同时“已勾选 + 已配置”的渠道。
|
||||
# 应用内 alerts.jsonl 记录与 SSE 已在前面完成, 不依赖 webhook 成败,
|
||||
# 失败由 webhook_adapter 记 WARNING(可见)。
|
||||
if feishu_url and "feishu" in channels:
|
||||
@@ -1453,6 +1627,26 @@ class QuoteService:
|
||||
if wecom_url and "wecom" in channels:
|
||||
_WEBHOOK_EXECUTOR.submit(webhook_adapter.send_wecom, wecom_url, title, body)
|
||||
enqueued += 1
|
||||
if custom_url and "custom" in channels:
|
||||
_WEBHOOK_EXECUTOR.submit(
|
||||
webhook_adapter.send_custom,
|
||||
custom_url,
|
||||
title,
|
||||
body,
|
||||
"monitor_alert",
|
||||
ev,
|
||||
custom_secret,
|
||||
)
|
||||
enqueued += 1
|
||||
if email_adapter.is_configured(email_config) and "email" in channels:
|
||||
_WEBHOOK_EXECUTOR.submit(
|
||||
email_adapter.send_email,
|
||||
email_config,
|
||||
email_password,
|
||||
title,
|
||||
body,
|
||||
)
|
||||
enqueued += 1
|
||||
if enqueued:
|
||||
logger.info("Webhook 已提交 %d 条 (异步投递, 按渠道独立投递, 失败记 WARNING)", enqueued)
|
||||
except Exception as e: # noqa: BLE001
|
||||
@@ -1476,10 +1670,7 @@ class QuoteService:
|
||||
for ev in all_alerts:
|
||||
# 通知标题: 用 source 分类 (策略/信号/价格/异动)
|
||||
source = ev.get("source", "")
|
||||
source_label = {
|
||||
"strategy": "策略", "signal": "信号",
|
||||
"price": "价格", "market": "异动", "sector": "板块",
|
||||
}.get(source, source or "通知")
|
||||
source_label = SOURCE_LABELS.get(source, source or "通知")
|
||||
|
||||
name = ev.get("name") or ""
|
||||
symbol = ev.get("symbol") or ""
|
||||
@@ -1490,6 +1681,8 @@ class QuoteService:
|
||||
body = f"{symbol} {name} {message}".strip()
|
||||
else:
|
||||
body = message or name
|
||||
# 补上触发时的现价/涨跌幅 (日期提醒无行情, 自然为空)
|
||||
body = _body_with_quote(body, ev)
|
||||
|
||||
title = f"TickFlow · {source_label}"
|
||||
notify_adapter.notify(title, body)
|
||||
@@ -1569,11 +1762,11 @@ class QuoteService:
|
||||
table = {"etf": "kline_etf_daily", "index": "kline_index_daily"}.get(asset_type, "kline_daily")
|
||||
daily_glob = str(self._repo.store.data_dir / table / "**" / "*.parquet")
|
||||
ohlcv_cols = ["symbol", "date", "open", "high", "low", "close", "volume", "amount", "quote_ts"]
|
||||
hist_df = (
|
||||
hist_df = guarded_collect(
|
||||
scan_daily_parquet(daily_glob)
|
||||
.filter(pl.col("date") >= cutoff)
|
||||
.sort(["symbol", "date"])
|
||||
.collect()
|
||||
.sort(["symbol", "date"]),
|
||||
priority="background",
|
||||
)
|
||||
if hist_df.is_empty():
|
||||
return
|
||||
|
||||
@@ -356,17 +356,30 @@ def _compute_batch(repo, enriched_dir, instruments, historical_shares,
|
||||
return df.filter((pl.col("date") >= batch_start) & (pl.col("date") <= batch_end))
|
||||
|
||||
|
||||
def _scan_enriched_fallback(repo, start: date, end: date) -> pl.DataFrame | None:
|
||||
"""缓存不覆盖时的慢路径: scan enriched parquet + 重算所需指标列。
|
||||
def _filter_excluded_symbols(df: pl.DataFrame, excluded_symbols: list[str]) -> pl.DataFrame:
|
||||
if excluded_symbols and "symbol" in df.columns:
|
||||
return df.filter(~pl.col("symbol").str.to_uppercase().is_in(excluded_symbols))
|
||||
return df
|
||||
|
||||
仅在 regime 首次全量回填或缓存未预热时触发。返回含信号列的多日 DataFrame。
|
||||
|
||||
def _scan_enriched_fallback(
|
||||
repo,
|
||||
start: date,
|
||||
end: date,
|
||||
*,
|
||||
index_pct_map: dict | None = None,
|
||||
excluded_symbols: list[str] | None = None,
|
||||
) -> pl.DataFrame | None:
|
||||
"""缓存不覆盖时的慢路径: 分批扫描并直接返回日级环境聚合。
|
||||
|
||||
仅在 regime 首次全量回填或缓存未预热时触发。
|
||||
|
||||
内存控制(关键, 两层优化):
|
||||
1. needed 白名单: regime 只需 change_pct/ma20/涨跌停信号等少数列, 不用 compute_all
|
||||
算 72 列全套指标(那会让全量峰值达 6.8GB)。
|
||||
2. 分批: 范围超过 batch_days 个交易日时按批切片, 每批带 warmup 前缀算完后 concat。
|
||||
2. 分批: 每批带 warmup 前缀算完后立即聚合为日级行, 不保留跨批个股明细。
|
||||
batch_days / warmup_days 由用户偏好控制(数据页「市场环境」卡片设置),
|
||||
实测默认值(60/40)全量(515万行)峰值约 1.9GB, 4GB 内存机器可稳跑。
|
||||
峰值随单批大小受控, 不再随完整历史长度线性增长。
|
||||
必须传入 instruments(涨跌停价表), 否则 compute_limit_signals 会跳过涨跌停信号。
|
||||
"""
|
||||
try:
|
||||
@@ -381,6 +394,7 @@ def _scan_enriched_fallback(repo, start: date, end: date) -> pl.DataFrame | None
|
||||
enriched_dir = repo.store.data_dir / "kline_daily_enriched"
|
||||
if not enriched_dir.exists():
|
||||
return None
|
||||
excluded_symbols = excluded_symbols or []
|
||||
instruments = repo.get_instruments()
|
||||
historical_shares = repo.get_historical_shares()
|
||||
|
||||
@@ -394,23 +408,31 @@ def _scan_enriched_fallback(repo, start: date, end: date) -> pl.DataFrame | None
|
||||
if len(target_dates) <= batch_days:
|
||||
df = _compute_batch(repo, enriched_dir, instruments, historical_shares,
|
||||
target_dates[0], target_dates[-1], warmup_days)
|
||||
return df if not df.is_empty() else None
|
||||
if df.is_empty():
|
||||
return None
|
||||
df = _filter_excluded_symbols(df, excluded_symbols)
|
||||
result = _aggregate_daily(df, index_pct_map)
|
||||
return result if not result.is_empty() else None
|
||||
|
||||
# 大范围: 按交易日分批, 逐批算 + concat
|
||||
# 大范围: 每批个股明细立即压缩为日级行, 只保留小型聚合结果。
|
||||
batches = [
|
||||
(target_dates[i], target_dates[min(i + batch_days - 1, len(target_dates) - 1)])
|
||||
for i in range(0, len(target_dates), batch_days)
|
||||
]
|
||||
logger.info("regime fallback: %d 天分 %d 批 (每批≤%d天 + %d天warmup)",
|
||||
logger.info("regime fallback: %d 天分 %d 批逐批聚合 (每批≤%d天 + %d天warmup)",
|
||||
len(target_dates), len(batches), batch_days, warmup_days)
|
||||
parts: list[pl.DataFrame] = []
|
||||
daily_parts: list[pl.DataFrame] = []
|
||||
for bs, be in batches:
|
||||
df = _compute_batch(repo, enriched_dir, instruments, historical_shares, bs, be, warmup_days)
|
||||
if not df.is_empty():
|
||||
parts.append(df)
|
||||
if not parts:
|
||||
if df.is_empty():
|
||||
continue
|
||||
df = _filter_excluded_symbols(df, excluded_symbols)
|
||||
daily = _aggregate_daily(df, index_pct_map)
|
||||
if not daily.is_empty():
|
||||
daily_parts.append(daily)
|
||||
if not daily_parts:
|
||||
return None
|
||||
return pl.concat(parts, how="vertical_relaxed")
|
||||
return pl.concat(daily_parts, how="vertical_relaxed")
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("regime scan_enriched_fallback failed: %s", e)
|
||||
return None
|
||||
@@ -440,15 +462,6 @@ def run_regime_batch(repo, start: date, end: date) -> pl.DataFrame:
|
||||
# 指数涨幅(主力指数)
|
||||
index_pct_map = _load_index_pct(repo, start, end)
|
||||
|
||||
# enriched 多日数据(优先缓存)
|
||||
df = repo.get_enriched_range(start, end)
|
||||
if df is None or df.is_empty():
|
||||
logger.info("regime batch: enriched cache miss [%s~%s], fallback to scan", start, end)
|
||||
df = _scan_enriched_fallback(repo, start, end)
|
||||
if df is None or df.is_empty():
|
||||
logger.info("regime batch: no enriched data for [%s~%s]", start, end)
|
||||
return pl.DataFrame()
|
||||
|
||||
# 口径: 默认剔除风险警示(ST)股(与主线统计同一开关) — 主板 ST 在 2026-07 前
|
||||
# 享 5% 涨跌幅且是跨行业状态桶, 混入会系统性抬高涨停宽度/高度(弱市炒 ST 尤甚)。
|
||||
# 涨跌家数/MA20 占比等宽度指标几乎不受影响。切换口径需全量重算 regime。
|
||||
@@ -457,14 +470,33 @@ def run_regime_batch(repo, start: date, end: date) -> pl.DataFrame:
|
||||
exclude_st = _prefs_st.get_sentiment_exclude_st()
|
||||
except Exception:
|
||||
exclude_st = True
|
||||
excluded_symbols: list[str] = []
|
||||
if exclude_st:
|
||||
from app.services.market_mainline import load_risk_warning_symbols
|
||||
|
||||
st_syms = load_risk_warning_symbols(repo.store.data_dir)
|
||||
if st_syms and "symbol" in df.columns:
|
||||
df = df.filter(
|
||||
~pl.col("symbol").str.to_uppercase().is_in(sorted(st_syms))
|
||||
)
|
||||
excluded_symbols = sorted(st_syms)
|
||||
|
||||
# enriched 多日数据(优先缓存)。慢路径在每批内部完成过滤和日级聚合,
|
||||
# 避免把所有批次的个股明细同时保留到最终 group_by。
|
||||
df = repo.get_enriched_range(start, end)
|
||||
if df is None or df.is_empty():
|
||||
logger.info("regime batch: enriched cache miss [%s~%s], fallback to scan", start, end)
|
||||
result = _scan_enriched_fallback(
|
||||
repo,
|
||||
start,
|
||||
end,
|
||||
index_pct_map=index_pct_map,
|
||||
excluded_symbols=excluded_symbols,
|
||||
)
|
||||
if result is None or result.is_empty():
|
||||
logger.info("regime batch: no enriched data for [%s~%s]", start, end)
|
||||
return pl.DataFrame()
|
||||
return result
|
||||
|
||||
df = _filter_excluded_symbols(df, excluded_symbols)
|
||||
if df.is_empty():
|
||||
return pl.DataFrame()
|
||||
|
||||
return _aggregate_daily(df, index_pct_map)
|
||||
|
||||
|
||||
@@ -35,12 +35,16 @@ logger = logging.getLogger(__name__)
|
||||
_CACHE_TTL = 120.0
|
||||
_cache: dict[str, dict] = {}
|
||||
_cache_ts: dict[str, float] = {}
|
||||
# 该条目实际覆盖的天数: enriched 只读 days 换算出的日历窗口, 缓存的"全量"因此
|
||||
# 以写入时的 days 为上限, 请求更长窗口时不能复用 (见 build_rps_rotation)。
|
||||
_cache_days: dict[str, int] = {}
|
||||
|
||||
|
||||
def invalidate_cache() -> None:
|
||||
"""清空轮动矩阵结果缓存(数据管道完成后调用, 避免返回旧数据)。"""
|
||||
_cache.clear()
|
||||
_cache_ts.clear()
|
||||
_cache_days.clear()
|
||||
|
||||
|
||||
def _latest_enriched_date(repo) -> date | None:
|
||||
@@ -95,13 +99,16 @@ def _load_concept_map_df(repo, kind: str = "concept") -> tuple[pl.DataFrame, int
|
||||
).unique()
|
||||
else:
|
||||
map_df = pl.DataFrame(schema={"_sym_up": pl.Utf8, kind: pl.Utf8})
|
||||
_map_cache[kind] = map_df
|
||||
# 缓存与返回值同构 ((map_df, count) 元组): 旧版只缓存裸 map_df, 命中路径
|
||||
# 返回 DataFrame 被调用方当元组解包, 600s 内二次访问必报错 (#186)
|
||||
payload = (map_df, len(members_seen))
|
||||
_map_cache[kind] = payload
|
||||
_map_ts[kind] = now
|
||||
return map_df, len(members_seen)
|
||||
return payload
|
||||
|
||||
|
||||
# 维度映射缓存: {kind: (map_df, count)}。按 kind 隔离(概念/行业分别缓存)。
|
||||
_map_cache: dict[str, pl.DataFrame] = {}
|
||||
_map_cache: dict[str, tuple[pl.DataFrame, int]] = {}
|
||||
_map_ts: dict[str, float] = {}
|
||||
|
||||
|
||||
@@ -135,18 +142,15 @@ def build_rps_rotation(repo, days: int = 12, kind: str = "concept", level: int |
|
||||
cache_key = f"{kind}|{level}|{latest.isoformat()}"
|
||||
now = time.time()
|
||||
cached = _cache.get(cache_key)
|
||||
if cached and (now - _cache_ts.get(cache_key, 0)) < _CACHE_TTL:
|
||||
if (
|
||||
cached
|
||||
and _cache_days.get(cache_key, 0) >= days
|
||||
and (now - _cache_ts.get(cache_key, 0)) < _CACHE_TTL
|
||||
):
|
||||
return _slice_cached(cached, days)
|
||||
|
||||
# 1. 维度映射(symbol → 维度成员), 已按 kind 缓存为 polars DataFrame。
|
||||
# 兼容返回裸 DataFrame 的实现: 元组解包会把两列拆成 Series(见
|
||||
# market_mainline.compute_mainline_range 同类处理)。
|
||||
loaded = _load_concept_map_df(repo, kind)
|
||||
if isinstance(loaded, tuple):
|
||||
map_df, member_count = loaded
|
||||
else:
|
||||
map_df = loaded
|
||||
member_count = loaded[kind].n_unique() if kind in loaded.columns else 0
|
||||
# 1. 维度映射(symbol → 维度成员), 已按 kind 缓存为 (map_df, count) 元组 (#186)。
|
||||
map_df, member_count = _load_concept_map_df(repo, kind)
|
||||
if map_df.is_empty():
|
||||
logger.info("rps_rotation: no %s data (ext dimension not fetched yet)", kind)
|
||||
return {"dates": [], "columns": {}, "concept_count": 0}
|
||||
@@ -205,9 +209,10 @@ def build_rps_rotation(repo, days: int = 12, kind: str = "concept", level: int |
|
||||
"concept_count": member_count,
|
||||
}
|
||||
|
||||
# 写缓存(存全量, 按需 slice)
|
||||
# 写缓存(存本次窗口的全量, 按需 slice; 覆盖天数一并记下)
|
||||
_cache[cache_key] = full
|
||||
_cache_ts[cache_key] = now
|
||||
_cache_days[cache_key] = days
|
||||
|
||||
return _slice_cached(full, days)
|
||||
|
||||
|
||||
@@ -23,6 +23,9 @@ logger = logging.getLogger(__name__)
|
||||
_history_cache: dict[tuple[str, date, int], tuple[float, pl.DataFrame]] = {}
|
||||
_HISTORY_CACHE_TTL = 120.0 # 秒
|
||||
|
||||
# load_prior_consecutive 最多回看多少个已存在的日分区 (缺列时继续往前找的上限)
|
||||
_PRIOR_PARTITION_SCAN = 10
|
||||
|
||||
|
||||
@dataclass
|
||||
class ScreenerResult:
|
||||
@@ -109,15 +112,14 @@ class ScreenerService:
|
||||
可直接从 parquet 读取, 无需 _load_enriched_for_date 的全量指标重算
|
||||
(历史日期该慢路径最坏会触发 9 次全市场 compute_enriched_full)。
|
||||
|
||||
选取逻辑与旧循环等价: 在 as_of 前 1~9 天内找到第一个存在的日分区
|
||||
(即前一交易日), 读取其 symbol + consec_col。存储列的值与重算值逐位一致
|
||||
(连板计数为 run-length, 150 天 warmup 完全覆盖 A 股最长连板, 二者相等)。
|
||||
由近到远取 as_of 之前已存在的日分区 (即前一交易日), 读取其
|
||||
symbol + consec_col。存储列的值与重算值逐位一致 (连板计数为 run-length,
|
||||
150 天 warmup 完全覆盖 A 股最长连板, 二者相等)。
|
||||
|
||||
返回列: symbol, prev_consec。找不到前一交易日时返回空 DataFrame。
|
||||
"""
|
||||
enriched_dir = self.repo.store.data_dir / self._enriched_dirname
|
||||
for delta in range(1, 10):
|
||||
candidate = as_of - timedelta(days=delta)
|
||||
for candidate in self._prior_partition_dates(as_of, _PRIOR_PARTITION_SCAN):
|
||||
target_parquet = enriched_dir / f"date={candidate.isoformat()}" / "part.parquet"
|
||||
if not target_parquet.exists():
|
||||
continue
|
||||
@@ -140,6 +142,32 @@ class ScreenerService:
|
||||
return pl.DataFrame()
|
||||
return pl.DataFrame()
|
||||
|
||||
def _prior_partition_dates(self, as_of: date, limit: int) -> list[date]:
|
||||
"""enriched 目录里早于 as_of 的分区日期, 由近到远最多 limit 个。
|
||||
|
||||
枚举分区目录而不是按自然日回看固定天数: 春节长假连着调休周末,
|
||||
相邻两个交易日能隔 10~11 个自然日 (如 2024-02-08 → 2024-02-19),
|
||||
固定窗口会整段落空。与 auction_benchmark._prev_trading_day
|
||||
「本地日K分区日期 = 已知交易日集合」同口径。
|
||||
"""
|
||||
enriched_dir = self.repo.store.data_dir / self._enriched_dirname
|
||||
days: list[date] = []
|
||||
try:
|
||||
entries = list(enriched_dir.iterdir())
|
||||
except OSError:
|
||||
return []
|
||||
for part in entries:
|
||||
if not part.name.startswith("date="):
|
||||
continue
|
||||
try:
|
||||
day = date.fromisoformat(part.name[5:])
|
||||
except ValueError:
|
||||
continue
|
||||
if day < as_of:
|
||||
days.append(day)
|
||||
days.sort(reverse=True)
|
||||
return days[:limit]
|
||||
|
||||
def _compute_enriched_full(self, df_target: pl.DataFrame, target_date: date) -> pl.DataFrame:
|
||||
"""从 14 列基础数据即时计算完整 enriched (含全部指标和信号)。
|
||||
|
||||
@@ -154,8 +182,10 @@ class ScreenerService:
|
||||
# 加载 warmup 历史 (目标日期前 ~120 天)
|
||||
enriched_dir = self.repo.store.data_dir / self._enriched_dirname
|
||||
start = target_date - timedelta(days=150)
|
||||
# turnover_rate 是 enriched 存储列, 必须随行透传: 否则即时计算后该列
|
||||
# 丢失, 自定义 SQL 用它做条件会 Binder Error 被吞成空结果 (#187)
|
||||
read_cols = ["symbol", "date", "open", "high", "low", "close", "volume",
|
||||
"amount", "raw_close", "raw_high", "raw_low"]
|
||||
"amount", "raw_close", "raw_high", "raw_low", "turnover_rate"]
|
||||
|
||||
try:
|
||||
lf = (
|
||||
@@ -245,8 +275,9 @@ class ScreenerService:
|
||||
start = target_date - timedelta(days=min((lookback_days + warmup) * 2, 180))
|
||||
|
||||
enriched_dir = self.repo.store.data_dir / self._enriched_dirname
|
||||
# 同 _compute_enriched_full: turnover_rate 存储列随行透传 (#187)
|
||||
read_cols = ["symbol", "date", "open", "high", "low", "close", "volume",
|
||||
"amount", "raw_close", "raw_high", "raw_low"]
|
||||
"amount", "raw_close", "raw_high", "raw_low", "turnover_rate"]
|
||||
|
||||
try:
|
||||
lf = (
|
||||
@@ -335,10 +366,14 @@ class ScreenerService:
|
||||
# 用独立的 :memory: 连接 (而非复用 repo 共享连接的 cursor): conditions 是用户
|
||||
# 传入的 SQL 片段, 隔离连接下注入至多能碰 read_csv/read_parquet 文件; 若复用共享
|
||||
# 连接则会把 app 已注册的真实业务表也暴露给注入, 扩大攻击面。隔离连接创建开销极低。
|
||||
# 再关闭 external_access, 让注入的文件读写函数 (read_parquet/COPY 等) 直接报错,
|
||||
# 视图数据仍通过 con.register 注入, 不受该开关影响 (#224)。
|
||||
con = None
|
||||
try:
|
||||
import duckdb
|
||||
con = duckdb.connect(database=":memory:")
|
||||
con = duckdb.connect(
|
||||
database=":memory:", config={"enable_external_access": False}
|
||||
)
|
||||
con.register("enriched", df.to_arrow())
|
||||
where = " AND ".join(f"({c})" for c in conditions)
|
||||
sql = f"SELECT * FROM enriched WHERE {where}"
|
||||
|
||||
@@ -16,8 +16,8 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import AsyncIterator
|
||||
from pathlib import Path
|
||||
from typing import AsyncIterator
|
||||
|
||||
import polars as pl
|
||||
|
||||
@@ -238,10 +238,10 @@ def _build_user_prompt(
|
||||
"请按系统提示词第 4 节的说明,在基本面/财务面维度给出\"接入中\"的友好提示,不要编造数据。)",
|
||||
])
|
||||
|
||||
from app.services.ai_provider import sanitize_focus
|
||||
safe_focus = sanitize_focus(focus)
|
||||
if safe_focus:
|
||||
parts.extend(["", f"本次分析请特别关注: {safe_focus}"])
|
||||
from app.services.ai_provider import build_focus_instruction
|
||||
focus_instruction = build_focus_instruction(focus, report_name="个股分析报告")
|
||||
if focus_instruction:
|
||||
parts.extend(["", focus_instruction])
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
@@ -325,11 +325,12 @@ async def analyze_stock_stream(
|
||||
# 不限制输出: 推理模型(deepseek reasoner 系)思考 token 计入 max_tokens
|
||||
# 预算, 固定上限会把正文挤光(实测 4500 全被推理吃掉 → 正文 0 字)。
|
||||
max_tokens=None,
|
||||
prefer_final_answer=True,
|
||||
):
|
||||
got_content = True
|
||||
yield json.dumps({"type": "delta", "content": delta}, ensure_ascii=False)
|
||||
|
||||
except Exception as e: # noqa: BLE001
|
||||
except Exception as e:
|
||||
logger.exception("AI stock analysis failed for %s: %s", symbol, e)
|
||||
yield json.dumps({"type": "error", "message": f"AI 分析失败: {e}"}, ensure_ascii=False)
|
||||
return
|
||||
|
||||
@@ -75,6 +75,15 @@ def read_cache(data_dir: Path) -> dict | None:
|
||||
|
||||
def clear_cache(data_dir: Path) -> None:
|
||||
"""删除策略结果缓存;策略代码 reload 后避免继续展示旧公式结果。"""
|
||||
import traceback
|
||||
|
||||
# 运维可见性: 策略页依赖本缓存秒加载, 被清空即整页回退到全量重算。
|
||||
# 记录调用链 (最近 5 帧), 排查"缓存莫名消失"类问题不需要复现现场。
|
||||
frames = traceback.extract_stack()[:-1]
|
||||
chain = " <- ".join(
|
||||
f"{f.filename.rsplit('/', 1)[-1]}:{f.lineno}:{f.name}" for f in frames[-5:]
|
||||
)
|
||||
logger.warning("策略缓存被清除, 调用链: %s", chain)
|
||||
path = _cache_path(data_dir)
|
||||
with _file_lock:
|
||||
path.unlink(missing_ok=True)
|
||||
|
||||
@@ -0,0 +1,182 @@
|
||||
"""策略 run_all 渐进式执行 — 单飞后台执行 + 快策略先返回。
|
||||
|
||||
页面进入策略页时 run_all 全量跑需要 ~2 分钟, 用户只能盯着空卡片等。此模块把
|
||||
执行拆成「同步等一小段 + 后台继续算」:
|
||||
|
||||
- 全局同一时刻只执行一个 run_all (polars/Numba 并发跑两份有崩死风险),
|
||||
请求先到先得, 后来者排队; 相同 key (资产/周期/日期/策略集) 的重复请求
|
||||
直接搭车现有执行, 不重复算。
|
||||
- 按历史耗时升序执行: 快策略 (秒级) 在首返时限内完成并随 HTTP 响应返回,
|
||||
慢策略 (分钟级) 留在后台慢慢算。
|
||||
- 每个策略算完立刻增量写入 strategy_cache, 前端轮询 cached-summary
|
||||
逐个点亮卡片数字。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_TIMINGS_FILENAME = "strategy_run_timings.json"
|
||||
_timings_lock = threading.Lock()
|
||||
|
||||
|
||||
def _timings_path(data_dir: Path) -> Path:
|
||||
return data_dir / "user_data" / _TIMINGS_FILENAME
|
||||
|
||||
|
||||
def load_run_timings(data_dir: Path) -> dict[str, float]:
|
||||
"""读取各策略上次执行耗时 (ms); 无文件/损坏时返回空。"""
|
||||
with _timings_lock:
|
||||
try:
|
||||
data = json.loads(_timings_path(data_dir).read_text(encoding="utf-8"))
|
||||
except (FileNotFoundError, ValueError, OSError):
|
||||
return {}
|
||||
if not isinstance(data, dict):
|
||||
return {}
|
||||
return {str(k): float(v) for k, v in data.items() if isinstance(v, (int, float))}
|
||||
|
||||
|
||||
def record_run_timings(data_dir: Path, elapsed_ms: dict[str, float]) -> None:
|
||||
"""批量记录策略耗时 (ms), 与已有文件合并后原子重写。"""
|
||||
if not elapsed_ms:
|
||||
return
|
||||
with _timings_lock:
|
||||
path = _timings_path(data_dir)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
merged: dict[str, float] = {}
|
||||
try:
|
||||
old = json.loads(path.read_text(encoding="utf-8"))
|
||||
if isinstance(old, dict):
|
||||
merged = {str(k): float(v) for k, v in old.items() if isinstance(v, (int, float))}
|
||||
except (FileNotFoundError, ValueError, OSError):
|
||||
pass
|
||||
merged.update({sid: float(ms) for sid, ms in elapsed_ms.items()})
|
||||
tmp = path.with_name(path.name + ".tmp")
|
||||
tmp.write_text(json.dumps(merged, ensure_ascii=False), encoding="utf-8")
|
||||
os.replace(tmp, path)
|
||||
|
||||
|
||||
def order_strategy_ids(all_ids: list[str], timings: dict[str, float]) -> list[str]:
|
||||
"""快策略先算: 有历史耗时的按耗时升序, 未知耗时的保持原顺序排在后面。"""
|
||||
known = sorted(
|
||||
(timings[sid], i, sid) for i, sid in enumerate(all_ids) if sid in timings
|
||||
)
|
||||
known_ids = {sid for _, _, sid in known}
|
||||
unknown = [sid for sid in all_ids if sid not in known_ids]
|
||||
return [sid for _, _, sid in known] + unknown
|
||||
|
||||
|
||||
class StrategyRunHandle:
|
||||
"""一次 run_all 的执行状态; 端点线程 (读) 与后台执行线程 (写) 共享。"""
|
||||
|
||||
def __init__(self, key: tuple, ordered_ids: list[str]) -> None:
|
||||
self.key = key
|
||||
self.started_at_ms = int(time.time() * 1000)
|
||||
self._lock = threading.Lock()
|
||||
self._results: dict[str, dict] = {}
|
||||
self._remaining: list[str] = list(ordered_ids)
|
||||
self._errors: dict[str, str] = {}
|
||||
self._error: str | None = None
|
||||
self._done = False
|
||||
|
||||
def complete(self, sid: str, payload: dict) -> None:
|
||||
with self._lock:
|
||||
self._results[sid] = payload
|
||||
if sid in self._remaining:
|
||||
self._remaining.remove(sid)
|
||||
|
||||
def fail_one(self, sid: str, message: str) -> None:
|
||||
"""单个策略失败: 记错误并移出待算队列, 不影响其余策略继续。"""
|
||||
with self._lock:
|
||||
self._errors[sid] = message
|
||||
if sid in self._remaining:
|
||||
self._remaining.remove(sid)
|
||||
|
||||
def fail(self, message: str) -> None:
|
||||
with self._lock:
|
||||
self._error = message
|
||||
self._done = True
|
||||
|
||||
def finish(self) -> None:
|
||||
with self._lock:
|
||||
self._done = True
|
||||
|
||||
def snapshot(self) -> dict:
|
||||
"""线程安全快照: 结果拷贝 + 剩余/逐策略错误/整体错误/完成状态。"""
|
||||
with self._lock:
|
||||
return {
|
||||
"results": dict(self._results),
|
||||
"pending": list(self._remaining),
|
||||
"errors": dict(self._errors),
|
||||
"error": self._error,
|
||||
"done": self._done,
|
||||
"started_at_ms": self.started_at_ms,
|
||||
}
|
||||
|
||||
|
||||
class StrategyRunManager:
|
||||
"""run_all 单飞管理器。
|
||||
|
||||
- 相同 key 且仍在执行 (含排队中) 的重复请求搭车现有执行, 不重复算
|
||||
(页面 reload / StrictMode / 反复切换); 已完成的不再搭车, 重跑即新执行。
|
||||
- 不同 key 在唯一 daemon 工作线程里排队; 端点在首返时限内等不到也只能
|
||||
先返回 pending, 前端靠轮询缓存拿最终结果。
|
||||
- 工作线程为 daemon: 进程退出不等待剩余计算 (缓存写入均为原子替换,
|
||||
中断只留部分结果, 下次进入页面补算)。
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._lock = threading.Lock()
|
||||
self._handles: dict[tuple, StrategyRunHandle] = {}
|
||||
self._queue: queue.Queue[tuple[StrategyRunHandle, Callable]] = queue.Queue()
|
||||
self._worker: threading.Thread | None = None
|
||||
|
||||
def get_or_submit(
|
||||
self,
|
||||
key: tuple,
|
||||
ordered_ids: list[str],
|
||||
job: Callable[[StrategyRunHandle], None],
|
||||
) -> StrategyRunHandle:
|
||||
with self._lock:
|
||||
# 顺手清理已完成的 handle, 防止字典随不同 key 无限增长
|
||||
for k in [k for k, h in self._handles.items() if h.snapshot()["done"]]:
|
||||
del self._handles[k]
|
||||
existing = self._handles.get(key)
|
||||
if existing is not None:
|
||||
return existing
|
||||
handle = StrategyRunHandle(key, ordered_ids)
|
||||
self._handles[key] = handle
|
||||
self._ensure_worker()
|
||||
self._queue.put((handle, job))
|
||||
return handle
|
||||
|
||||
def _ensure_worker(self) -> None:
|
||||
with self._lock:
|
||||
if self._worker is None or not self._worker.is_alive():
|
||||
self._worker = threading.Thread(
|
||||
target=self._run_loop, name="runall", daemon=True
|
||||
)
|
||||
self._worker.start()
|
||||
|
||||
def _run_loop(self) -> None:
|
||||
while True:
|
||||
handle, job = self._queue.get()
|
||||
try:
|
||||
job(handle)
|
||||
except Exception as e:
|
||||
logger.exception("run_all 后台执行失败: %s", e)
|
||||
handle.fail(str(e))
|
||||
else:
|
||||
handle.finish()
|
||||
|
||||
|
||||
# 进程级单例: 与 strategy_cache 的模块级锁同风格, 生命周期跟随进程
|
||||
MANAGER = StrategyRunManager()
|
||||
@@ -113,9 +113,11 @@ def is_trading_day(now: datetime | None = None) -> bool | None:
|
||||
return False
|
||||
|
||||
with _CACHE_LOCK:
|
||||
# 「未知」(None) 也是一个结论, 同样按 TTL 缓存 —— 它正是 _TTL_UNKNOWN_S 要
|
||||
# 挡住的场景 (未配 fuyao 且 tickflow 不可用时, 轮询每拍都会重打一次探测)。
|
||||
# _CACHE.day 只在探测写回时设置, 因此「当天已探过」用它判定即可。
|
||||
if (
|
||||
_CACHE.day == now.date()
|
||||
and _CACHE.verdict is not None
|
||||
and (time.monotonic() - _CACHE.probed_at) < _ttl_of(_CACHE.verdict)
|
||||
):
|
||||
return _CACHE.verdict
|
||||
|
||||
@@ -183,15 +183,23 @@ def add_batch(
|
||||
symbols: list[str],
|
||||
note: str = "",
|
||||
group_id: str | None = None,
|
||||
group_ids: list[str] | None = None,
|
||||
) -> tuple[list[dict], int]:
|
||||
"""批量添加并保持既有语义:每个新处理的标的移动到列表最前面。
|
||||
|
||||
group_id 为可选的初始分组(如从某分组页添加时); 重复添加的标的保留
|
||||
既有全部分组, 仅在显式传入 group_id 且尚未属于该组时并入。
|
||||
分组为可选的初始分组:``group_id`` 单组(如从某分组页添加)或 ``group_ids``
|
||||
多组(如批量导入同时并入多个分组)。重复添加的标的保留既有全部分组,
|
||||
仅把尚未属于的传入分组并入;二者可同时使用、内部去重。
|
||||
"""
|
||||
with _LOCK:
|
||||
groups = _read_groups()
|
||||
_validate_group_id(group_id, groups)
|
||||
# 合并单/多组参数并去重;逐组校验存在性
|
||||
apply_ids: list[str] = []
|
||||
for gid in (group_ids or []) + ([group_id] if group_id is not None else []):
|
||||
if gid in apply_ids:
|
||||
continue
|
||||
_validate_group_id(gid, groups)
|
||||
apply_ids.append(gid)
|
||||
rows = _read_entries().to_dicts()
|
||||
added = 0
|
||||
for symbol in symbols:
|
||||
@@ -200,8 +208,9 @@ def add_batch(
|
||||
added += 1
|
||||
rows = [row for row in rows if row["symbol"] != symbol]
|
||||
gids = list((existing or {}).get("group_ids") or [])
|
||||
if group_id is not None and group_id not in gids:
|
||||
gids.append(group_id)
|
||||
for gid in apply_ids:
|
||||
if gid not in gids:
|
||||
gids.append(gid)
|
||||
rows.insert(0, {
|
||||
"symbol": symbol,
|
||||
"added_at": datetime.utcnow().isoformat(timespec="seconds"),
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
"""自选股 CSV/TXT 与粘贴代码批量导入:解码 → 抽代码 → instruments 校验。
|
||||
|
||||
国内行情软件(同花顺/东财/通达信)导出的自选多为 CSV/TXT,且常为 GBK 系编码
|
||||
(参见 ext_data.ensure_utf8_csv 的说明)。本模块把上传字节 / 粘贴文本解析为与截图
|
||||
OCR 一致的候选结构,前端复用同一套勾选确认流程;写入目标统一走自选分组语义
|
||||
(watchlist.add_batch 的 group_ids M:N 并入),本模块不落盘、不建标签。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
import io
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from app.services.watchlist_ocr.pipeline import (
|
||||
_CODE_RE,
|
||||
ImportCandidate,
|
||||
build_instrument_lookups,
|
||||
extract_codes,
|
||||
resolve_candidates,
|
||||
)
|
||||
|
||||
_CJK_RE = re.compile(f"[{chr(0x4E00)}-{chr(0x9FFF)}]") # CJK 统一表意文字块
|
||||
|
||||
# 编码回退链:UTF-8(含 BOM)→ GB18030(GB18030 是 GBK 超集,无需单独回退)
|
||||
_ENCODINGS = ("utf-8-sig", "gb18030")
|
||||
|
||||
|
||||
def _finalize(
|
||||
provider: str,
|
||||
text: str,
|
||||
codes: list[str],
|
||||
candidates: list[dict[str, Any]],
|
||||
) -> dict[str, Any]:
|
||||
"""组装与截图 OCR 一致的候选响应,统一 matched/unmatched 计数口径。"""
|
||||
matched_count = sum(1 for c in candidates if c["matched"])
|
||||
return {
|
||||
"provider": provider,
|
||||
"raw_text": text,
|
||||
"codes": codes,
|
||||
"candidates": candidates,
|
||||
"matched_count": matched_count,
|
||||
"unmatched_count": len(candidates) - matched_count,
|
||||
}
|
||||
|
||||
|
||||
def decode_csv_bytes(raw: bytes) -> str:
|
||||
"""把上传字节解码为文本,兼容 UTF-8 / GBK 系编码。"""
|
||||
if not raw:
|
||||
raise ValueError("空文件")
|
||||
last_err: Exception | None = None
|
||||
for enc in _ENCODINGS:
|
||||
try:
|
||||
return raw.decode(enc)
|
||||
except (UnicodeDecodeError, LookupError) as e:
|
||||
last_err = e
|
||||
raise ValueError("无法识别文件编码,请另存为 UTF-8 或 GBK 后重试") from last_err
|
||||
|
||||
|
||||
def _is_code_cell(cell: str) -> bool:
|
||||
# 调用方(parse_csv_rows)已 strip 过单元格
|
||||
return bool(_CODE_RE.fullmatch(cell))
|
||||
|
||||
|
||||
def _pick_name(cells: list[str]) -> str | None:
|
||||
"""取行内首个含 ≥2 个汉字且非六位代码的单元格作为名称候选。"""
|
||||
for cell in cells:
|
||||
if not cell or _is_code_cell(cell):
|
||||
continue
|
||||
if len(_CJK_RE.findall(cell)) >= 2:
|
||||
return cell
|
||||
return None
|
||||
|
||||
|
||||
def parse_csv_rows(text: str) -> list[tuple[list[str], str | None]]:
|
||||
"""解析 CSV/TXT 文本为 [(行内代码列表, 名称候选), ...]。
|
||||
|
||||
- 自动识别逗号 / Tab 分隔(同花顺/通达信导出常见 Tab)。
|
||||
- 逐行取所有六位数字作为代码候选;无代码行保留给名称兜底(是否输出由
|
||||
import_watchlist_csv 决定:表头等名称命不中主数据的行会被忽略)。
|
||||
- 返回列表保持文件行序。
|
||||
"""
|
||||
if not text.strip():
|
||||
return []
|
||||
|
||||
first = next((ln for ln in text.splitlines() if ln.strip()), "")
|
||||
delimiter = "\t" if first.count("\t") > first.count(",") else ","
|
||||
reader = csv.reader(io.StringIO(text), delimiter=delimiter)
|
||||
|
||||
rows: list[tuple[list[str], str | None]] = []
|
||||
for raw_row in reader:
|
||||
cells = [c.strip() for c in raw_row if c is not None]
|
||||
if not cells:
|
||||
continue
|
||||
codes: list[str] = []
|
||||
for cell in cells:
|
||||
codes.extend(m.group(1) for m in _CODE_RE.finditer(cell))
|
||||
rows.append((codes, _pick_name(cells)))
|
||||
return rows
|
||||
|
||||
|
||||
def import_watchlist_csv(
|
||||
raw: bytes,
|
||||
data_dir: Path,
|
||||
*,
|
||||
existing_symbols: set[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""解析 CSV/TXT 字节并返回候选列表(不写入自选)。返回结构与 OCR 一致。"""
|
||||
text = decode_csv_bytes(raw)
|
||||
return _resolve_rows(text, data_dir, existing_symbols=existing_symbols)
|
||||
|
||||
|
||||
def import_watchlist_codes(
|
||||
text: str,
|
||||
data_dir: Path,
|
||||
*,
|
||||
existing_symbols: set[str] | None = None,
|
||||
max_codes: int = 1000,
|
||||
) -> dict[str, Any]:
|
||||
"""解析粘贴的证券代码并返回候选列表(不写入自选)。
|
||||
|
||||
与 CSV 行级解析不同:粘贴文本里的多个代码可能挤在同一行/同一段(逗号、空格、
|
||||
换行分隔),必须按 ``extract_codes`` 全量抽码、去重保序,逐码生成候选,
|
||||
否则会把同行多码压成单候选而静默丢码。仅与 CSV 路径共享 lookups/resolve。
|
||||
"""
|
||||
codes = extract_codes(text)
|
||||
if not codes:
|
||||
return _finalize("codes", text, [], [])
|
||||
if len(codes) > max_codes:
|
||||
raise ValueError(f"一次最多导入 {max_codes} 个股票代码,已识别 {len(codes)} 个")
|
||||
|
||||
code_to_symbol, symbol_to_name = build_instrument_lookups(data_dir)
|
||||
candidates = resolve_candidates(codes, code_to_symbol, symbol_to_name, existing_symbols)
|
||||
return _finalize("codes", text, codes, [c.to_dict() for c in candidates])
|
||||
|
||||
|
||||
def _resolve_rows(
|
||||
text: str,
|
||||
data_dir: Path,
|
||||
*,
|
||||
existing_symbols: set[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""逐行把文本解析为候选(CSV/TXT 用;行内多码取首个已匹配者)。"""
|
||||
rows = parse_csv_rows(text)
|
||||
|
||||
code_to_symbol, symbol_to_name = build_instrument_lookups(data_dir)
|
||||
# 名称兜底反向表:CSV 可能只有名称列(无代码)
|
||||
name_to_symbol: dict[str, str] = {}
|
||||
for symbol, name in symbol_to_name.items():
|
||||
name_to_symbol.setdefault(name, symbol)
|
||||
|
||||
existing = existing_symbols or set()
|
||||
# 全部唯一代码按出现顺序一次性构建候选(复用 OCR 的构造逻辑,单一来源)
|
||||
unique_codes: list[str] = []
|
||||
seen_codes: set[str] = set()
|
||||
for row_codes, _ in rows:
|
||||
for c in row_codes:
|
||||
if c not in seen_codes:
|
||||
seen_codes.add(c)
|
||||
unique_codes.append(c)
|
||||
cand_by_code = {
|
||||
c.code: c
|
||||
for c in resolve_candidates(unique_codes, code_to_symbol, symbol_to_name, existing)
|
||||
}
|
||||
|
||||
candidates: list[dict[str, Any]] = []
|
||||
seen_symbols: set[str] = set() # 已发出的已匹配 symbol
|
||||
emitted_unmatched: set[str] = set() # 已发出的未匹配 code
|
||||
for row_codes, row_name in rows:
|
||||
# 行内代码优先:取第一个已匹配主数据的(避免价格/成交量数字误报)
|
||||
matched = next((cand_by_code[c] for c in row_codes if cand_by_code[c].matched), None)
|
||||
symbol = matched.symbol if matched else (name_to_symbol.get(row_name) if row_name else None)
|
||||
if symbol:
|
||||
if symbol in seen_symbols:
|
||||
continue
|
||||
seen_symbols.add(symbol)
|
||||
cand = matched or ImportCandidate(
|
||||
code=row_codes[0] if row_codes else "",
|
||||
symbol=symbol,
|
||||
name=symbol_to_name.get(symbol),
|
||||
matched=True,
|
||||
already_in_watchlist=symbol in existing,
|
||||
)
|
||||
else:
|
||||
# 名称兜底失败且无代码 → 表头/杂项行,忽略
|
||||
if not row_codes:
|
||||
continue
|
||||
code = row_codes[0]
|
||||
if code in emitted_unmatched:
|
||||
continue
|
||||
emitted_unmatched.add(code)
|
||||
cand = ImportCandidate(
|
||||
code=code,
|
||||
symbol=None,
|
||||
name=row_name,
|
||||
matched=False,
|
||||
already_in_watchlist=False,
|
||||
)
|
||||
candidates.append(cand.to_dict())
|
||||
|
||||
return _finalize("csv", text, unique_codes, candidates)
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Webhook 推送适配器 — 把告警事件推送到外部 IM / 量化软件。
|
||||
|
||||
职责: 把后端产生的告警事件, 通过用户配置的 Webhook 地址推送到外部。
|
||||
目前支持飞书群推送 Webhook; QMT / ptrade 等量化通道为待定。
|
||||
目前支持飞书、企业微信和通用第三方 JSON Webhook。
|
||||
|
||||
飞书自定义机器人接入:
|
||||
1. 飞书群 → 群设置 → 群推送 Webhook → 添加「自定义机器人」
|
||||
@@ -17,8 +17,10 @@ from __future__ import annotations
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from urllib.parse import urlparse
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -93,7 +95,7 @@ def _truncate_to_bytes(text: str, max_bytes: int, suffix: str = "…") -> str:
|
||||
_FEISHU_MAX_ATTEMPTS = 3
|
||||
|
||||
|
||||
def _post_feishu(webhook_url: str, payload: dict, secret: str) -> bool:
|
||||
def _post_feishu(webhook_url: str, payload: dict, secret: str, max_attempts: int = _FEISHU_MAX_ATTEMPTS) -> bool:
|
||||
"""发送飞书 webhook 请求并判定成败 (供 text / card 共用)。
|
||||
|
||||
成功响应: HTTP 200 且业务 code=0 (或非 JSON/非 dict 的 200)。
|
||||
@@ -102,11 +104,14 @@ def _post_feishu(webhook_url: str, payload: dict, secret: str) -> bool:
|
||||
一次瞬时 5xx/timeout 若不重试, 该告警会被冷却窗口(默认 1h)压掉, 离屏用户彻底
|
||||
收不到推送。永久失败 (4xx / 业务 code≠0, 如签名错、URL 失效) 不重试。最终失败
|
||||
记 WARNING (而非之前的 debug), 保证「推送丢了」在日志里可见。
|
||||
|
||||
max_attempts: 尝试次数, 默认 3 (生产推送语义)。诊断用途(如手动测试配置)可传 1,
|
||||
避免失败时等满退避重试。
|
||||
"""
|
||||
import httpx
|
||||
|
||||
last_err = ""
|
||||
for attempt in range(1, _FEISHU_MAX_ATTEMPTS + 1):
|
||||
for attempt in range(1, max_attempts + 1):
|
||||
try:
|
||||
# 启用签名校验时, 请求体须带 timestamp + sign (每次重试都重算, 防时间戳过期)
|
||||
if secret:
|
||||
@@ -136,14 +141,14 @@ def _post_feishu(webhook_url: str, payload: dict, secret: str) -> bool:
|
||||
except Exception as e: # noqa: BLE001 — 网络/超时, 可重试
|
||||
last_err = str(e)
|
||||
|
||||
if attempt < _FEISHU_MAX_ATTEMPTS:
|
||||
if attempt < max_attempts:
|
||||
time.sleep(min(2 ** (attempt - 1), 3)) # 退避: 1s, 2s
|
||||
|
||||
logger.warning("飞书 Webhook 推送最终失败(已重试 %d 次): %s", _FEISHU_MAX_ATTEMPTS, last_err)
|
||||
logger.warning("飞书 Webhook 推送最终失败(已重试 %d 次): %s", max_attempts, last_err)
|
||||
return False
|
||||
|
||||
|
||||
def send_feishu(webhook_url: str, title: str, body: str, secret: str = "") -> bool:
|
||||
def send_feishu(webhook_url: str, title: str, body: str, secret: str = "", max_attempts: int = _FEISHU_MAX_ATTEMPTS) -> bool:
|
||||
"""推送一条文本消息到飞书群推送 Webhook。
|
||||
|
||||
Args:
|
||||
@@ -151,6 +156,7 @@ def send_feishu(webhook_url: str, title: str, body: str, secret: str = "") -> bo
|
||||
title: 消息标题 (与正文拼接为一条文本)
|
||||
body: 消息正文
|
||||
secret: 签名密钥 (机器人启用了「签名校验」时必填; 留空则不带签名)
|
||||
max_attempts: 尝试次数 (诊断用途可传 1, 默认保持生产重试语义)
|
||||
|
||||
Returns:
|
||||
True=成功送达, False=失败或 URL 非法。
|
||||
@@ -164,7 +170,7 @@ def send_feishu(webhook_url: str, title: str, body: str, secret: str = "") -> bo
|
||||
return False
|
||||
|
||||
payload: dict = {"msg_type": "text", "content": {"text": text}}
|
||||
return _post_feishu(webhook_url, payload, secret)
|
||||
return _post_feishu(webhook_url, payload, secret, max_attempts)
|
||||
|
||||
|
||||
def send_feishu_card(webhook_url: str, title: str, subtitle: str, body_md: str, secret: str = "") -> bool:
|
||||
@@ -340,3 +346,76 @@ def send_wecom_markdown(webhook_url: str, title: str, body_md: str) -> bool:
|
||||
payload: dict = {"msgtype": "markdown", "markdown": {"content": content}}
|
||||
return _post_wecom(webhook_url, payload)
|
||||
|
||||
|
||||
# ================================================================
|
||||
# 通用第三方 JSON Webhook
|
||||
# ================================================================
|
||||
|
||||
_CUSTOM_MAX_ATTEMPTS = 3
|
||||
|
||||
|
||||
def is_valid_custom_url(url: str) -> bool:
|
||||
"""Accept absolute HTTP(S) URLs, including LAN endpoints used by local deployments."""
|
||||
try:
|
||||
parsed = urlparse((url or "").strip())
|
||||
except ValueError:
|
||||
return False
|
||||
return parsed.scheme in {"http", "https"} and bool(parsed.netloc) and not parsed.username
|
||||
|
||||
|
||||
def send_custom(
|
||||
webhook_url: str,
|
||||
title: str,
|
||||
body: str,
|
||||
event_type: str,
|
||||
data: dict | None = None,
|
||||
secret: str = "",
|
||||
max_attempts: int = _CUSTOM_MAX_ATTEMPTS,
|
||||
) -> bool:
|
||||
"""POST a stable JSON envelope to a user-configured third-party system.
|
||||
|
||||
When ``secret`` is configured the raw request body is signed with HMAC-SHA256.
|
||||
The receiver can validate ``X-TickFlow-Timestamp`` and
|
||||
``X-TickFlow-Signature: sha256=<hex>`` before accepting the event.
|
||||
"""
|
||||
if not is_valid_custom_url(webhook_url):
|
||||
return False
|
||||
|
||||
timestamp = str(int(time.time()))
|
||||
payload = {
|
||||
"event": str(event_type or "notification"),
|
||||
"timestamp": int(timestamp),
|
||||
"title": str(title or ""),
|
||||
"body": str(body or ""),
|
||||
"data": data or {},
|
||||
}
|
||||
encoded = json.dumps(
|
||||
payload, ensure_ascii=False, separators=(",", ":"), sort_keys=True, default=str,
|
||||
).encode("utf-8")
|
||||
headers = {"Content-Type": "application/json", "User-Agent": "TickFlow-Webhook/1.0"}
|
||||
if secret:
|
||||
digest = hmac.new(secret.encode("utf-8"), encoded, hashlib.sha256).hexdigest()
|
||||
headers["X-TickFlow-Timestamp"] = timestamp
|
||||
headers["X-TickFlow-Signature"] = f"sha256={digest}"
|
||||
|
||||
import httpx
|
||||
|
||||
last_err = ""
|
||||
for attempt in range(1, max_attempts + 1):
|
||||
try:
|
||||
response = httpx.post(
|
||||
webhook_url, content=encoded, headers=headers, timeout=5.0,
|
||||
)
|
||||
if 200 <= response.status_code < 300:
|
||||
return True
|
||||
last_err = f"HTTP {response.status_code}: {response.text[:200]}"
|
||||
if response.status_code < 500:
|
||||
logger.warning("第三方 Webhook 推送失败(不重试): %s", last_err)
|
||||
return False
|
||||
except Exception as exc: # Network failures are retryable and must not escape.
|
||||
last_err = str(exc)
|
||||
if attempt < max_attempts:
|
||||
time.sleep(min(2 ** (attempt - 1), 3))
|
||||
|
||||
logger.warning("第三方 Webhook 推送最终失败(已重试 %d 次): %s", max_attempts, last_err)
|
||||
return False
|
||||
|
||||
@@ -346,6 +346,7 @@ META = {{...}},{entrypoint_requirement}。只输出完整 Python 代码。
|
||||
"numpy",
|
||||
"app.backtest.matrix",
|
||||
"app.strategy.builtin.factor_rank_research",
|
||||
"app.strategy.market_data", # 新增: 策略可读取指数/ETF 日K
|
||||
"datetime",
|
||||
"__future__",
|
||||
})
|
||||
|
||||
@@ -80,7 +80,12 @@ def merge_results(
|
||||
ordered = sorted(symbols, key=lambda s: res.scores[s], reverse=True)
|
||||
count = len(ordered)
|
||||
for rank, sym in enumerate(ordered, start=1):
|
||||
norm[sym] = 1 - (rank - 1) / max(count - 1, 1)
|
||||
# 单候选无法排名, 必须用中性分: 当成"最优=1"会凭空抬高融合分,
|
||||
# 而回测合并 (merge_signal_matrices 的 n <= 1 分支) 用的是中性分,
|
||||
# 两条路径同一天同一标的会给出不同评分与排序。
|
||||
norm[sym] = (
|
||||
_NEUTRAL_NORM if count <= 1 else 1 - (rank - 1) / (count - 1)
|
||||
)
|
||||
else:
|
||||
# 子策略未产出 score: 命中即中性分, 不奖励也不惩罚。
|
||||
for row in res.rows:
|
||||
|
||||
@@ -28,6 +28,10 @@ logger = logging.getLogger(__name__)
|
||||
PREFIX = "csg_" # 自定义信号列名前缀
|
||||
ID_RE = re.compile(r"^[a-z0-9_]{1,40}$")
|
||||
OPS = {">", ">=", "<", "<=", "==", "!="}
|
||||
# string 扩展字段 (概念/行业归属等) 的运算符: contains 为字面量包含
|
||||
# (非正则, 用户输入不进入 pattern 编译), ==/!= 为字符串精确比较。
|
||||
STRING_OPS = {"contains", "==", "!="}
|
||||
_MAX_STR_RIGHT = 64
|
||||
|
||||
# 字段白名单:只允许这些列出现在条件里(防注入)。均为数值型。
|
||||
# 与 ENRICHED_COLUMNS 的数值列保持一致,排除 symbol/date/name 等非数值列。
|
||||
@@ -63,9 +67,66 @@ _OP_BUILDERS = {
|
||||
"<=": lambda c, v: c <= v,
|
||||
"==": lambda c, v: c == v,
|
||||
"!=": lambda c, v: c != v,
|
||||
# literal=True: 右值按字面量匹配, 不当正则编译 (用户输入含 .* 等也安全)
|
||||
"contains": lambda c, v: c.cast(pl.Utf8).str.contains(v, literal=True),
|
||||
}
|
||||
|
||||
|
||||
def _string_ext_fields() -> frozenset[str]:
|
||||
"""string 扩展字段列名 (概念/行业等); 解析/校验按字段 dtype 分发。"""
|
||||
try:
|
||||
from app.factors.ext_factors import ext_string_fields
|
||||
|
||||
return ext_string_fields()
|
||||
except Exception:
|
||||
return frozenset()
|
||||
|
||||
|
||||
def allowed_fields() -> frozenset[str]:
|
||||
"""条件可引用字段 = 物化列白名单 并入 注册表因子 与 string 扩展字段。
|
||||
|
||||
因子列在历史路径 (compute_signals) 由 materialize_factor_columns 复用
|
||||
评分物化管线补算; 盘中单日快照无滚动窗口, 依赖因子的信号被 inject 以
|
||||
缺列告警跳过 (与日期偏移条件同样的优雅降级)。
|
||||
string 扩展字段 (ext_{表}_{字段}, 概念/行业归属) 只支持 contains/==/!=,
|
||||
在帧组装时由 attach_ext_columns 注入, 不注册为因子 (数值口径约束)。
|
||||
"""
|
||||
from app.factors.registry import all_factors
|
||||
|
||||
return frozenset(ALLOWED_FIELDS | {spec.id for spec in all_factors()} | _string_ext_fields())
|
||||
|
||||
|
||||
def materialize_factor_columns(
|
||||
df: pl.DataFrame,
|
||||
exprs: dict[str, pl.Expr],
|
||||
needed: set[str] | None = None,
|
||||
) -> pl.DataFrame:
|
||||
"""把信号表达式引用、且 df 缺失的注册表因子列补算出来。
|
||||
|
||||
复用评分物化路径 (materialize_scoring_columns) — 与检验/评分同一条计算
|
||||
逻辑, 不引入第二套实现。非注册表列不在此处理 (缺列仍由 inject 告警跳过)。
|
||||
"""
|
||||
if df.is_empty() or not exprs:
|
||||
return df
|
||||
cols = set(df.columns)
|
||||
missing: set[str] = set()
|
||||
for name, roots in expression_dependencies(exprs).items():
|
||||
if needed is not None and name not in needed:
|
||||
continue
|
||||
missing.update(root for root in roots if root not in cols)
|
||||
if not missing:
|
||||
return df
|
||||
from app.factors.registry import all_factors
|
||||
|
||||
factor_ids = {spec.id for spec in all_factors()}
|
||||
to_compute = missing & factor_ids
|
||||
if not to_compute:
|
||||
return df
|
||||
from app.strategy.scoring import materialize_scoring_columns
|
||||
|
||||
return materialize_scoring_columns(df, sorted(to_compute))
|
||||
|
||||
|
||||
# ── 持久化(镜像 strategy/config.py 的写法)──────────────
|
||||
def _dir(data_dir: Path) -> Path:
|
||||
d = data_dir / "user_data" / "custom_signals"
|
||||
@@ -120,22 +181,35 @@ def _parse_days(c: dict, key: str, i: int) -> int:
|
||||
return n
|
||||
|
||||
|
||||
def _parse_right(right: str) -> tuple[str, object]:
|
||||
"""解析右值。返回 ('field', colname) 或 ('const', float)。
|
||||
def _parse_right(right: str, *, string_mode: bool = False) -> tuple[str, object]:
|
||||
"""解析右值。返回 ('field', colname) / ('const', float) / ('const_str', str)。
|
||||
|
||||
接受三种形式:
|
||||
数值模式接受三种形式:
|
||||
- 数字 (int / float / 数字字符串) → 常量
|
||||
- "field:字段名" → 字段引用
|
||||
- 裸字段名 (在白名单内) → 自动视为字段引用
|
||||
(AI 生成偶尔漏写 field: 前缀; 白名单字段名不可能是数字, 无歧义)
|
||||
|
||||
string 模式 (左字段是 string 扩展字段): 只接受非空字符串字面量
|
||||
(概念/行业名), 不支持字段引用 —— "字段A包含字段B" 无业务语义且
|
||||
会与 field: 前缀解析产生歧义。
|
||||
"""
|
||||
if string_mode:
|
||||
if not isinstance(right, str) or not right.strip():
|
||||
raise ValueError("字符串条件的右值必须是非空字符串 (如概念/行业名)")
|
||||
if right.startswith("field:"):
|
||||
raise ValueError("字符串条件不支持字段引用右值, 请填字符串字面量")
|
||||
if len(right) > _MAX_STR_RIGHT:
|
||||
raise ValueError(f"字符串右值过长 (≤{_MAX_STR_RIGHT} 字符): {right[:20]}…")
|
||||
return ("const_str", right.strip())
|
||||
if isinstance(right, (int, float)):
|
||||
return ("const", float(right))
|
||||
if not isinstance(right, str):
|
||||
raise ValueError(f"非法右值: {right!r}")
|
||||
allowed = allowed_fields()
|
||||
if right.startswith("field:"):
|
||||
col = right[len("field:"):]
|
||||
if col not in ALLOWED_FIELDS:
|
||||
if col not in allowed:
|
||||
raise ValueError(f"右值字段不在白名单: {col}")
|
||||
return ("field", col)
|
||||
# 纯数字
|
||||
@@ -144,7 +218,7 @@ def _parse_right(right: str) -> tuple[str, object]:
|
||||
except ValueError:
|
||||
pass
|
||||
# 裸字段名 — 兜底容错, 仍受白名单约束
|
||||
if right in ALLOWED_FIELDS:
|
||||
if right in allowed:
|
||||
return ("field", right)
|
||||
raise ValueError(f"非法右值(应为 field:xxx 或数字): {right!r}")
|
||||
|
||||
@@ -158,20 +232,36 @@ def validate(sig: dict) -> None:
|
||||
raise ValueError("信号 name 不能为空")
|
||||
if sig.get("kind") not in ("entry", "exit", "both"):
|
||||
raise ValueError("kind 必须是 entry / exit / both")
|
||||
timeframe = sig.get("timeframe", TIMEFRAME_DAILY)
|
||||
if timeframe not in (TIMEFRAME_DAILY, TIMEFRAME_INTRADAY):
|
||||
raise ValueError(f"timeframe 必须是 {TIMEFRAME_DAILY} / {TIMEFRAME_INTRADAY}: {timeframe!r}")
|
||||
conds = sig.get("conditions")
|
||||
if not isinstance(conds, list) or len(conds) == 0:
|
||||
raise ValueError("conditions 不能为空")
|
||||
if len(conds) > 8:
|
||||
raise ValueError("conditions 最多 8 条")
|
||||
if timeframe == TIMEFRAME_INTRADAY:
|
||||
_validate_intraday(sig)
|
||||
return
|
||||
string_fields = _string_ext_fields()
|
||||
for i, c in enumerate(conds):
|
||||
if not isinstance(c, dict):
|
||||
raise ValueError(f"第 {i+1} 个条件格式错误")
|
||||
left = c.get("left", "")
|
||||
if left not in ALLOWED_FIELDS:
|
||||
if left not in allowed_fields():
|
||||
raise ValueError(f"第 {i+1} 个条件: 字段 {left!r} 不在白名单")
|
||||
if c.get("op") not in OPS:
|
||||
is_str = left in string_fields
|
||||
if is_str:
|
||||
if c.get("op") not in STRING_OPS:
|
||||
raise ValueError(
|
||||
f"第 {i+1} 个条件: 字符串字段 {left!r} 仅支持 "
|
||||
f"{'/'.join(sorted(STRING_OPS))} 运算符"
|
||||
)
|
||||
elif c.get("op") == "contains":
|
||||
raise ValueError(f"第 {i+1} 个条件: contains 仅用于字符串扩展字段")
|
||||
elif c.get("op") not in OPS:
|
||||
raise ValueError(f"第 {i+1} 个条件: 运算符 {c.get('op')!r} 非法")
|
||||
_parse_right(c.get("right")) # 会校验右值字段/数字
|
||||
_parse_right(c.get("right"), string_mode=is_str) # 会校验右值字段/数字/字符串
|
||||
_parse_days(c, "leftDays", i) # 左字段偏移
|
||||
_parse_days(c, "rightDays", i) # 右字段偏移
|
||||
|
||||
@@ -200,6 +290,7 @@ def build_expressions(signals: list[dict], allow_shift: bool = True) -> dict[str
|
||||
- 编译失败的信号被跳过并告警(不影响其它信号)。
|
||||
"""
|
||||
out: dict[str, pl.Expr] = {}
|
||||
string_fields = _string_ext_fields()
|
||||
for sig in signals:
|
||||
if sig.get("enabled") is False:
|
||||
continue
|
||||
@@ -215,7 +306,12 @@ def build_expressions(signals: list[dict], allow_shift: bool = True) -> dict[str
|
||||
raise ValueError("盘中实时路径不支持日期偏移条件, 已跳过")
|
||||
left = c["left"]
|
||||
op = c["op"]
|
||||
kind, val = _parse_right(c["right"])
|
||||
is_str = left in string_fields
|
||||
if is_str and op not in STRING_OPS:
|
||||
raise ValueError(f"字符串字段 {left!r} 不支持运算符 {op!r}")
|
||||
if op == "contains" and not is_str:
|
||||
raise ValueError(f"contains 仅用于字符串扩展字段: {left!r}")
|
||||
kind, val = _parse_right(c["right"], string_mode=is_str)
|
||||
right_expr = _col(val, right_days) if kind == "field" else val
|
||||
parts.append(_OP_BUILDERS[op](_col(left, left_days), right_expr))
|
||||
combined = parts[0]
|
||||
@@ -273,3 +369,169 @@ def _expr_root_columns(expr: pl.Expr) -> set[str]:
|
||||
return set(names)
|
||||
except Exception:
|
||||
return set()
|
||||
|
||||
|
||||
# ══ 盘中信号(timeframe="intraday")═════════════════════════
|
||||
# 与日线自定义信号同一套 left/op/right 条件结构, 但:
|
||||
# - 字段白名单换成分钟特征(intraday_features.INTRADAY_FEATURES);
|
||||
# - 运算符额外支持 cross_up / cross_down(序列上穿/下穿 另一序列或阈值);
|
||||
# - 不支持 leftDays/rightDays 日期偏移;
|
||||
# - 信号列名前缀 csgi_, 注入对象是分钟特征帧而非日线 enriched。
|
||||
# 语义: 信号输出 = 当日条件组合的上升沿(false→true), 首根 bar 不触发。
|
||||
|
||||
from app.strategy.intraday_features import INTRADAY_FEATURES # noqa: E402
|
||||
|
||||
TIMEFRAME_DAILY = "daily"
|
||||
TIMEFRAME_INTRADAY = "intraday"
|
||||
INTRADAY_PREFIX = "csgi_"
|
||||
INTRADAY_OPS = OPS | {"cross_up", "cross_down"}
|
||||
_EDGE_GROUP = ["symbol", "date"]
|
||||
|
||||
|
||||
def intraday_column_name(signal_id: str) -> str:
|
||||
"""盘中信号 id → 分钟帧列名(加 csgi_ 前缀)。"""
|
||||
return f"{INTRADAY_PREFIX}{signal_id}"
|
||||
|
||||
|
||||
def _parse_right_intraday(right: object) -> tuple[str, object]:
|
||||
"""盘中条件的右值: ('const', float) 或 ('field', 特征名)。"""
|
||||
if isinstance(right, (int, float)):
|
||||
return ("const", float(right))
|
||||
if not isinstance(right, str):
|
||||
raise ValueError(f"非法右值: {right!r}")
|
||||
if right.startswith("field:"):
|
||||
col = right[len("field:"):]
|
||||
if col not in INTRADAY_FEATURES:
|
||||
raise ValueError(f"盘中右值字段不在白名单: {col}")
|
||||
return ("field", col)
|
||||
try:
|
||||
return ("const", float(right))
|
||||
except ValueError:
|
||||
pass
|
||||
if right in INTRADAY_FEATURES:
|
||||
return ("field", right)
|
||||
raise ValueError(f"非法盘中右值(应为 field:特征 或数字): {right!r}")
|
||||
|
||||
|
||||
def _validate_intraday(sig: dict) -> None:
|
||||
"""校验盘中信号定义, 非法抛 ValueError。"""
|
||||
conds = sig.get("conditions")
|
||||
for i, c in enumerate(conds):
|
||||
if not isinstance(c, dict):
|
||||
raise ValueError(f"第 {i+1} 个条件格式错误")
|
||||
left = c.get("left", "")
|
||||
if left not in INTRADAY_FEATURES:
|
||||
raise ValueError(f"第 {i+1} 个条件: 盘中字段 {left!r} 不在白名单")
|
||||
if c.get("op") not in INTRADAY_OPS:
|
||||
raise ValueError(f"第 {i+1} 个条件: 运算符 {c.get('op')!r} 非法(盘中额外支持 cross_up/cross_down)")
|
||||
_parse_right_intraday(c.get("right"))
|
||||
if int(c.get("leftDays", 0) or 0) or int(c.get("rightDays", 0) or 0):
|
||||
raise ValueError(f"第 {i+1} 个条件: 盘中信号不支持日期偏移(leftDays/rightDays)")
|
||||
min_bars = sig.get("min_bars", 0)
|
||||
try:
|
||||
n = int(min_bars)
|
||||
except (TypeError, ValueError):
|
||||
raise ValueError(f"min_bars 必须是整数: {min_bars!r}") # noqa: B904
|
||||
if n < 0 or n > 240:
|
||||
raise ValueError(f"min_bars 必须在 0..240 之间: {n}")
|
||||
|
||||
|
||||
def build_intraday_expressions(signals: list[dict]) -> dict[str, pl.Expr]:
|
||||
"""把盘中信号编译为特征帧上的「条件」表达式(AND 组合, 未做上升沿)。
|
||||
|
||||
表达式在 intraday_features.build_feature_frame 产出的帧上求值;
|
||||
上升沿须通过 apply_intraday_edges 在 DataFrame 层两步计算 —
|
||||
对已含 .over() 窗口的组合表达式直接 shift().over() 是窗口嵌套,
|
||||
Polars 会返回全 null。编译失败的信号跳过并告警。
|
||||
"""
|
||||
out: dict[str, pl.Expr] = {}
|
||||
for sig in signals:
|
||||
if sig.get("enabled") is False or sig.get("timeframe") != TIMEFRAME_INTRADAY:
|
||||
continue
|
||||
try:
|
||||
parts: list[pl.Expr] = []
|
||||
for c in sig["conditions"]:
|
||||
left = pl.col(c["left"])
|
||||
kind, val = _parse_right_intraday(c["right"])
|
||||
op = c["op"]
|
||||
if op == "cross_up":
|
||||
# 前一根 bar 未满足 且 当前 bar 满足; 右值为常量时不 shift 字面量
|
||||
if kind == "field":
|
||||
prev_ok = left.shift(1).over(_EDGE_GROUP) <= pl.col(val).shift(1).over(_EDGE_GROUP)
|
||||
cur_ok = left > pl.col(val)
|
||||
else:
|
||||
prev_ok = left.shift(1).over(_EDGE_GROUP) <= val
|
||||
cur_ok = left > val
|
||||
parts.append(prev_ok & cur_ok)
|
||||
elif op == "cross_down":
|
||||
if kind == "field":
|
||||
prev_ok = left.shift(1).over(_EDGE_GROUP) >= pl.col(val).shift(1).over(_EDGE_GROUP)
|
||||
cur_ok = left < pl.col(val)
|
||||
else:
|
||||
prev_ok = left.shift(1).over(_EDGE_GROUP) >= val
|
||||
cur_ok = left < val
|
||||
parts.append(prev_ok & cur_ok)
|
||||
else:
|
||||
right = pl.col(val) if kind == "field" else val
|
||||
parts.append(_OP_BUILDERS[op](left, right))
|
||||
combined = parts[0]
|
||||
for p in parts[1:]:
|
||||
combined = combined & p
|
||||
out[intraday_column_name(sig["id"])] = combined
|
||||
except Exception as e:
|
||||
logger.warning("intraday signal compile failed %s: %s", sig.get("id"), e)
|
||||
return out
|
||||
|
||||
|
||||
def apply_intraday_edges(frame: pl.DataFrame, exprs: dict[str, pl.Expr]) -> pl.DataFrame:
|
||||
"""对特征帧求值盘中信号: 先算条件列, 再取「当日条件上升沿」为布尔列。
|
||||
|
||||
上升沿: 条件 false→true 的那根 bar 为 true; 首根 bar(前值为 null)不触发;
|
||||
条件含 null(特征不足)视为 false。四条消费路径(监控/实盘/回测/回放)
|
||||
必须共用本函数, 保证口径一致。
|
||||
"""
|
||||
if frame.is_empty() or not exprs:
|
||||
return frame
|
||||
df = frame.with_columns([e.fill_null(False).alias(n) for n, e in exprs.items()])
|
||||
return df.with_columns([
|
||||
(
|
||||
pl.col(n)
|
||||
& ~pl.col(n).shift(1).over(_EDGE_GROUP).fill_null(True)
|
||||
).cast(pl.Boolean).alias(n)
|
||||
for n in exprs
|
||||
])
|
||||
|
||||
|
||||
# ── 盘中信号定义加载(带指纹缓存: 引擎/监控高频路径用) ──────────
|
||||
_intraday_cache: dict[Path, tuple[object, list[dict]]] = {}
|
||||
|
||||
|
||||
def _dir_fingerprint(d: Path) -> tuple:
|
||||
"""目录内 *.json 的 (文件名, mtime) 指纹 — 创建/删除/编辑都会变化。"""
|
||||
try:
|
||||
return tuple(sorted((f.name, f.stat().st_mtime_ns) for f in d.glob("*.json")))
|
||||
except OSError:
|
||||
return ()
|
||||
|
||||
|
||||
def load_intraday_all(data_dir: Path) -> list[dict]:
|
||||
"""读取全部启用的盘中信号定义(带缓存)。
|
||||
|
||||
盘中评估与引擎注入每分钟执行, 不宜每次全量读盘; save/delete 端点
|
||||
调用 invalidate_intraday_cache() 主动失效。
|
||||
"""
|
||||
d = _dir(data_dir)
|
||||
fp = _dir_fingerprint(d)
|
||||
cached = _intraday_cache.get(data_dir)
|
||||
if cached is not None and cached[0] == fp:
|
||||
return cached[1]
|
||||
sigs = [
|
||||
s for s in load_all(data_dir)
|
||||
if s.get("timeframe") == TIMEFRAME_INTRADAY and s.get("enabled") is not False
|
||||
]
|
||||
_intraday_cache[data_dir] = (fp, sigs)
|
||||
return sigs
|
||||
|
||||
|
||||
def invalidate_intraday_cache() -> None:
|
||||
_intraday_cache.clear()
|
||||
|
||||
@@ -32,8 +32,12 @@ _FENCED_JSON_RE = re.compile(r"```(?:json)?\s*\n?(.*?)```", re.DOTALL)
|
||||
|
||||
|
||||
def _format_fields() -> str:
|
||||
"""按类别格式化白名单字段(key(中文标签)),供 LLM 参考。"""
|
||||
allowed = custom_signals.ALLOWED_FIELDS
|
||||
"""按类别格式化白名单字段(key(中文标签)), 供 LLM 参考.
|
||||
|
||||
行情/指标类物理列之后追加注册表因子, 分组与 /api/custom-signals/options
|
||||
的 factor 分组一致: 因子是预计算因子值, 同样可作为条件字段比较.
|
||||
"""
|
||||
allowed = custom_signals.allowed_fields()
|
||||
lines: list[str] = []
|
||||
quote = sorted(f for f in _QUOTE_FIELDS if f in allowed)
|
||||
lines.append(
|
||||
@@ -46,6 +50,27 @@ def _format_fields() -> str:
|
||||
f"{label}: "
|
||||
+ ", ".join(f"{f}({ENRICHED_COLUMNS.get(f, f)})" for f in fields)
|
||||
)
|
||||
from app.factors.registry import all_factors
|
||||
|
||||
factor_groups: dict[str, list[str]] = {}
|
||||
for spec in all_factors():
|
||||
if spec.id in custom_signals.ALLOWED_FIELDS:
|
||||
continue # 已作为物理列出现在清单里
|
||||
label = spec.label
|
||||
if spec.asset_types == frozenset({"stock"}):
|
||||
label += "·仅股票"
|
||||
factor_groups.setdefault(spec.group, []).append(f"{spec.id}({label})")
|
||||
for group, items in sorted(factor_groups.items()):
|
||||
lines.append(f"因子·{group}: " + ", ".join(sorted(items)))
|
||||
# string 扩展字段 (概念/行业归属): 只支持 contains/==/!=, 右值为字符串字面量
|
||||
from app.factors.ext_factors import ext_string_field_entries
|
||||
|
||||
str_entries = ext_string_field_entries()
|
||||
if str_entries:
|
||||
lines.append(
|
||||
"字符串字段(仅 contains/==/!=): "
|
||||
+ ", ".join(f"{e['key']}({e['label']})" for e in str_entries)
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
@@ -53,12 +78,15 @@ _SYSTEM_TEMPLATE = """你是A股量化信号设计专家。用户会描述一个
|
||||
|
||||
可用字段(白名单,只能使用以下字段,禁止自造或使用白名单之外的字段):
|
||||
{fields}
|
||||
其中「因子·」开头的行是平台预计算的因子值(动量/波动/量价等衍生特征),可直接比较数值构造条件。
|
||||
|
||||
运算符(op):> >= < <= == !=
|
||||
字符串字段额外支持 contains(包含子串, 如概念/行业归属判断), 右值为字符串字面量, 如 "AI"、"半导体".
|
||||
|
||||
右值(right):
|
||||
- 数字:写字符串形式,如 "2"、"3000"、"0.05"
|
||||
- 另一字段:必须带 "field:" 前缀,如 "field:ma20";严禁裸写字段名,如 "macd_dea" 应写成 "field:macd_dea"
|
||||
- 字符串字面量: 仅当左字段是「字符串字段」时使用(配合 contains/==/!=), 如所属概念包含AI写成 {{"left": "字符串字段", "op": "contains", "right": "AI", "leftDays": 0, "rightDays": 0}}
|
||||
|
||||
日期偏移(leftDays / rightDays):取 N 个交易日前的值,0 = 当日最新;范围 0~{max_days}。只有明确需要「前N日」时才使用偏移。
|
||||
|
||||
@@ -132,8 +160,8 @@ def _normalize_condition(c: object) -> dict:
|
||||
if not isinstance(right, str) or not right.strip():
|
||||
raise ValueError(f"右值非法: {right!r}")
|
||||
right = right.strip()
|
||||
# 兜底: AI 偶尔漏写 field: 前缀的裸字段名, 补全为规范形式
|
||||
if not right.startswith("field:") and right in custom_signals.ALLOWED_FIELDS:
|
||||
# 兜底: AI 偶尔漏写 field: 前缀的裸字段名, 补全为规范形式 (含因子字段)
|
||||
if not right.startswith("field:") and right in custom_signals.allowed_fields():
|
||||
right = f"field:{right}"
|
||||
return {
|
||||
"left": str(left),
|
||||
|
||||
+156
-28
@@ -12,6 +12,7 @@ import sys
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass, field, replace
|
||||
from datetime import date
|
||||
from pathlib import Path
|
||||
@@ -20,6 +21,7 @@ from typing import Any
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
|
||||
from app.config import settings
|
||||
from app.strategy.scoring import (
|
||||
SCORING_DIRECTION_LOW,
|
||||
effective_scoring,
|
||||
@@ -951,6 +953,19 @@ class StrategyEngine:
|
||||
strategy_id=strategy_id,
|
||||
exit_signal_hits=exit_signal_hits,
|
||||
)
|
||||
# 盘中信号列注入(csgi_): 实盘扫描与分钟回测共用本路径 — 与监控评估
|
||||
# 同一特征构造器, 单点注入保证三处口径一致。
|
||||
history = self._inject_intraday_signal_columns(history)
|
||||
missing_csgi = [
|
||||
name for name in s.required_features
|
||||
if name.startswith("csgi_") and name not in history.columns
|
||||
]
|
||||
if missing_csgi:
|
||||
raise ValueError(
|
||||
"策略引用了未定义的盘中信号: "
|
||||
+ ", ".join(sorted(missing_csgi))
|
||||
+ " — 请先在「自定义信号」中创建(timeframe=intraday)后再运行"
|
||||
)
|
||||
if s.minute_daily_bars > 0:
|
||||
df = s.filter_minute_history_fn(history, params, daily=context.daily_history)
|
||||
else:
|
||||
@@ -982,6 +997,15 @@ class StrategyEngine:
|
||||
+ ", ".join(sorted(missing_csg))
|
||||
+ " — 请先在「自定义信号」管理中创建对应信号后再运行"
|
||||
)
|
||||
missing_csgi = [
|
||||
name for name in s.required_features
|
||||
if name.startswith("csgi_")
|
||||
]
|
||||
if missing_csgi:
|
||||
raise ValueError(
|
||||
"盘中信号仅可用于分钟策略(timeframes=['1m']), 日线策略不支持: "
|
||||
+ ", ".join(sorted(missing_csgi))
|
||||
)
|
||||
df = s.filter_history_fn(df, params)
|
||||
if "date" in df.columns:
|
||||
df = df.filter(pl.col("date") == as_of)
|
||||
@@ -1097,8 +1121,15 @@ class StrategyEngine:
|
||||
overrides_map: dict | None = None,
|
||||
*,
|
||||
strategy_ids: list[str] | None = None,
|
||||
parallel: bool = True,
|
||||
) -> dict[str, StrategyResult]:
|
||||
"""批量执行策略;当前数据、历史和矩阵均来自同一个调用上下文。"""
|
||||
"""批量执行策略;当前数据、历史和矩阵均来自同一个调用上下文。
|
||||
|
||||
parallel=True 时用有界线程池并发执行: 策略对 context 是只读纯函数
|
||||
(polars 计算释放 GIL), 并发不改变结果, 逐策略耗时日志不变。composite
|
||||
子策略的递归 run_all 以 parallel=False 调用, 保证嵌套时线程总数仍
|
||||
不超过 worker 上限, 不随叠加层数放大。
|
||||
"""
|
||||
if context.current is None:
|
||||
raise ValueError("strategy run_all context requires current data")
|
||||
df = context.current
|
||||
@@ -1120,37 +1151,16 @@ class StrategyEngine:
|
||||
raise ValueError("selected strategies require history data")
|
||||
|
||||
shared_matrix = context.market
|
||||
matrix_strats = [
|
||||
(sid, strategy)
|
||||
for sid, strategy in selected
|
||||
if strategy.execution_backend == "matrix_native"
|
||||
]
|
||||
if (
|
||||
shared_matrix is None
|
||||
and matrix_strats
|
||||
and shared_history is not None
|
||||
and not shared_history.is_empty()
|
||||
):
|
||||
from app.backtest.matrix import build_market_data_matrix
|
||||
|
||||
field_columns: set[str] = set()
|
||||
for sid, strategy in matrix_strats:
|
||||
field_columns.update(
|
||||
self._matrix_field_columns(
|
||||
strategy,
|
||||
overrides_map.get(sid),
|
||||
params_map.get(sid),
|
||||
)
|
||||
)
|
||||
shared_matrix = build_market_data_matrix(
|
||||
shared_history,
|
||||
field_columns=field_columns,
|
||||
if shared_matrix is None:
|
||||
shared_matrix = self.build_shared_matrix(
|
||||
context, selected, params_map, overrides_map
|
||||
)
|
||||
|
||||
results: dict[str, StrategyResult] = {}
|
||||
|
||||
for sid, _ in selected:
|
||||
results[sid] = self.run(
|
||||
def _execute(sid: str) -> tuple[str, StrategyResult]:
|
||||
started = time.perf_counter()
|
||||
result = self.run(
|
||||
sid,
|
||||
replace(
|
||||
context,
|
||||
@@ -1161,9 +1171,77 @@ class StrategyEngine:
|
||||
params=params_map.get(sid),
|
||||
overrides=overrides_map.get(sid),
|
||||
)
|
||||
elapsed_ms = (time.perf_counter() - started) * 1000
|
||||
# >=1s 打 INFO 供热点归因 (哪些策略吃掉了 run_all 的大头), 其余 DEBUG 防噪。
|
||||
log_fn = logger.info if elapsed_ms >= 1000 else logger.debug
|
||||
log_fn(
|
||||
"run_all: strategy %s took %.0fms (total=%d)",
|
||||
sid,
|
||||
elapsed_ms,
|
||||
result.total,
|
||||
)
|
||||
return sid, result
|
||||
|
||||
workers = min(settings.strategy_run_all_workers, len(selected))
|
||||
if parallel and workers > 1:
|
||||
with ThreadPoolExecutor(
|
||||
max_workers=workers, thread_name_prefix="strategy-run"
|
||||
) as pool:
|
||||
futures = [pool.submit(_execute, sid) for sid, _ in selected]
|
||||
# 按原顺序收集: 首个失败策略的异常语义与串行执行一致。
|
||||
for future in futures:
|
||||
sid, result = future.result()
|
||||
results[sid] = result
|
||||
else:
|
||||
for sid, _ in selected:
|
||||
sid, result = _execute(sid)
|
||||
results[sid] = result
|
||||
|
||||
return results
|
||||
|
||||
def build_shared_matrix(
|
||||
self,
|
||||
context: StrategyDataContext,
|
||||
selected: list[tuple[str, StrategyDef]],
|
||||
params_map: dict | None = None,
|
||||
overrides_map: dict | None = None,
|
||||
):
|
||||
"""按所选策略的字段并集构建市场数据矩阵; 无矩阵策略或无历史时返回 None。
|
||||
|
||||
渐进式 run_all (逐策略执行) 也用它一次建好并集矩阵后放入 context.market,
|
||||
避免每个 matrix_native 策略重复构建同一份大矩阵 (全市场历史, 秒级)。
|
||||
"""
|
||||
params_map = params_map or {}
|
||||
overrides_map = overrides_map or {}
|
||||
matrix_strats = [
|
||||
(sid, strategy)
|
||||
for sid, strategy in selected
|
||||
if strategy.execution_backend == "matrix_native"
|
||||
]
|
||||
history = context.history
|
||||
if not matrix_strats or history is None or history.is_empty():
|
||||
return None
|
||||
|
||||
from app.backtest.matrix import build_market_data_matrix
|
||||
|
||||
field_columns: set[str] = set()
|
||||
for sid, strategy in matrix_strats:
|
||||
field_columns.update(
|
||||
self._matrix_field_columns(
|
||||
strategy,
|
||||
overrides_map.get(sid),
|
||||
params_map.get(sid),
|
||||
)
|
||||
)
|
||||
matrix_t0 = time.perf_counter()
|
||||
matrix = build_market_data_matrix(history, field_columns=field_columns)
|
||||
logger.info(
|
||||
"run_all: shared matrix built in %.0fms (fields=%d)",
|
||||
(time.perf_counter() - matrix_t0) * 1000,
|
||||
len(field_columns),
|
||||
)
|
||||
return matrix
|
||||
|
||||
@staticmethod
|
||||
def _matrix_field_columns(
|
||||
strategy: StrategyDef,
|
||||
@@ -1244,6 +1322,11 @@ class StrategyEngine:
|
||||
basic_filter = dict(strategy.basic_filter or {})
|
||||
if overrides.get("basic_filter"):
|
||||
basic_filter.update(overrides["basic_filter"])
|
||||
# 策略扫描的运行期过滤同样要按资产类型中和股票专属键 (boards/价格界),
|
||||
# 否则 ETF 候选在矩阵掩码阶段被静默清零 (#215); 函数级导入避免
|
||||
# engine ↔ backtest.strategy 的模块级循环依赖 (与上方 matrix 导入同模式)
|
||||
from app.backtest.strategy import _basic_filter_for_asset
|
||||
basic_filter = _basic_filter_for_asset(basic_filter, context.asset_type)
|
||||
scoring = effective_scoring(strategy.meta.get("scoring"), overrides)
|
||||
asset_mask = None
|
||||
if pool:
|
||||
@@ -1393,6 +1476,9 @@ class StrategyEngine:
|
||||
params_map={},
|
||||
overrides_map=overrides_map,
|
||||
strategy_ids=child_ids,
|
||||
# 嵌套调用串行: 父级 worker 已并发, 子级再开池会使线程总数随叠加
|
||||
# 层数放大 (4×4×...), 超出并发闸与核数的合理范围。
|
||||
parallel=False,
|
||||
)
|
||||
ordered_results = [child_results[cid] for cid in child_ids]
|
||||
|
||||
@@ -1550,6 +1636,48 @@ class StrategyEngine:
|
||||
"turnover_rate", "change_pct", "pre_close",
|
||||
)
|
||||
|
||||
def _user_data_dir(self) -> Path | None:
|
||||
"""从策略目录推导 data_dir(…/strategies/custom → data_dir)。推不出则跳过注入。"""
|
||||
for d in self._strategy_dirs:
|
||||
if d.name == "custom" and d.parent.name == "strategies":
|
||||
return d.parent.parent
|
||||
return None
|
||||
|
||||
def _inject_intraday_signal_columns(self, minute_df: pl.DataFrame) -> pl.DataFrame:
|
||||
"""向当日分钟K帧注入自定义盘中信号列(csgi_, 当日条件上升沿)。
|
||||
|
||||
单点注入: 实盘分钟扫描与分钟回测 worker 共用本方法, 特征计算与
|
||||
监控评估同源(intraday_features), 保证口径一致。无定义/帧为空时原样返回。
|
||||
"""
|
||||
if minute_df is None or minute_df.is_empty() or "datetime" not in minute_df.columns:
|
||||
return minute_df
|
||||
data_dir = self._user_data_dir()
|
||||
if data_dir is None:
|
||||
return minute_df
|
||||
try:
|
||||
from app.strategy import custom_signals
|
||||
from app.strategy.intraday_features import build_feature_frame
|
||||
|
||||
definitions = custom_signals.load_intraday_all(data_dir)
|
||||
if not definitions:
|
||||
return minute_df
|
||||
exprs = custom_signals.build_intraday_expressions(definitions)
|
||||
if not exprs:
|
||||
return minute_df
|
||||
frame = build_feature_frame(minute_df)
|
||||
if frame.is_empty():
|
||||
return minute_df
|
||||
evaluated = custom_signals.apply_intraday_edges(frame, exprs).select(
|
||||
["symbol", "datetime", *exprs.keys()]
|
||||
)
|
||||
return minute_df.join(evaluated, on=["symbol", "datetime"], how="left").with_columns(
|
||||
[pl.col(name).fill_null(False).cast(pl.Boolean).alias(name) for name in exprs]
|
||||
)
|
||||
except Exception as e:
|
||||
# 注入失败不阻断策略执行: 未注入列会由 required_features 校验兜底报错
|
||||
logger.warning("intraday signal inject failed: %s", e)
|
||||
return minute_df
|
||||
|
||||
@staticmethod
|
||||
def _join_basic_columns(df: pl.DataFrame, current: pl.DataFrame) -> pl.DataFrame:
|
||||
"""把 enriched 快照列按 symbol 联到分钟策略输出上, 只补 df 缺失的列。"""
|
||||
|
||||
@@ -0,0 +1,178 @@
|
||||
"""盘中信号特征帧 — 当日已完成分钟 K → 数值特征序列(每根已完成 bar 一行)。
|
||||
|
||||
单一口径源: 监控评估(quote_service) / 分钟策略执行(引擎注入) / 分钟回测 / 回放验证
|
||||
共用本模块构造特征, 保证四条路径对同一段分钟数据产出完全一致的特征值。
|
||||
|
||||
设计:
|
||||
- 会话对齐: 滚动窗口只在本时段(09:30-11:30 / 13:00-15:00)内回看,
|
||||
不跨午休、不跨日; 日累计类特征(vwap/当日高低)按交易日分组。
|
||||
- null 语义: 窗口不足、基准为零、缺昨收、开盘未满 30 分钟 → 特征为 null,
|
||||
任何条件对 null 判 false, 绝不把数据不足伪装成 0。
|
||||
- 纯函数: 不做 IO, 分钟帧由调用方传入(cutoff 过滤也由调用方决定)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
import polars as pl
|
||||
|
||||
from app.market_time import CN_TZ
|
||||
|
||||
# ── 特征白名单(供 custom_signals 校验与 /options 展示) ──────────
|
||||
# 字段 → 中文标签。数值均为「每根已完成 bar 一个值」的序列。
|
||||
INTRADAY_FEATURES: dict[str, str] = {
|
||||
"price": "现价",
|
||||
"vwap": "分时均价",
|
||||
"pct_vs_prev_close": "相对昨收涨跌幅",
|
||||
"pct_from_open": "相对开盘涨跌幅",
|
||||
"vol_ratio_1m_today": "1分钟放量比(今日基准)",
|
||||
"vol_ratio_3m_today": "3分钟放量比(今日基准)",
|
||||
"vol_ratio_5m_today": "5分钟放量比(今日基准)",
|
||||
"day_high_dist": "距当日最高价",
|
||||
"day_low_dist": "距当日最低价",
|
||||
"open_30m_high_dist": "距开盘30分钟最高价",
|
||||
"open_30m_low_dist": "距开盘30分钟最低价",
|
||||
}
|
||||
|
||||
# 滚动窗口特征的窗口长度(字段名后缀 → bar 数)
|
||||
_VOL_WINDOWS = {1: "vol_ratio_1m_today", 3: "vol_ratio_3m_today", 5: "vol_ratio_5m_today"}
|
||||
|
||||
_REQUIRED_COLS = ("symbol", "datetime", "close", "volume", "amount")
|
||||
_DAY_KEY = ["symbol", "date"]
|
||||
_SESSION_KEY = ["symbol", "date", "session"]
|
||||
|
||||
|
||||
def _naive(dt: datetime) -> datetime | None:
|
||||
"""统一为北京墙钟 naive(与分钟存储契约一致)。"""
|
||||
if not isinstance(dt, datetime):
|
||||
return None
|
||||
if dt.tzinfo is not None:
|
||||
return dt.astimezone(CN_TZ).replace(tzinfo=None)
|
||||
return dt
|
||||
|
||||
|
||||
def build_feature_frame(
|
||||
minute_df: pl.DataFrame,
|
||||
*,
|
||||
prev_close: dict[str, float] | None = None,
|
||||
cutoff: datetime | None = None,
|
||||
) -> pl.DataFrame:
|
||||
"""把分钟 K 帧编译为特征帧。
|
||||
|
||||
参数:
|
||||
minute_df: 列含 symbol/datetime/open/high/low/close/volume/amount(后四列必需,
|
||||
OHLC 缺失时相关特征降级为 null)。可包含多个交易日, 特征按日分组。
|
||||
prev_close: 映射 symbol → 昨收(已复权口径需与分钟价一致); 缺失标的的
|
||||
pct_vs_prev_close 为 null。
|
||||
cutoff: 只使用严格早于 cutoff 的 bar(盘中传入当前分钟; 回放/回测传 None)。
|
||||
|
||||
返回: symbol/datetime + INTRADAY_FEATURES 全部特征列(Float64, 可 null)。
|
||||
"""
|
||||
empty = pl.DataFrame(
|
||||
schema={"symbol": pl.Utf8, "datetime": pl.Datetime, "date": pl.Date}
|
||||
| {name: pl.Float64 for name in INTRADAY_FEATURES}
|
||||
)
|
||||
if minute_df is None or minute_df.is_empty() or not set(_REQUIRED_COLS).issubset(minute_df.columns):
|
||||
return empty
|
||||
|
||||
df = minute_df
|
||||
if "symbol" in df.columns:
|
||||
df = df.with_columns(pl.col("symbol").cast(pl.Utf8))
|
||||
dt_expr = pl.col("datetime")
|
||||
if df.schema["datetime"].time_zone is not None:
|
||||
dt_expr = dt_expr.dt.convert_time_zone(CN_TZ.key).dt.replace_time_zone(None)
|
||||
df = df.with_columns(dt_expr.alias("datetime"))
|
||||
if cutoff is not None:
|
||||
cut = _naive(cutoff)
|
||||
if cut is not None:
|
||||
df = df.filter(pl.col("datetime") < cut)
|
||||
df = df.drop_nulls("datetime").sort(["symbol", "datetime"])
|
||||
if df.is_empty():
|
||||
return empty
|
||||
|
||||
df = df.with_columns(
|
||||
pl.col("datetime").dt.date().alias("date"),
|
||||
# 会话归属: 13:00 及以后为午后续时段, 滚动窗口不与上午合并
|
||||
pl.when(pl.col("datetime").dt.hour() >= 13).then(1).otherwise(0).alias("session"),
|
||||
)
|
||||
df = df.with_columns(pl.int_range(pl.len()).over(_SESSION_KEY).alias("session_idx"))
|
||||
|
||||
cols = {"price": pl.col("close").cast(pl.Float64)}
|
||||
|
||||
# ── 日累计特征(跨上午/下午累计) ──
|
||||
if {"volume", "amount"}.issubset(df.columns):
|
||||
cum_vol = pl.col("volume").cast(pl.Float64).cum_sum().over(_DAY_KEY)
|
||||
cum_amt = pl.col("amount").cast(pl.Float64).cum_sum().over(_DAY_KEY)
|
||||
cols["vwap"] = pl.when(cum_vol > 0).then(cum_amt / (cum_vol * 100.0))
|
||||
else:
|
||||
cols["vwap"] = pl.lit(None, dtype=pl.Float64)
|
||||
|
||||
if prev_close:
|
||||
pc = pl.DataFrame(
|
||||
{"symbol": list(prev_close.keys()), "_prev_close": [float(v) for v in prev_close.values()]}
|
||||
)
|
||||
df = df.join(pc, on="symbol", how="left")
|
||||
cols["pct_vs_prev_close"] = pl.when(
|
||||
pl.col("_prev_close").is_not_null() & (pl.col("_prev_close") > 0)
|
||||
).then(pl.col("close") / pl.col("_prev_close") - 1.0)
|
||||
else:
|
||||
cols["pct_vs_prev_close"] = pl.lit(None, dtype=pl.Float64)
|
||||
|
||||
if "open" in df.columns:
|
||||
day_open = pl.col("open").cast(pl.Float64).first().over(_DAY_KEY)
|
||||
cols["pct_from_open"] = pl.when(day_open > 0).then(pl.col("close") / day_open - 1.0)
|
||||
else:
|
||||
cols["pct_from_open"] = pl.lit(None, dtype=pl.Float64)
|
||||
|
||||
if "high" in df.columns:
|
||||
day_high = pl.col("high").cast(pl.Float64).cum_max().over(_DAY_KEY)
|
||||
cols["day_high_dist"] = pl.when(day_high > 0).then(pl.col("close") / day_high - 1.0)
|
||||
else:
|
||||
cols["day_high_dist"] = pl.lit(None, dtype=pl.Float64)
|
||||
|
||||
if "low" in df.columns:
|
||||
day_low = pl.col("low").cast(pl.Float64).cum_min().over(_DAY_KEY)
|
||||
cols["day_low_dist"] = pl.when(day_low > 0).then(pl.col("close") / day_low - 1.0)
|
||||
else:
|
||||
cols["day_low_dist"] = pl.lit(None, dtype=pl.Float64)
|
||||
|
||||
# ── 滚动放量比(今日基准): 当前 N 根 bar 量和 / 此前 N 根 bar 量和 ──
|
||||
# 滚动窗口按(symbol, date, session)分组 → 不跨午休、不跨日; 窗口不满自然为 null。
|
||||
if "volume" in df.columns:
|
||||
vol = pl.col("volume").cast(pl.Float64)
|
||||
for n, name in _VOL_WINDOWS.items():
|
||||
win = vol.rolling_sum(n).over(_SESSION_KEY)
|
||||
prev_win = win.shift(n).over(_SESSION_KEY)
|
||||
cols[name] = pl.when(prev_win > 0).then(win / prev_win)
|
||||
else:
|
||||
for name in _VOL_WINDOWS.values():
|
||||
cols[name] = pl.lit(None, dtype=pl.Float64)
|
||||
|
||||
# ── 开盘 30 分钟高低点: 上午时段第 30 根 bar 的累计高/低, 全日广播 ──
|
||||
if "high" in df.columns:
|
||||
marker_h = (
|
||||
pl.when((pl.col("session") == 0) & (pl.col("session_idx") == 29))
|
||||
.then(pl.col("high").cast(pl.Float64).cum_max().over(_DAY_KEY))
|
||||
.otherwise(None)
|
||||
.forward_fill()
|
||||
.over(_DAY_KEY)
|
||||
)
|
||||
cols["open_30m_high_dist"] = pl.when(marker_h > 0).then(pl.col("close") / marker_h - 1.0)
|
||||
else:
|
||||
cols["open_30m_high_dist"] = pl.lit(None, dtype=pl.Float64)
|
||||
|
||||
if "low" in df.columns:
|
||||
marker_l = (
|
||||
pl.when((pl.col("session") == 0) & (pl.col("session_idx") == 29))
|
||||
.then(pl.col("low").cast(pl.Float64).cum_min().over(_DAY_KEY))
|
||||
.otherwise(None)
|
||||
.forward_fill()
|
||||
.over(_DAY_KEY)
|
||||
)
|
||||
cols["open_30m_low_dist"] = pl.when(marker_l > 0).then(pl.col("close") / marker_l - 1.0)
|
||||
else:
|
||||
cols["open_30m_low_dist"] = pl.lit(None, dtype=pl.Float64)
|
||||
|
||||
return df.with_columns([expr.cast(pl.Float64).alias(name) for name, expr in cols.items()]).select(
|
||||
["symbol", "date", "datetime", *INTRADAY_FEATURES.keys()]
|
||||
)
|
||||
@@ -1,13 +1,25 @@
|
||||
"""监控中心专用的日内分时穿越信号。"""
|
||||
"""监控中心专用的日内分时信号评估器。
|
||||
|
||||
v2: 特征计算与条件求值统一走 intraday_features 特征帧 + custom_signals 的
|
||||
盘中表达式编译 — 与分钟策略执行/分钟回测/回放验证同一条口径。
|
||||
|
||||
- 内置 4 个分时穿越信号(signal_intraday_*)由同一表达式机制生成, 列名不变,
|
||||
存量监控规则零迁移;
|
||||
- 自定义盘中信号(timeframe="intraday", csgi_ 前缀)与内置信号一并评估注入。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
import polars as pl
|
||||
|
||||
from app.market_time import CN_TZ
|
||||
from app.strategy import custom_signals
|
||||
from app.strategy.intraday_features import build_feature_frame
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
INTRADAY_SIGNAL_LABELS: dict[str, str] = {
|
||||
"signal_intraday_avg_cross_up": "分时价格上穿均价",
|
||||
@@ -16,22 +28,41 @@ INTRADAY_SIGNAL_LABELS: dict[str, str] = {
|
||||
"signal_intraday_zero_cross_down": "分时价格下穿0轴",
|
||||
}
|
||||
INTRADAY_SIGNAL_FIELDS = frozenset(INTRADAY_SIGNAL_LABELS)
|
||||
_LEGACY_MIN_BARS = 2 # 旧实现要求至少两根已完成 bar 才判穿越, 语义保持
|
||||
|
||||
|
||||
def uses_intraday_signals(rule: dict) -> bool:
|
||||
"""规则是否引用盘中信号列(内置 4 个或自定义 csgi_)。"""
|
||||
return any(
|
||||
c.get("op") == "truth" and c.get("field") in INTRADAY_SIGNAL_FIELDS
|
||||
(
|
||||
isinstance(c, dict)
|
||||
and c.get("op") == "truth"
|
||||
and (c.get("field") in INTRADAY_SIGNAL_FIELDS or str(c.get("field", "")).startswith(custom_signals.INTRADAY_PREFIX))
|
||||
)
|
||||
for c in rule.get("conditions", [])
|
||||
if isinstance(c, dict)
|
||||
)
|
||||
|
||||
|
||||
def _finite(value: Any) -> float | None:
|
||||
try:
|
||||
number = float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return number if math.isfinite(number) else None
|
||||
def _legacy_builtin_definitions() -> list[dict]:
|
||||
"""内置 4 个分时穿越信号的等价定义(与 v1 逐字节同口径)。
|
||||
|
||||
v1 语义: 上穿 = 前一根 bar 未满足且当前 bar 满足 —— 与
|
||||
build_intraday_expressions 的「条件上升沿」完全一致。
|
||||
"""
|
||||
return [
|
||||
{"id": "signal_intraday_avg_cross_up", "timeframe": "intraday", "enabled": True,
|
||||
"conditions": [{"left": "price", "op": "cross_up", "right": "field:vwap"}],
|
||||
"min_bars": _LEGACY_MIN_BARS},
|
||||
{"id": "signal_intraday_avg_cross_down", "timeframe": "intraday", "enabled": True,
|
||||
"conditions": [{"left": "price", "op": "cross_down", "right": "field:vwap"}],
|
||||
"min_bars": _LEGACY_MIN_BARS},
|
||||
{"id": "signal_intraday_zero_cross_up", "timeframe": "intraday", "enabled": True,
|
||||
"conditions": [{"left": "pct_vs_prev_close", "op": "cross_up", "right": 0}],
|
||||
"min_bars": _LEGACY_MIN_BARS},
|
||||
{"id": "signal_intraday_zero_cross_down", "timeframe": "intraday", "enabled": True,
|
||||
"conditions": [{"left": "pct_vs_prev_close", "op": "cross_down", "right": 0}],
|
||||
"min_bars": _LEGACY_MIN_BARS},
|
||||
]
|
||||
|
||||
|
||||
def _naive_datetime(value: Any) -> datetime | None:
|
||||
@@ -43,7 +74,7 @@ def _naive_datetime(value: Any) -> datetime | None:
|
||||
|
||||
|
||||
class IntradaySignalEvaluator:
|
||||
"""按已完成的一分钟 K 线生成边沿触发信号。"""
|
||||
"""按已完成的一分钟 K 线评估盘中信号(边沿触发, 新 bar 出现才可能触发)。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._last_bar: dict[tuple[str, str], datetime] = {}
|
||||
@@ -56,86 +87,92 @@ class IntradaySignalEvaluator:
|
||||
prev_close: dict[str, float],
|
||||
asset_type: str,
|
||||
now: datetime,
|
||||
signals: list[dict] | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""返回本分钟触发信号的行列表(每 symbol 一行, 仅新出现的 bar 触发)。"""
|
||||
active_keys = {(asset_type, symbol) for symbol in symbols}
|
||||
self._last_bar = {
|
||||
key: value for key, value in self._last_bar.items()
|
||||
if key[0] != asset_type or key in active_keys
|
||||
}
|
||||
required = {"symbol", "datetime", "close", "volume", "amount"}
|
||||
if not symbols or minute_df.is_empty() or not required.issubset(minute_df.columns):
|
||||
definitions = _legacy_builtin_definitions() + list(signals or [])
|
||||
if not symbols:
|
||||
return []
|
||||
|
||||
frame = build_feature_frame(
|
||||
minute_df.filter(pl.col("symbol").cast(pl.Utf8).is_in(sorted(symbols))),
|
||||
prev_close=prev_close,
|
||||
cutoff=now,
|
||||
)
|
||||
if frame.is_empty():
|
||||
return []
|
||||
|
||||
exprs = custom_signals.build_intraday_expressions(definitions)
|
||||
if not exprs:
|
||||
return []
|
||||
# 内置 4 信号保留历史列名(不带 csgi_ 前缀) — 存量监控规则零迁移
|
||||
for legacy_id in INTRADAY_SIGNAL_FIELDS:
|
||||
prefixed = custom_signals.intraday_column_name(legacy_id)
|
||||
if prefixed in exprs:
|
||||
exprs[legacy_id] = exprs.pop(prefixed)
|
||||
min_bars_by_col = {
|
||||
custom_signals.intraday_column_name(d["id"]): int(d.get("min_bars", 0) or 0)
|
||||
for d in definitions
|
||||
}
|
||||
min_bars_by_col.update({
|
||||
name: _LEGACY_MIN_BARS for name in INTRADAY_SIGNAL_FIELDS
|
||||
})
|
||||
|
||||
evaluated = custom_signals.apply_intraday_edges(frame, exprs)
|
||||
# min_bars 门槛: 当日已完成 bar 数不足时强制不触发
|
||||
evaluated = evaluated.with_columns(
|
||||
pl.int_range(pl.len()).over(["symbol", "date"]).alias("_bar_idx")
|
||||
)
|
||||
for name, min_bars in min_bars_by_col.items():
|
||||
if name in evaluated.columns and min_bars > 0:
|
||||
evaluated = evaluated.with_columns(
|
||||
pl.when(pl.col("_bar_idx") + 1 >= min_bars)
|
||||
.then(pl.col(name))
|
||||
.otherwise(False)
|
||||
.alias(name)
|
||||
)
|
||||
|
||||
cutoff = _naive_datetime(now)
|
||||
if cutoff is None:
|
||||
return []
|
||||
cutoff = cutoff.replace(second=0, microsecond=0)
|
||||
scoped = minute_df.filter(pl.col("symbol").cast(pl.Utf8).is_in(sorted(symbols)))
|
||||
if scoped.is_empty():
|
||||
return []
|
||||
|
||||
results: list[dict[str, Any]] = []
|
||||
for part in scoped.partition_by("symbol", maintain_order=False):
|
||||
signal_cols = [name for name in exprs if name in evaluated.columns]
|
||||
for part in evaluated.partition_by("symbol", maintain_order=False):
|
||||
part = part.sort("datetime")
|
||||
symbol = str(part["symbol"][0])
|
||||
points: list[tuple[datetime, float, float | None]] = []
|
||||
cumulative_amount = 0.0
|
||||
cumulative_volume = 0.0
|
||||
for row in part.iter_rows(named=True):
|
||||
bar_time = _naive_datetime(row.get("datetime"))
|
||||
price = _finite(row.get("close"))
|
||||
volume = _finite(row.get("volume"))
|
||||
amount = _finite(row.get("amount"))
|
||||
if bar_time is None or bar_time.date() != cutoff.date() or bar_time >= cutoff or price is None:
|
||||
continue
|
||||
if volume is not None and volume > 0 and amount is not None and amount >= 0:
|
||||
cumulative_volume += volume
|
||||
cumulative_amount += amount
|
||||
average = (
|
||||
cumulative_amount / (cumulative_volume * 100.0)
|
||||
if cumulative_volume > 0 and cumulative_amount > 0
|
||||
else None
|
||||
)
|
||||
points.append((bar_time, price, average))
|
||||
|
||||
if not points:
|
||||
last_time = part["datetime"][-1]
|
||||
if cutoff is not None and last_time.date() != cutoff.date():
|
||||
continue
|
||||
current = points[-1]
|
||||
key = (asset_type, symbol)
|
||||
last_bar = self._last_bar.get(key)
|
||||
self._last_bar[key] = current[0]
|
||||
if last_bar is None or last_bar.date() != current[0].date() or current[0] <= last_bar:
|
||||
last_seen = self._last_bar.get(key)
|
||||
self._last_bar[key] = last_time
|
||||
# 只有出现新 bar 才可能触发; 首次见到该标的只建状态不发信号
|
||||
if last_seen is None or last_time <= last_seen or last_time.date() != last_seen.date():
|
||||
continue
|
||||
if len(points) < 2:
|
||||
continue
|
||||
|
||||
previous = points[-2]
|
||||
baseline = _finite(prev_close.get(symbol))
|
||||
avg_up = previous[2] is not None and current[2] is not None and previous[1] <= previous[2] and current[1] > current[2]
|
||||
avg_down = previous[2] is not None and current[2] is not None and previous[1] >= previous[2] and current[1] < current[2]
|
||||
zero_up = baseline is not None and baseline > 0 and previous[1] <= baseline and current[1] > baseline
|
||||
zero_down = baseline is not None and baseline > 0 and previous[1] >= baseline and current[1] < baseline
|
||||
if avg_up or avg_down or zero_up or zero_down:
|
||||
results.append({
|
||||
"symbol": symbol,
|
||||
"signal_intraday_avg_cross_up": avg_up,
|
||||
"signal_intraday_avg_cross_down": avg_down,
|
||||
"signal_intraday_zero_cross_up": zero_up,
|
||||
"signal_intraday_zero_cross_down": zero_down,
|
||||
})
|
||||
row = {name: bool(part[name][-1]) for name in signal_cols}
|
||||
if any(row.values()):
|
||||
row["symbol"] = symbol
|
||||
results.append(row)
|
||||
return results
|
||||
|
||||
@staticmethod
|
||||
def inject(df: pl.DataFrame, signals: list[dict[str, Any]]) -> pl.DataFrame:
|
||||
existing = [field for field in INTRADAY_SIGNAL_FIELDS if field in df.columns]
|
||||
"""把本分钟触发的信号以布尔列注入 enriched 快照(缺省 False)。"""
|
||||
fields = sorted(INTRADAY_SIGNAL_FIELDS | {f for s in signals for f in s if f != "symbol"})
|
||||
existing = [field for field in fields if field in df.columns]
|
||||
out = df.drop(existing) if existing else df
|
||||
if signals:
|
||||
out = out.join(pl.DataFrame(signals), on="symbol", how="left")
|
||||
else:
|
||||
out = out.with_columns([
|
||||
pl.lit(False).alias(field) for field in INTRADAY_SIGNAL_FIELDS
|
||||
])
|
||||
return out.with_columns([
|
||||
pl.col(field).fill_null(False).cast(pl.Boolean).alias(field)
|
||||
for field in INTRADAY_SIGNAL_FIELDS
|
||||
cols = sorted({f for s in signals for f in s if f != "symbol"})
|
||||
out = out.join(pl.DataFrame(signals).select(["symbol", *cols]), on="symbol", how="left")
|
||||
out = out.with_columns([
|
||||
(
|
||||
pl.col(field).fill_null(False).cast(pl.Boolean).alias(field)
|
||||
if field in out.columns
|
||||
else pl.lit(False, dtype=pl.Boolean).alias(field)
|
||||
)
|
||||
for field in fields
|
||||
])
|
||||
return out
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
"""批次登记域 — 薄"批次"页 (持仓提醒): 只生成监控规则, 不做会计。
|
||||
|
||||
每行一个买入批次 → 派生两条规则: lot_{id}_p (price 止盈止损) / lot_{id}_d (date 到期提醒)。
|
||||
记账/加减仓属"交易口径", 不在本模块 (issue #230)。纯函数 + 文件存储, 镜像 monitor_rules.py,
|
||||
不做 API、不做引擎重载。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
from datetime import date as _date
|
||||
from pathlib import Path
|
||||
|
||||
from app.services.fs_utils import atomic_write_text
|
||||
from app.strategy import monitor_rules
|
||||
from app.strategy.monitor import MonitorRuleEngine # 复用条件文本拼装 (静态方法)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# id 需满足规则 id 同款正则, 且为后缀留位: 派生规则 {id}_p/_d 不得超过 40 字符
|
||||
_ID = monitor_rules.ID_RE
|
||||
_MAX_ID_LEN = 40 - 2 # 派生规则 id 后缀 "_p" / "_d"
|
||||
|
||||
|
||||
def _dir(data_dir: Path) -> Path:
|
||||
d = data_dir / "user_data" / "lots"
|
||||
d.mkdir(parents=True, exist_ok=True)
|
||||
return d
|
||||
|
||||
|
||||
def _path(data_dir: Path, lot_id: str) -> Path:
|
||||
return _dir(data_dir) / f"{lot_id}.json"
|
||||
|
||||
|
||||
def _now_iso() -> str:
|
||||
return datetime.now(UTC).isoformat()
|
||||
|
||||
|
||||
# ── 校验与归一化 ────────────────────────────────────────
|
||||
def validate_lot(lot: dict) -> None:
|
||||
"""校验批次字段, 非法抛 ValueError (中文信息)。"""
|
||||
lot_id = lot.get("id")
|
||||
if lot_id is not None and (
|
||||
not isinstance(lot_id, str) or not _ID.match(lot_id) or len(lot_id) > _MAX_ID_LEN
|
||||
):
|
||||
raise ValueError(f"批次 id 非法 (仅小写字母数字下划线, 且需为派生规则 id 留位): {lot_id!r}")
|
||||
if not (lot.get("symbol") or "").strip():
|
||||
raise ValueError("symbol 不能为空")
|
||||
cost = lot.get("cost_price")
|
||||
if isinstance(cost, bool) or not isinstance(cost, (int, float)) or cost <= 0:
|
||||
raise ValueError("cost_price 必须是正数")
|
||||
for key, label in (("qty", "数量"), ("target_pct", "止盈%"), ("stop_pct", "止损%")):
|
||||
v = lot.get(key, 0)
|
||||
if isinstance(v, bool) or not isinstance(v, (int, float)) or v < 0:
|
||||
raise ValueError(f"{label} 不能为负数")
|
||||
lead = lot.get("lead_days", 0)
|
||||
if isinstance(lead, bool) or not isinstance(lead, int) or lead < 0:
|
||||
raise ValueError("lead_days 必须是非负整数")
|
||||
for key, label in (("buy_date", "买入日期"), ("remind_date", "到期日")):
|
||||
raw = lot.get(key)
|
||||
if raw not in (None, ""):
|
||||
try:
|
||||
_date.fromisoformat(raw)
|
||||
except ValueError:
|
||||
raise ValueError(f"{label} 必须是 YYYY-MM-DD: {raw!r}") from None
|
||||
if not (lot.get("target_pct", 0) > 0 or lot.get("stop_pct", 0) > 0 or lot.get("remind_date")):
|
||||
raise ValueError("止盈% / 止损% / 到期日 至少设置一项 (否则无监控点)")
|
||||
|
||||
|
||||
def normalize_lot(lot: dict) -> dict:
|
||||
"""补全默认字段, 返回规范化后的批次 (不校验)。"""
|
||||
d = dict(lot)
|
||||
d["symbol"] = (d.get("symbol") or "").strip()
|
||||
d.setdefault("qty", 0)
|
||||
d.setdefault("cost_price", 0)
|
||||
d.setdefault("buy_date", None)
|
||||
d.setdefault("target_pct", 0)
|
||||
d.setdefault("stop_pct", 0)
|
||||
d.setdefault("remind_date", None)
|
||||
d.setdefault("lead_days", 1)
|
||||
d.setdefault("created_at", _now_iso())
|
||||
return d
|
||||
|
||||
|
||||
# ── 持久化 ─────────────────────────────────────────────
|
||||
def load_all(data_dir: Path) -> list[dict]:
|
||||
"""读取全部批次。损坏的文件被跳过。"""
|
||||
out: list[dict] = []
|
||||
for f in sorted(_dir(data_dir).glob("lot_*.json")):
|
||||
try:
|
||||
out.append(normalize_lot(json.loads(f.read_text(encoding="utf-8"))))
|
||||
except Exception as e:
|
||||
logger.warning("lot load failed %s: %s", f.name, e)
|
||||
return out
|
||||
|
||||
|
||||
def save_one(data_dir: Path, lot: dict) -> None:
|
||||
p = _path(data_dir, lot["id"])
|
||||
p.parent.mkdir(parents=True, exist_ok=True)
|
||||
atomic_write_text(p, json.dumps(lot, ensure_ascii=False, indent=2))
|
||||
|
||||
|
||||
def delete_one(data_dir: Path, lot_id: str) -> bool:
|
||||
p = _path(data_dir, lot_id)
|
||||
if p.exists():
|
||||
p.unlink()
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
# ── 批次 → 监控规则 (纯映射) ───────────────────────────
|
||||
def lot_to_rules(lot: dict) -> tuple[dict | None, dict | None]:
|
||||
"""批次 → (price 止盈止损规则, date 到期规则); 无对应监控点时返回 None。
|
||||
|
||||
纯映射不做 I/O; 规则 id 派生自批次 id ({lot_id}_p/_d), 保证稳定可级联。
|
||||
"""
|
||||
symbol = lot["symbol"]
|
||||
lot_id = lot["id"]
|
||||
cost = float(lot["cost_price"])
|
||||
target = float(lot.get("target_pct", 0))
|
||||
stop = float(lot.get("stop_pct", 0))
|
||||
qty = float(lot.get("qty", 0) or 0)
|
||||
qty_text = f" · {qty:g}股" if qty > 0 else ""
|
||||
|
||||
conds: list[dict] = []
|
||||
if target > 0:
|
||||
conds.append({"field": "close", "op": ">=", "value": round(cost * (1 + target / 100), 4)})
|
||||
if stop > 0:
|
||||
conds.append({"field": "close", "op": "<=", "value": round(cost * (1 - stop / 100), 4)})
|
||||
price_rule = None
|
||||
if conds:
|
||||
msg = f"批次止盈止损 · 成本{cost:g}"
|
||||
if target > 0:
|
||||
msg += f" · 止盈{target:g}%"
|
||||
if stop > 0:
|
||||
msg += f" · 止损{stop:g}%"
|
||||
msg += qty_text
|
||||
cond_text = MonitorRuleEngine._format_conditions_text({"logic": "or"}, conds)
|
||||
if cond_text:
|
||||
msg += f" · {cond_text}"
|
||||
price_rule = {
|
||||
"id": f"{lot_id}_p",
|
||||
"name": f"批次止盈止损 · {symbol}",
|
||||
"type": "price",
|
||||
"asset_type": "stock",
|
||||
"scope": "symbols",
|
||||
"symbols": [symbol],
|
||||
"conditions": conds,
|
||||
"logic": "or",
|
||||
"cooldown_seconds": 86400,
|
||||
"severity": "warn",
|
||||
"message": msg,
|
||||
"enabled": True,
|
||||
"lot_id": lot_id,
|
||||
}
|
||||
|
||||
date_rule = None
|
||||
if lot.get("remind_date"):
|
||||
lead = int(lot.get("lead_days", 1))
|
||||
date_rule = {
|
||||
"id": f"{lot_id}_d",
|
||||
"name": f"批次到期 · {symbol}",
|
||||
"type": "date",
|
||||
"asset_type": "stock",
|
||||
"scope": "symbols",
|
||||
"symbols": [symbol],
|
||||
"remind_date": lot["remind_date"],
|
||||
"lead_days": lead,
|
||||
"cooldown_seconds": 86400,
|
||||
"severity": "info",
|
||||
# 提前天数由引擎 evaluate_date_rules 统一追加, 这里只放静态部分
|
||||
"message": f"批次到期提醒 · {lot['remind_date']}{qty_text}",
|
||||
"enabled": True,
|
||||
"lot_id": lot_id,
|
||||
}
|
||||
return price_rule, date_rule
|
||||
@@ -0,0 +1,130 @@
|
||||
"""策略可访问的指数/ETF 日K读取模块 — 白名单放行的只读数据入口。
|
||||
|
||||
供 Custom/AI 策略在 filter_history 内读取任意指数(及 ETF)的完整日K。
|
||||
策略通过白名单 import 本模块, 调用纯读函数; 禁止写操作或任意文件访问。
|
||||
|
||||
设计要点:
|
||||
- 模块自身是框架侧信任代码, 对策略的沙箱逃逸拦截(ai_generator._validate_safety)照旧生效。
|
||||
- repo 线程安全懒加载(首次调用才构建); DataStore() 默认 settings.data_dir, 与 main.py 同源。
|
||||
- 未知 symbol / 数据缺失 → 返回空 DataFrame(不抛), 与 repo 语义一致。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from datetime import date
|
||||
from typing import Any
|
||||
|
||||
import polars as pl
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 完整历史默认区间下界(A股数据远晚于此, 仅作"全量"占位)。
|
||||
_FULL_START = date(1990, 1, 1)
|
||||
|
||||
# ── repo 懒加载(线程安全) ─────────────────────────────
|
||||
_repo = None
|
||||
_lock = threading.Lock()
|
||||
|
||||
|
||||
def _get_repo():
|
||||
global _repo
|
||||
if _repo is None:
|
||||
with _lock:
|
||||
if _repo is None:
|
||||
from app.tickflow.repository import DataStore, KlineRepository
|
||||
_repo = KlineRepository(DataStore())
|
||||
return _repo
|
||||
|
||||
|
||||
def _set_repo(repo: Any) -> None:
|
||||
"""测试注入: 用 fake repo 替换单例。"""
|
||||
global _repo
|
||||
with _lock:
|
||||
_repo = repo
|
||||
|
||||
|
||||
def _reset_repo() -> None:
|
||||
"""测试清理: 重置单例, 下次调用重新懒加载。"""
|
||||
global _repo
|
||||
with _lock:
|
||||
_repo = None
|
||||
|
||||
|
||||
# ── 参数规范化 ─────────────────────────────────────────
|
||||
def _norm_date(value, default: date) -> date:
|
||||
if value is None:
|
||||
return default
|
||||
if isinstance(value, str):
|
||||
return date.fromisoformat(value)
|
||||
return value
|
||||
|
||||
|
||||
def _validate_symbol(symbol: Any) -> bool:
|
||||
return isinstance(symbol, str) and bool(symbol.strip())
|
||||
|
||||
|
||||
# ── 公开只读 API ───────────────────────────────────────
|
||||
def get_index_daily(symbol, start=None, end=None, columns=None):
|
||||
"""读取指数日K(含技术指标)。未知 symbol / 无数据返回空 DataFrame。"""
|
||||
if not _validate_symbol(symbol):
|
||||
logger.warning("market_data: 非法指数 symbol %r", symbol)
|
||||
return pl.DataFrame()
|
||||
s = _norm_date(start, _FULL_START)
|
||||
e = _norm_date(end, date.today())
|
||||
try:
|
||||
return _get_repo().get_index_daily(symbol, s, e, columns)
|
||||
except Exception as exc:
|
||||
logger.warning("market_data get_index_daily failed %s: %s", symbol, exc)
|
||||
return pl.DataFrame()
|
||||
|
||||
|
||||
def get_etf_daily(symbol, start=None, end=None, columns=None):
|
||||
"""读取 ETF 日K(含技术指标)。同 get_index_daily 语义。"""
|
||||
if not _validate_symbol(symbol):
|
||||
logger.warning("market_data: 非法 ETF symbol %r", symbol)
|
||||
return pl.DataFrame()
|
||||
s = _norm_date(start, _FULL_START)
|
||||
e = _norm_date(end, date.today())
|
||||
try:
|
||||
return _get_repo().get_etf_daily(symbol, s, e, columns)
|
||||
except Exception as exc:
|
||||
logger.warning("market_data get_etf_daily failed %s: %s", symbol, exc)
|
||||
return pl.DataFrame()
|
||||
|
||||
|
||||
def get_daily(symbol, start=None, end=None, columns=None):
|
||||
"""按资产类型自动分派读取日K: 指数 → get_index_daily; ETF → get_etf_daily; 股票 → get_daily。"""
|
||||
if not _validate_symbol(symbol):
|
||||
logger.warning("market_data: 非法 symbol %r", symbol)
|
||||
return pl.DataFrame()
|
||||
s = _norm_date(start, _FULL_START)
|
||||
e = _norm_date(end, date.today())
|
||||
repo = _get_repo()
|
||||
try:
|
||||
asset_type = repo.resolve_asset_type(symbol)
|
||||
if asset_type == "index":
|
||||
return repo.get_index_daily(symbol, s, e, columns)
|
||||
if asset_type == "etf":
|
||||
return repo.get_etf_daily(symbol, s, e, columns)
|
||||
return repo.get_daily(symbol, s, e, columns)
|
||||
except Exception as exc:
|
||||
logger.warning("market_data get_daily failed %s: %s", symbol, exc)
|
||||
return pl.DataFrame()
|
||||
|
||||
|
||||
def list_index_symbols() -> list[dict]:
|
||||
"""列出已收录的指数符号(含名称)。无数据返回空列表。"""
|
||||
try:
|
||||
df = _get_repo().get_instruments_asset("index")
|
||||
except Exception as exc:
|
||||
logger.warning("market_data list_index_symbols failed: %s", exc)
|
||||
return []
|
||||
if df.is_empty() or "symbol" not in df.columns:
|
||||
return []
|
||||
name_col = "name" if "name" in df.columns else None
|
||||
cols = ["symbol"] + ([name_col] if name_col else [])
|
||||
return [
|
||||
{"symbol": row["symbol"], "name": row.get("name")}
|
||||
for row in df.select(cols).iter_rows(named=True)
|
||||
]
|
||||
@@ -25,6 +25,7 @@ from app.market_time import cn_today
|
||||
from app.strategy import config as _strategy_config
|
||||
from app.strategy.custom_signals import _OP_BUILDERS # type: ignore # 复用运算符构造器
|
||||
from app.strategy.intraday_signals import INTRADAY_SIGNAL_LABELS, uses_intraday_signals
|
||||
from app.strategy.monitor_rules import date_rule_in_window
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -71,6 +72,17 @@ def _signal_cn_name(name: str) -> str:
|
||||
return _SIGNAL_CN.get(name, name)
|
||||
|
||||
|
||||
def format_alert_quote(price, change_pct) -> str:
|
||||
"""告警正文尾部: '现价 1650.0 · +10.0%'。price/pct 均可缺; pct 为小数制。"""
|
||||
parts = []
|
||||
if price is not None:
|
||||
parts.append(f"现价 {price}")
|
||||
if change_pct is not None:
|
||||
sign = "+" if change_pct >= 0 else ""
|
||||
parts.append(f"{sign}{change_pct * 100:.1f}%")
|
||||
return " · ".join(parts)
|
||||
|
||||
|
||||
@dataclass
|
||||
class StrategyAlert:
|
||||
"""策略告警"""
|
||||
@@ -323,6 +335,10 @@ class MonitorRuleEngine:
|
||||
self._rules: dict[str, dict] = {} # rule_id → rule
|
||||
# (rule_id, symbol, event_type) → 上次触发时间戳(秒)。用于 cooldown 去重。
|
||||
self._last_fire: dict[tuple[str, str, str], float] = {}
|
||||
# date 规则每个交易日只在首个轮询评估一次; 规则集变更时失效重评
|
||||
self._date_eval_day: str | None = None
|
||||
self._date_eval_rules_version = -1
|
||||
self._rules_version = 0 # set/add/remove/clear 递增, 供 date 缓存失效
|
||||
self._strategy_engine = None # 延迟注入, type=strategy 规则用它跑选股
|
||||
# symbol → 股票名 (enriched DataFrame 已 drop name 列, 触发时从此映射回填)
|
||||
self._name_map: dict[str, str] = {}
|
||||
@@ -428,6 +444,8 @@ class MonitorRuleEngine:
|
||||
rule.get("threshold_pct"),
|
||||
rule.get("window_minutes"),
|
||||
rule.get("abnormal_window"),
|
||||
rule.get("remind_date"),
|
||||
rule.get("lead_days"),
|
||||
)
|
||||
|
||||
def set_rules(self, rules: list[dict]) -> None:
|
||||
@@ -476,12 +494,14 @@ class MonitorRuleEngine:
|
||||
if key[0] in active_ids
|
||||
}
|
||||
logger.info("MonitorRuleEngine: 装载 %d 条规则", len(self._rules))
|
||||
self._rules_version += 1
|
||||
|
||||
def add_rule(self, rule: dict) -> None:
|
||||
if rule.get("enabled") is not False:
|
||||
self._rules[rule["id"]] = rule
|
||||
else:
|
||||
self._rules.pop(rule["id"], None)
|
||||
self._rules_version += 1
|
||||
|
||||
def remove_rule(self, rule_id: str) -> None:
|
||||
self._rules.pop(rule_id, None)
|
||||
@@ -498,6 +518,7 @@ class MonitorRuleEngine:
|
||||
self._sector_condition_state = {
|
||||
k: v for k, v in self._sector_condition_state.items() if k[0] != rule_id
|
||||
}
|
||||
self._rules_version += 1
|
||||
|
||||
def clear(self) -> None:
|
||||
self._rules.clear()
|
||||
@@ -506,6 +527,7 @@ class MonitorRuleEngine:
|
||||
self._strategy_signal_state.clear()
|
||||
self._strategy_signal_seen.clear()
|
||||
self._sector_condition_state.clear()
|
||||
self._rules_version += 1
|
||||
|
||||
@property
|
||||
def rules(self) -> dict[str, dict]:
|
||||
@@ -669,7 +691,9 @@ class MonitorRuleEngine:
|
||||
for rule_id, rule in list(self._rules.items()):
|
||||
if rule.get("asset_type", "stock") != asset_type:
|
||||
continue
|
||||
if rule.get("type") in ("sector", "abnormal"):
|
||||
if rule.get("type") in ("sector", "abnormal", "date"):
|
||||
# 三者不走行情 DataFrame 评估, 各走 evaluate_sectors / evaluate_abnormal /
|
||||
# evaluate_date_rules 专用路径
|
||||
continue
|
||||
try:
|
||||
events.extend(self._evaluate_rule(df, rule, now))
|
||||
@@ -683,6 +707,80 @@ class MonitorRuleEngine:
|
||||
|
||||
return events
|
||||
|
||||
def evaluate_date_rules(self, now: float | None = None) -> list[dict]:
|
||||
"""纯日历评估 date 规则: 窗口命中 + 每天最多一次, 无行情条件。
|
||||
|
||||
由行情轮询在盘中调用 (quote_service._evaluate_monitors), 事件与 _evaluate_rule 同构。
|
||||
窗口按自然日; 到期落在休市/节假日时需 lead_days 覆盖 (交易日历口径待 issue 定夺)。
|
||||
每个交易日只在首个轮询完整评估一次, 其余轮次命中缓存直接跳过。
|
||||
"""
|
||||
now = now if now is not None else time.time()
|
||||
today_iso = cn_today().isoformat()
|
||||
if self._date_eval_day == today_iso and self._date_eval_rules_version == self._rules_version:
|
||||
return []
|
||||
# 跨天首轮清掉已过期日期的按天 cooldown 键, 避免 _last_fire 无限累积
|
||||
self._last_fire = {
|
||||
key: value
|
||||
for key, value in self._last_fire.items()
|
||||
if not (key[1].startswith("_date_") and key[1] != f"_date_{today_iso}")
|
||||
}
|
||||
|
||||
today_d = _dt.date.fromisoformat(today_iso)
|
||||
events: list[dict] = []
|
||||
for rule in list(self._rules.values()):
|
||||
if rule.get("type") != "date" or rule.get("enabled") is False:
|
||||
continue
|
||||
remind = rule.get("remind_date") or ""
|
||||
if not date_rule_in_window(remind, int(rule.get("lead_days", 0)), today_iso):
|
||||
continue
|
||||
# 按天隔离: 窗口内每天最多触发一次
|
||||
key = (rule["id"], f"_date_{today_iso}", "date")
|
||||
cooldown = int(rule.get("cooldown_seconds") or 86400)
|
||||
last = self._last_fire.get(key)
|
||||
if last is not None and (now - last) < cooldown:
|
||||
continue
|
||||
self._last_fire[key] = now
|
||||
|
||||
symbols = [s for s in rule.get("symbols", []) if s]
|
||||
single_symbol = symbols[0] if len(symbols) == 1 else None
|
||||
msg = rule.get("message") or f"日期提醒 · {today_iso}"
|
||||
try:
|
||||
remain = (_dt.date.fromisoformat(remind) - today_d).days
|
||||
except ValueError:
|
||||
remain = 0
|
||||
msg += " · 今日到期" if remain <= 0 else f" · {remain}天后到期"
|
||||
# 单标的由 ev.symbol 携带; 仅多标的时拼列表
|
||||
if len(symbols) > 1:
|
||||
shown = "、".join(symbols[:3]) + ("等" if len(symbols) > 3 else "")
|
||||
msg = f"{msg} · {shown}"
|
||||
|
||||
ev = {
|
||||
"ts": int(now * 1000),
|
||||
"rule_id": rule["id"],
|
||||
"rule_name": rule.get("name", ""),
|
||||
"strategy_id": None,
|
||||
"source": "date",
|
||||
"type": "date_reminder",
|
||||
"symbol": single_symbol or "",
|
||||
"name": (self._name_map.get(single_symbol) or single_symbol) if single_symbol else None,
|
||||
"message": msg,
|
||||
"price": None,
|
||||
"change_pct": None,
|
||||
"signals": [],
|
||||
"severity": rule.get("severity", "info"),
|
||||
"conditions": [],
|
||||
"logic": "and",
|
||||
}
|
||||
events.append(ev)
|
||||
if self._alert_handler:
|
||||
try:
|
||||
self._alert_handler(ev)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("alert handler failed: %s", e)
|
||||
self._date_eval_day = today_iso
|
||||
self._date_eval_rules_version = self._rules_version
|
||||
return events
|
||||
|
||||
def evaluate_sectors(
|
||||
self,
|
||||
stock_df: pl.DataFrame,
|
||||
@@ -1623,12 +1721,7 @@ class MonitorRuleEngine:
|
||||
# signal / price / market: 条件摘要 + 现价 + 涨跌幅
|
||||
# 条件摘要: 把 conditions (truth/比较) 拼成可读串, 如 "MA20金叉 且 量比>2"
|
||||
cond_text = self._format_conditions_text(rule, conditions)
|
||||
price_text = f"现价 {price}" if price is not None else ""
|
||||
pct_text = ""
|
||||
if pct is not None:
|
||||
sign = "+" if pct >= 0 else ""
|
||||
pct_text = f"{sign}{pct * 100:.1f}%"
|
||||
tail = " · ".join(s for s in (price_text, pct_text) if s)
|
||||
tail = format_alert_quote(price, pct)
|
||||
if cond_text and tail:
|
||||
return f"{cond_text} · {tail}"
|
||||
return cond_text or tail or "监控触发"
|
||||
|
||||
@@ -18,9 +18,10 @@ import json
|
||||
import logging
|
||||
import math
|
||||
import re
|
||||
from datetime import datetime, timezone
|
||||
from datetime import date, datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
|
||||
from app.services.fs_utils import atomic_write_text
|
||||
from app.strategy.custom_signals import ALLOWED_FIELDS
|
||||
from app.strategy.intraday_signals import uses_intraday_signals
|
||||
|
||||
@@ -28,7 +29,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 常量 ────────────────────────────────────────────────
|
||||
ID_RE = re.compile(r"^[a-z0-9_]{1,40}$")
|
||||
RULE_TYPES = {"strategy", "signal", "price", "market", "ladder", "sector", "abnormal", "volume_delta"}
|
||||
RULE_TYPES = {"strategy", "signal", "price", "market", "ladder", "sector", "abnormal", "volume_delta", "date"}
|
||||
SCOPES = {"symbols", "all", "sector", "watchlist_group"}
|
||||
LOGICS = {"and", "or"}
|
||||
DIRECTIONS = {"entry", "exit", "both"}
|
||||
@@ -100,7 +101,7 @@ def load_one(data_dir: Path, rule_id: str) -> dict | None:
|
||||
def save_one(data_dir: Path, rule: dict) -> None:
|
||||
p = _path(data_dir, rule["id"])
|
||||
p.parent.mkdir(parents=True, exist_ok=True)
|
||||
p.write_text(json.dumps(rule, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
atomic_write_text(p, json.dumps(rule, ensure_ascii=False, indent=2))
|
||||
|
||||
|
||||
def delete_one(data_dir: Path, rule_id: str) -> bool:
|
||||
@@ -117,6 +118,22 @@ def _is_signal_field(field: str) -> bool:
|
||||
return any(field.startswith(p) for p in _SIGNAL_PREFIXES)
|
||||
|
||||
|
||||
def date_rule_in_window(remind_date: str, lead_days: int, today: str) -> bool:
|
||||
"""提醒窗口 [remind_date - lead_days, remind_date] 是否包含 today (均 YYYY-MM-DD)。
|
||||
|
||||
只判自然日历窗口; 是否在交易时段由调用方决定。到期落在休市/节假日不会顺延,
|
||||
需 lead_days 覆盖 (交易日历口径待 issue 定夺)。非法输入一律返回 False (fail-safe)。
|
||||
"""
|
||||
try:
|
||||
remind = date.fromisoformat(remind_date)
|
||||
today_d = date.fromisoformat(today)
|
||||
lead = max(0, int(lead_days or 0))
|
||||
except (ValueError, TypeError):
|
||||
return False
|
||||
start = remind - timedelta(days=lead)
|
||||
return start <= today_d <= remind
|
||||
|
||||
|
||||
def validate(rule: dict) -> None:
|
||||
"""校验一条监控规则,非法则抛 ValueError (含中文信息)。"""
|
||||
rid = rule.get("id", "")
|
||||
@@ -236,6 +253,20 @@ def validate(rule: dict) -> None:
|
||||
raise ValueError(f"basic_filter.{key} 必须是正数字或 null")
|
||||
else:
|
||||
raise ValueError(f"basic_filter 不支持字段: {key}")
|
||||
elif rule.get("type") == "date":
|
||||
# 日期提醒: 纯日历, 锚定标的 (scope=symbols) 避免无对象的空提醒
|
||||
remind = rule.get("remind_date")
|
||||
if not isinstance(remind, str) or not remind.strip():
|
||||
raise ValueError("日期提醒规则必须指定 remind_date")
|
||||
try:
|
||||
date.fromisoformat(remind.strip())
|
||||
except ValueError:
|
||||
raise ValueError(f"remind_date 必须是 YYYY-MM-DD 日期: {remind!r}") from None
|
||||
lead = rule.get("lead_days", 0)
|
||||
if isinstance(lead, bool) or not isinstance(lead, int) or lead < 0:
|
||||
raise ValueError("lead_days 必须是非负整数 (提前提醒天数)")
|
||||
if rule.get("conditions"):
|
||||
raise ValueError("日期提醒规则不支持行情 conditions")
|
||||
else:
|
||||
# 信号/价格/市场类型: 需要 conditions
|
||||
conds = rule.get("conditions")
|
||||
@@ -349,6 +380,12 @@ def normalize(rule: dict) -> dict:
|
||||
r["scope"] = "all"
|
||||
r["symbols"] = []
|
||||
r["group_id"] = None
|
||||
# date 专属默认字段 (日期提醒): 纯日历窗口, 无行情条件, 每天至多一次
|
||||
if r.get("type") == "date":
|
||||
r["conditions"] = []
|
||||
r.setdefault("remind_date", None)
|
||||
r["lead_days"] = int(r.get("lead_days") or 0)
|
||||
r["cooldown_seconds"] = 86400
|
||||
# abnormal 专属默认字段 (异动边缘监控)
|
||||
r.setdefault("abnormal_window", "any")
|
||||
r.setdefault("logic", "and")
|
||||
@@ -357,7 +394,7 @@ def normalize(rule: dict) -> dict:
|
||||
r.setdefault("message", "")
|
||||
r.setdefault("webhook_url", "")
|
||||
r.setdefault("webhook_enabled", False)
|
||||
# webhook_channels: 命中时推送的外部渠道 (合法值 'feishu' | 'wecom')。
|
||||
# webhook_channels: 命中时推送的外部渠道。
|
||||
# 向后兼容: 老规则只有 webhook_enabled 布尔 (当时勾选即飞书+企业微信双推),
|
||||
# 这里把 webhook_enabled=True 但未带 webhook_channels 的老规则迁移为 ['feishu','wecom'],
|
||||
# 还原其当时的实际行为, 用户无感知。
|
||||
@@ -365,7 +402,9 @@ def normalize(rule: dict) -> dict:
|
||||
r["webhook_channels"] = ["feishu", "wecom"] if r.get("webhook_enabled") else []
|
||||
else:
|
||||
# 防御性过滤, 只保留合法渠道
|
||||
r["webhook_channels"] = [c for c in r["webhook_channels"] if c in ("feishu", "wecom")]
|
||||
r["webhook_channels"] = [
|
||||
c for c in r["webhook_channels"] if c in ("feishu", "wecom", "custom", "email")
|
||||
]
|
||||
r.setdefault("created_at", datetime.now(timezone.utc).isoformat())
|
||||
return r
|
||||
|
||||
|
||||
+126
-56
@@ -6,61 +6,22 @@ from typing import Any
|
||||
|
||||
import polars as pl
|
||||
|
||||
from app.factors.registry import (
|
||||
factor_dependencies as _registry_factor_dependencies,
|
||||
)
|
||||
from app.factors.registry import get_factor as _registry_get_factor
|
||||
from app.factors.registry import scoring_warmups as _registry_scoring_warmups
|
||||
from app.factors.registry import virtual_dependencies as _registry_virtual_dependencies
|
||||
|
||||
SCORING_DIRECTION_HIGH = "high"
|
||||
SCORING_DIRECTION_LOW = "low"
|
||||
SCORING_DIRECTIONS = frozenset({SCORING_DIRECTION_HIGH, SCORING_DIRECTION_LOW})
|
||||
|
||||
VIRTUAL_SCORING_DEPENDENCIES: dict[str, frozenset[str]] = {
|
||||
**{
|
||||
f"ma{period}_bias": frozenset({"close", f"ma{period}"})
|
||||
for period in (5, 10, 20, 30, 60)
|
||||
},
|
||||
**{
|
||||
f"ema{period}_bias": frozenset({"close", f"ema{period}"})
|
||||
for period in (5, 10, 20, 30, 60)
|
||||
},
|
||||
"macd_dif_pct": frozenset({"close", "macd_dif"}),
|
||||
"macd_dea_pct": frozenset({"close", "macd_dea"}),
|
||||
"macd_hist_pct": frozenset({"close", "macd_hist"}),
|
||||
"boll_position": frozenset({"close", "boll_upper", "boll_lower"}),
|
||||
"atr_pct": frozenset({"close", "atr_14"}),
|
||||
"boll_width": frozenset({"ma20", "boll_upper", "boll_lower"}),
|
||||
"vol_ratio_10d": frozenset({"volume"}),
|
||||
"vol_trend_5_10": frozenset({"vol_ma5", "vol_ma10"}),
|
||||
"turnover_ratio_5d": frozenset({"turnover_rate"}),
|
||||
"log_amount": frozenset({"amount"}),
|
||||
"amount_ratio_5d": frozenset({"amount"}),
|
||||
"gap_return": frozenset({"open", "prev_close"}),
|
||||
"intraday_return": frozenset({"open", "close"}),
|
||||
"close_position": frozenset({"high", "low", "close"}),
|
||||
"distance_to_high_60d": frozenset({"close", "high_60d"}),
|
||||
"distance_from_low_60d": frozenset({"close", "low_60d"}),
|
||||
"max_ret_20d": frozenset({"close"}),
|
||||
"ret_skew_20d": frozenset({"close"}),
|
||||
"up_days_20d": frozenset({"close"}),
|
||||
"amihud_20d": frozenset({"close", "amount"}),
|
||||
"turnover_z_60d": frozenset({"turnover_rate"}),
|
||||
"vol_price_corr_20d": frozenset({"close", "volume"}),
|
||||
"vwap_bias": frozenset({"close", "volume", "amount"}),
|
||||
"vol_trend_5_60": frozenset({"volume"}),
|
||||
"limit_up_count_20d": frozenset({"consecutive_limit_ups"}),
|
||||
"limit_up_count_60d": frozenset({"consecutive_limit_ups"}),
|
||||
}
|
||||
# P1 起依赖声明与预热窗口的单一权威来源为 app/factors/registry.py;
|
||||
# 本常量为兼容别名, 键集合与历史版本逐项一致 (见 tests/test_factor_registry.py 快照测试)。
|
||||
VIRTUAL_SCORING_DEPENDENCIES: dict[str, frozenset[str]] = dict(_registry_virtual_dependencies())
|
||||
|
||||
_ROLLING_SCORING_WARMUP: dict[str, int] = {
|
||||
"vol_ratio_10d": 11,
|
||||
"turnover_ratio_5d": 6,
|
||||
"amount_ratio_5d": 6,
|
||||
"max_ret_20d": 21,
|
||||
"ret_skew_20d": 21,
|
||||
"up_days_20d": 21,
|
||||
"amihud_20d": 21,
|
||||
"turnover_z_60d": 61,
|
||||
"vol_price_corr_20d": 21,
|
||||
"vol_trend_5_60": 60,
|
||||
"limit_up_count_20d": 21,
|
||||
"limit_up_count_60d": 61,
|
||||
}
|
||||
_ROLLING_SCORING_WARMUP: dict[str, int] = dict(_registry_scoring_warmups())
|
||||
|
||||
|
||||
def effective_scoring(
|
||||
@@ -89,27 +50,61 @@ def effective_scoring_directions(overrides: Mapping[str, Any] | None) -> dict[st
|
||||
|
||||
|
||||
def scoring_warmup_bars(scoring: Mapping[str, Any]) -> int:
|
||||
return max(
|
||||
(_ROLLING_SCORING_WARMUP.get(str(name), 1) for name, weight in scoring.items() if weight),
|
||||
default=1,
|
||||
)
|
||||
warmups: list[int] = [
|
||||
_ROLLING_SCORING_WARMUP.get(str(name), 1)
|
||||
for name, weight in scoring.items()
|
||||
if weight
|
||||
]
|
||||
# composite/custom 因子的预热来自注册表 (P3)
|
||||
for name, weight in scoring.items():
|
||||
if not weight:
|
||||
continue
|
||||
spec = _registry_get_factor(str(name))
|
||||
if spec is not None and spec.kind in ("custom", "composite"):
|
||||
warmups.append(spec.warmup_bars)
|
||||
return max(warmups, default=1)
|
||||
|
||||
|
||||
def scoring_dependencies(scoring: Mapping[str, Any]) -> set[str]:
|
||||
"""把受控虚拟评分字段展开为实际数据依赖。"""
|
||||
"""把受控虚拟评分字段展开为实际数据依赖 (含 composite/custom 递归展开)。"""
|
||||
dependencies: set[str] = set()
|
||||
for name, weight in scoring.items():
|
||||
if not weight:
|
||||
continue
|
||||
dependencies.update(VIRTUAL_SCORING_DEPENDENCIES.get(str(name), {str(name)}))
|
||||
dependencies.update(_registry_factor_dependencies([str(name)]))
|
||||
return dependencies
|
||||
|
||||
|
||||
def _composite_value_expr(available: set[str], name: str) -> pl.Expr | None:
|
||||
"""复合因子值 = Σ w_i * 截面 zscore(成员值); 成员可为已物化列或虚拟因子。"""
|
||||
spec = _registry_get_factor(name)
|
||||
if spec is None or not spec.components:
|
||||
return None
|
||||
total: pl.Expr | None = None
|
||||
for member_id, weight in spec.components:
|
||||
member_expr = (
|
||||
pl.col(member_id)
|
||||
if member_id in available
|
||||
else scoring_value_expr(available, member_id)
|
||||
)
|
||||
if member_expr is None:
|
||||
return None
|
||||
mean = member_expr.mean().over("date")
|
||||
std = member_expr.std().over("date")
|
||||
piece = pl.when(std > 0).then((member_expr - mean) / std).otherwise(None) * weight
|
||||
total = piece if total is None else total + piece
|
||||
return total
|
||||
|
||||
|
||||
def scoring_value_expr(columns: Collection[str], name: str) -> pl.Expr | None:
|
||||
"""返回评分值表达式;依赖不完整时返回 None。"""
|
||||
available = set(columns)
|
||||
if name in available:
|
||||
return pl.col(name)
|
||||
# composite 在 VIRTUAL 字典门控之前分派 (依赖经注册表递归展开) —— P3
|
||||
spec = _registry_get_factor(name)
|
||||
if spec is not None and spec.kind == "composite":
|
||||
return _composite_value_expr(available, name)
|
||||
dependencies = VIRTUAL_SCORING_DEPENDENCIES.get(name)
|
||||
if dependencies is None or not dependencies.issubset(available):
|
||||
return None
|
||||
@@ -202,6 +197,66 @@ def scoring_value_expr(columns: Collection[str], name: str) -> pl.Expr | None:
|
||||
window = 20 if name == "limit_up_count_20d" else 60
|
||||
hit = (pl.col("consecutive_limit_ups").fill_null(0) > 0).cast(pl.Float64)
|
||||
return hit.rolling_sum(window, min_samples=window).over("symbol")
|
||||
# ── 扩充批次 (2026-09-05): 全部滚动窗口默认 min_samples=窗口长 (fail-closed) ──
|
||||
if name == "log_float_mv":
|
||||
# 换手率 = 成交量/流通股本 → 股本 = volume/turnover_rate, 市值 = close x 股本
|
||||
return (
|
||||
pl.when((pl.col("turnover_rate") > 0) & (pl.col("volume") > 0))
|
||||
.then((pl.col("close") * pl.col("volume") / pl.col("turnover_rate")).log())
|
||||
.otherwise(None)
|
||||
)
|
||||
if name == "momentum_120d":
|
||||
return _relative(
|
||||
pl.col("close"),
|
||||
pl.col("close").shift(120),
|
||||
).over("symbol")
|
||||
if name == "mom_accel_20_60":
|
||||
return pl.col("momentum_20d") - pl.col("momentum_60d")
|
||||
if name == "rsi_14_delta_5d":
|
||||
return pl.col("rsi_14") - pl.col("rsi_14").shift(5).over("symbol")
|
||||
if name == "overnight_ret_20d":
|
||||
overnight = _relative(pl.col("open"), pl.col("prev_close"))
|
||||
return overnight.rolling_sum(20, min_samples=20).over("symbol")
|
||||
if name == "intraday_ret_20d":
|
||||
intraday = _relative(pl.col("close"), pl.col("open"))
|
||||
return intraday.rolling_sum(20, min_samples=20).over("symbol")
|
||||
if name == "downside_vol_20d":
|
||||
downside = (_daily_change_expr().clip(upper_bound=0.0) ** 2)
|
||||
return downside.rolling_mean(20, min_samples=20).sqrt().over("symbol")
|
||||
if name == "vol_regime_5_60":
|
||||
change = _daily_change_expr()
|
||||
fast = change.rolling_std(5, min_samples=5)
|
||||
slow = change.rolling_std(60, min_samples=60)
|
||||
return _ratio(fast, slow).over("symbol")
|
||||
if name == "amplitude_trend_20_60":
|
||||
fast = pl.col("amplitude").rolling_mean(20, min_samples=20)
|
||||
slow = pl.col("amplitude").rolling_mean(60, min_samples=60)
|
||||
return _relative(fast, slow).over("symbol")
|
||||
if name == "obv_trend_20d":
|
||||
change = _daily_change_expr()
|
||||
signed = change.sign() * pl.col("volume")
|
||||
total = signed.rolling_sum(20, min_samples=20)
|
||||
scale = pl.col("volume").rolling_mean(20, min_samples=20) * 20.0
|
||||
return _ratio(total, scale).over("symbol")
|
||||
if name == "amount_mean_20d":
|
||||
return (pl.col("amount") / 1e8).rolling_mean(20, min_samples=20).over("symbol")
|
||||
if name == "turnover_mean_20d":
|
||||
return pl.col("turnover_rate").rolling_mean(20, min_samples=20).over("symbol")
|
||||
if name == "turnover_std_20d":
|
||||
mean = pl.col("turnover_rate").rolling_mean(20, min_samples=20)
|
||||
std = pl.col("turnover_rate").rolling_std(20, min_samples=20)
|
||||
return _ratio(std, mean).over("symbol")
|
||||
if name == "position_240d":
|
||||
high = pl.col("close").rolling_max(240, min_samples=240)
|
||||
low = pl.col("close").rolling_min(240, min_samples=240)
|
||||
return _ratio(pl.col("close") - low, high - low).over("symbol")
|
||||
if name == "distance_to_high_240d":
|
||||
return _relative(
|
||||
pl.col("close"),
|
||||
pl.col("close").rolling_max(240, min_samples=240),
|
||||
).over("symbol")
|
||||
if name == "kdj_kd_diff":
|
||||
return pl.col("kdj_k") - pl.col("kdj_d")
|
||||
return None
|
||||
|
||||
|
||||
@@ -233,6 +288,21 @@ def materialize_scoring_columns(
|
||||
frame: pl.DataFrame,
|
||||
names: Collection[str],
|
||||
) -> pl.DataFrame:
|
||||
# custom (DSL) 因子先物化: frame_transform 可能需要多阶段临时列 (嵌套窗口规避),
|
||||
# 与单表达式路径不同, 必须整体走帧变换 —— 与检验/试算共用同一条计算路径 (P3)。
|
||||
from app.factors.dsl import FACTOR_COLUMN, compile_formula_cached
|
||||
|
||||
for name in names:
|
||||
spec = _registry_get_factor(str(name))
|
||||
if spec is None or spec.kind != "custom" or name in frame.columns:
|
||||
continue
|
||||
compiled = compile_formula_cached(spec.formula_text)
|
||||
if compiled.frame_transform is None:
|
||||
continue
|
||||
transformed = compiled.frame_transform(frame)
|
||||
if transformed is None:
|
||||
continue
|
||||
frame = transformed.with_columns(pl.col(FACTOR_COLUMN).alias(str(name))).drop(FACTOR_COLUMN)
|
||||
expressions = [
|
||||
expression.alias(name)
|
||||
for name in names
|
||||
|
||||
@@ -305,11 +305,12 @@ def detect_capabilities(force: bool = False) -> CapabilitySet:
|
||||
|
||||
# 数据集 → 能力映射: 第三方源声明某数据集且被选为当前 provider 时补授的能力。
|
||||
# 实时行情无对应能力键 (权限由 QuoteService.is_realtime_allowed 判定);
|
||||
# 五档盘口/WebSocket 暂无第三方数据集契约, 不增广。
|
||||
# WebSocket 暂无第三方数据集契约, 不增广。
|
||||
_DATASET_CAP_MAP: tuple[tuple[str, Cap], ...] = (
|
||||
("daily", Cap.KLINE_DAILY_BATCH),
|
||||
("adj_factor", Cap.ADJ_FACTOR),
|
||||
("minute", Cap.KLINE_MINUTE_BATCH),
|
||||
("depth5", Cap.DEPTH5_BATCH),
|
||||
("financial", Cap.FINANCIAL),
|
||||
("full_minute", Cap.INTRADAY_UNIVERSE),
|
||||
)
|
||||
@@ -327,6 +328,7 @@ def _augment_custom_sources(capset: CapabilitySet) -> None:
|
||||
"daily": daily_provider,
|
||||
"adj_factor": adj_provider,
|
||||
"minute": preferences.get_minute_data_provider(),
|
||||
"depth5": preferences.get_depth5_data_provider(),
|
||||
"financial": preferences.get_financial_provider(),
|
||||
"full_minute": preferences.get_full_minute_data_provider(),
|
||||
}
|
||||
@@ -548,8 +550,6 @@ def _compute_label_and_missing(
|
||||
|
||||
base_caps = _tier_caps_set(tiers, base)
|
||||
missing = sorted(c.value for c in (base_caps - held))
|
||||
extras = base_caps and (held - base_caps) or set() # extras 是超出该档的部分
|
||||
|
||||
# 实际超出 = held 中"既不属于本档、也不属于本档下方任何档"的 cap
|
||||
# 简化:extras = held - base_caps
|
||||
extras_set = held - base_caps
|
||||
|
||||
+146
-124
@@ -34,6 +34,7 @@ from app.enriched_generation import (
|
||||
)
|
||||
from app.market_time import cn_today
|
||||
from app.parquet import scan_enriched_parquet
|
||||
from app.polars_guard import guarded_collect
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -538,6 +539,12 @@ class KlineRepository:
|
||||
self._index_enriched_cache_date = None
|
||||
|
||||
def _refresh_enriched(self) -> None:
|
||||
from app.services.heavy_job_limiter import shared_heavy_job_limiter
|
||||
|
||||
with shared_heavy_job_limiter.slot("exclusive"):
|
||||
self._refresh_enriched_impl()
|
||||
|
||||
def _refresh_enriched_impl(self) -> None:
|
||||
"""从 parquet 加载 enriched 最新日到内存 + 构建聚合表。
|
||||
|
||||
enriched parquet 仅存 14 列基础数据。启动时读入历史数据并即时计算完整指标,
|
||||
@@ -582,7 +589,7 @@ class KlineRepository:
|
||||
# 300 日历天 ≈ 210 交易日, 覆盖 filter_history 最大 lookback(90) + warmup(60)
|
||||
try:
|
||||
from datetime import timedelta
|
||||
from app.indicators.pipeline import compute_indicators, compute_signals, compute_limit_signals
|
||||
from app.indicators.pipeline import compute_enriched_history_window
|
||||
start_full = latest - timedelta(days=300)
|
||||
read_cols = [c for c in ["symbol", "date", "open", "high", "low", "close",
|
||||
"volume", "amount", "raw_close", "raw_high", "raw_low"]
|
||||
@@ -595,47 +602,28 @@ class KlineRepository:
|
||||
|
||||
step = time.perf_counter()
|
||||
logger.info("enriched refresh step start: collect history from %s", start_full)
|
||||
df_hist = lf.select(read_cols).collect()
|
||||
df_hist = guarded_collect(lf.select(read_cols), priority="background")
|
||||
logger.info("enriched refresh step done: collect history rows=%d (%.2fs)", len(df_hist), time.perf_counter() - step)
|
||||
if not df_hist.is_empty():
|
||||
instruments = self._instruments_cache if self._instruments_cache is not None else pl.DataFrame()
|
||||
|
||||
# 分批计算并关联元数据, 保留完整历史, 限制宽表临时副本。
|
||||
step = time.perf_counter()
|
||||
logger.info("enriched refresh step start: compute indicators")
|
||||
df_full = compute_indicators(df_hist)
|
||||
logger.info("enriched refresh step done: compute indicators rows=%d (%.2fs)", len(df_full), time.perf_counter() - step)
|
||||
|
||||
# 异动偏离列 (deviate_Nd = 个股动量 - 基准指数动量), 运行时附着
|
||||
from app.indicators.pipeline import attach_deviation_columns
|
||||
df_full = attach_deviation_columns(df_full, self.store.data_dir)
|
||||
|
||||
step = time.perf_counter()
|
||||
logger.info("enriched refresh step start: compute signals")
|
||||
df_full = compute_signals(df_full)
|
||||
logger.info("enriched refresh step done: compute signals (%.2fs)", time.perf_counter() - step)
|
||||
if instruments is not None and not instruments.is_empty():
|
||||
step = time.perf_counter()
|
||||
logger.info("enriched refresh step start: compute limit signals")
|
||||
df_full = compute_limit_signals(
|
||||
df_full,
|
||||
instruments,
|
||||
historical_shares=self.get_historical_shares(),
|
||||
)
|
||||
logger.info("enriched refresh step done: compute limit signals (%.2fs)", time.perf_counter() - step)
|
||||
|
||||
# JOIN instruments 到完整历史 (filter_history/basic_filter 需要 name/股本等列)
|
||||
if instruments is not None and not instruments.is_empty():
|
||||
inst_cols = [c for c in ["name", "total_shares", "float_shares"]
|
||||
if c in instruments.columns and c not in df_full.columns]
|
||||
if inst_cols:
|
||||
step = time.perf_counter()
|
||||
logger.info("enriched refresh step start: join instruments")
|
||||
df_full = df_full.join(
|
||||
instruments.select(["symbol", *inst_cols]).unique(subset=["symbol"]),
|
||||
on="symbol",
|
||||
how="left",
|
||||
)
|
||||
logger.info("enriched refresh step done: join instruments (%.2fs)", time.perf_counter() - step)
|
||||
logger.info("enriched refresh step start: compute window (batched)")
|
||||
df_full = compute_enriched_history_window(
|
||||
df_hist,
|
||||
self.store.data_dir,
|
||||
instruments=instruments,
|
||||
historical_shares=(
|
||||
self.get_historical_shares()
|
||||
if instruments is not None and not instruments.is_empty()
|
||||
else None
|
||||
),
|
||||
include_instrument_metadata=True,
|
||||
)
|
||||
del df_hist
|
||||
logger.info("enriched refresh step done: compute window rows=%d (%.2fs)",
|
||||
len(df_full), time.perf_counter() - step)
|
||||
|
||||
# 缓存完整历史 (含指标+必要基础信息) 供 filter_history/backtest 直接复用
|
||||
if self.get_matrix_data_generation("stock") != refresh_generation:
|
||||
@@ -777,9 +765,9 @@ class KlineRepository:
|
||||
needed = [c for c in base_cols if c in hist_all.columns]
|
||||
step = time.perf_counter()
|
||||
logger.info("live agg step start: slice history cache")
|
||||
df_hist = hist_all.filter(
|
||||
df_hist = hist_all.select(needed).filter(
|
||||
(pl.col("date") >= start_60d) & (pl.col("date") <= latest)
|
||||
).select(needed).sort(["symbol", "date"])
|
||||
).sort(["symbol", "date"])
|
||||
logger.info("live agg step done: slice history cache rows=%d (%.2fs)", len(df_hist), time.perf_counter() - step)
|
||||
|
||||
state_cols = [
|
||||
@@ -899,7 +887,7 @@ class KlineRepository:
|
||||
c for c in ["symbol", "consecutive_limit_ups", "consecutive_limit_downs"]
|
||||
if c in lf.collect_schema().names()
|
||||
]
|
||||
consec_source = lf.select("date", *consec_cols).collect()
|
||||
consec_source = guarded_collect(lf.select("date", *consec_cols), priority="background")
|
||||
if len(consec_cols) == 3:
|
||||
consec_df = _last_available_rows(
|
||||
consec_source.select("date", *consec_cols), latest,
|
||||
@@ -988,7 +976,7 @@ class KlineRepository:
|
||||
"raw_close", "raw_high", "raw_low",
|
||||
"consecutive_limit_ups", "consecutive_limit_downs"]
|
||||
if c in lf.collect_schema().names()]
|
||||
df_hist = lf.select(read_cols).collect()
|
||||
df_hist = guarded_collect(lf.select(read_cols), priority="background")
|
||||
|
||||
if df_hist.is_empty():
|
||||
return df_hist, pl.DataFrame()
|
||||
@@ -1035,13 +1023,13 @@ class KlineRepository:
|
||||
read_cols = [c for c in ["symbol", "date", "open", "high", "low", "close",
|
||||
"volume", "amount", "raw_close", "raw_high", "raw_low"]
|
||||
if c in df_latest.columns]
|
||||
df_hist = (
|
||||
df_hist = guarded_collect(
|
||||
scan_enriched_parquet(self._etf_enriched_glob,
|
||||
cast_options=pl.ScanCastOptions(integer_cast="allow-float"))
|
||||
.filter(pl.col("date") >= start_full)
|
||||
.select(read_cols)
|
||||
.sort(["symbol", "date"])
|
||||
.collect()
|
||||
.sort(["symbol", "date"]),
|
||||
priority="background",
|
||||
)
|
||||
if df_hist.is_empty():
|
||||
self._etf_enriched_cache = df_latest.sort(["symbol"])
|
||||
@@ -1079,13 +1067,13 @@ class KlineRepository:
|
||||
read_cols = [c for c in ["symbol", "date", "open", "high", "low", "close",
|
||||
"volume", "amount"]
|
||||
if c in df_latest.columns]
|
||||
df_hist = (
|
||||
df_hist = guarded_collect(
|
||||
scan_enriched_parquet(self._index_enriched_glob,
|
||||
cast_options=pl.ScanCastOptions(integer_cast="allow-float"))
|
||||
.filter(pl.col("date") >= start_full)
|
||||
.select(read_cols)
|
||||
.sort(["symbol", "date"])
|
||||
.collect()
|
||||
.sort(["symbol", "date"]),
|
||||
priority="background",
|
||||
)
|
||||
if df_hist.is_empty():
|
||||
self._index_enriched_cache = df_latest.sort(["symbol"])
|
||||
@@ -1099,7 +1087,7 @@ class KlineRepository:
|
||||
def _refresh_instruments(self) -> None:
|
||||
"""加载 instruments 到内存。"""
|
||||
try:
|
||||
df = pl.scan_parquet(self._inst_glob).collect()
|
||||
df = guarded_collect(pl.scan_parquet(self._inst_glob), priority="background")
|
||||
if not df.is_empty():
|
||||
self._instruments_cache = df
|
||||
self._name_map_cache = None
|
||||
@@ -1110,7 +1098,7 @@ class KlineRepository:
|
||||
def _refresh_index_instruments(self) -> None:
|
||||
"""加载指数 instruments 到内存。"""
|
||||
try:
|
||||
df = pl.scan_parquet(self._index_inst_glob).collect()
|
||||
df = guarded_collect(pl.scan_parquet(self._index_inst_glob), priority="background")
|
||||
if not df.is_empty():
|
||||
self._index_instruments_cache = df
|
||||
self._index_symbol_set_cache = None
|
||||
@@ -1123,7 +1111,7 @@ class KlineRepository:
|
||||
"""加载 ETF instruments 到内存;兼容旧版 instruments_index 中的 ETF。"""
|
||||
parts: list[pl.DataFrame] = []
|
||||
try:
|
||||
df = pl.scan_parquet(self._etf_inst_glob).collect()
|
||||
df = guarded_collect(pl.scan_parquet(self._etf_inst_glob), priority="background")
|
||||
if not df.is_empty():
|
||||
parts.append(df)
|
||||
except Exception as e: # noqa: BLE001
|
||||
@@ -1241,16 +1229,20 @@ class KlineRepository:
|
||||
if cache_min > start or cache_max < end:
|
||||
return None
|
||||
|
||||
df = cache.filter((pl.col("date") >= start) & (pl.col("date") <= end))
|
||||
if symbols is not None:
|
||||
df = df.filter(pl.col("symbol").is_in(symbols))
|
||||
if columns and not df.is_empty():
|
||||
df = cache
|
||||
if columns:
|
||||
existing = [c for c in columns if c in df.columns]
|
||||
if "symbol" not in existing and "symbol" in df.columns:
|
||||
existing.insert(0, "symbol")
|
||||
if "date" not in existing and "date" in df.columns:
|
||||
existing.insert(1, "date")
|
||||
df = df.select(existing)
|
||||
df = df.select(list(dict.fromkeys(existing)))
|
||||
df = df.filter((pl.col("date") >= start) & (pl.col("date") <= end))
|
||||
if symbols is not None:
|
||||
df = df.filter(pl.col("symbol").is_in(symbols))
|
||||
if columns:
|
||||
# 保持旧接口空结果的完整 schema, 非空时沿用请求列校验。
|
||||
df = cache.clear() if df.is_empty() else df.select(existing)
|
||||
return df.sort(["symbol", "date"])
|
||||
|
||||
def get_live_agg(self) -> pl.DataFrame:
|
||||
@@ -1577,10 +1569,12 @@ class KlineRepository:
|
||||
) -> pl.DataFrame:
|
||||
"""分钟K查询 — Polars scan_parquet + predicate pushdown。"""
|
||||
try:
|
||||
return pl.scan_parquet(self._minute_glob_for(asset_type)).filter(
|
||||
(pl.col("symbol") == symbol)
|
||||
& (pl.col("datetime").dt.date() == trade_date)
|
||||
).sort("datetime").collect()
|
||||
return guarded_collect(
|
||||
pl.scan_parquet(self._minute_glob_for(asset_type)).filter(
|
||||
(pl.col("symbol") == symbol)
|
||||
& (pl.col("datetime").dt.date() == trade_date)
|
||||
).sort("datetime")
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("分钟K查询失败: %s", e)
|
||||
return pl.DataFrame()
|
||||
@@ -1599,10 +1593,12 @@ class KlineRepository:
|
||||
if not symbols:
|
||||
return pl.DataFrame()
|
||||
try:
|
||||
return pl.scan_parquet(self._minute_glob_for(asset_type)).filter(
|
||||
pl.col("symbol").is_in(symbols)
|
||||
& (pl.col("datetime").dt.date() == trade_date)
|
||||
).sort(["symbol", "datetime"]).collect()
|
||||
return guarded_collect(
|
||||
pl.scan_parquet(self._minute_glob_for(asset_type)).filter(
|
||||
pl.col("symbol").is_in(symbols)
|
||||
& (pl.col("datetime").dt.date() == trade_date)
|
||||
).sort(["symbol", "datetime"])
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("批量分钟K查询失败: %s", e)
|
||||
return pl.DataFrame()
|
||||
@@ -1625,15 +1621,15 @@ class KlineRepository:
|
||||
lf = pl.scan_parquet(self._minute_glob_for(asset_type))
|
||||
available = set(lf.collect_schema().names())
|
||||
select_cols = [c for c in ["symbol", "datetime", "open", "high", "low", "close", "volume", "amount"] if c in available]
|
||||
return (
|
||||
return guarded_collect(
|
||||
lf.select(select_cols)
|
||||
.filter(
|
||||
pl.col("symbol").is_in(symbols)
|
||||
& (pl.col("datetime").dt.date() >= start)
|
||||
& (pl.col("datetime").dt.date() <= end)
|
||||
)
|
||||
.sort(["symbol", "datetime"])
|
||||
.collect(streaming=True)
|
||||
.sort(["symbol", "datetime"]),
|
||||
streaming=True,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("分钟K范围查询失败: %s", e)
|
||||
@@ -1670,11 +1666,11 @@ class KlineRepository:
|
||||
lf = pl.scan_parquet(parts)
|
||||
available = set(lf.collect_schema().names())
|
||||
select_cols = [c for c in ["symbol", "datetime", "open", "high", "low", "close", "volume", "amount"] if c in available]
|
||||
return (
|
||||
return guarded_collect(
|
||||
lf.select(select_cols)
|
||||
.filter(pl.col("symbol").is_in(symbols))
|
||||
.sort(["symbol", "datetime"])
|
||||
.collect(streaming=True)
|
||||
.sort(["symbol", "datetime"]),
|
||||
streaming=True,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("分钟K按日期查询失败: %s", e)
|
||||
@@ -1744,7 +1740,7 @@ class KlineRepository:
|
||||
schema_names = lf.collect_schema().names()
|
||||
existing = [c for c in columns if c in schema_names]
|
||||
lf = lf.select(existing)
|
||||
return lf.collect()
|
||||
return guarded_collect(lf)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("日K查询失败: %s", e)
|
||||
return pl.DataFrame()
|
||||
@@ -1761,7 +1757,7 @@ class KlineRepository:
|
||||
schema_names = lf.collect_schema().names()
|
||||
existing = [c for c in columns if c in schema_names]
|
||||
lf = lf.select(existing)
|
||||
return lf.collect()
|
||||
return guarded_collect(lf)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("日K批量查询失败: %s", e)
|
||||
return pl.DataFrame()
|
||||
@@ -1778,7 +1774,7 @@ class KlineRepository:
|
||||
schema_names = lf.collect_schema().names()
|
||||
existing = [c for c in columns if c in schema_names]
|
||||
lf = lf.select(existing)
|
||||
return lf.collect()
|
||||
return guarded_collect(lf)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("指数日K查询失败: %s", e)
|
||||
return pl.DataFrame()
|
||||
@@ -1795,7 +1791,7 @@ class KlineRepository:
|
||||
schema_names = lf.collect_schema().names()
|
||||
existing = [c for c in columns if c in schema_names]
|
||||
lf = lf.select(existing)
|
||||
return lf.collect()
|
||||
return guarded_collect(lf)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.debug("ETF 日K查询跳过: %s", e)
|
||||
return pl.DataFrame()
|
||||
@@ -2162,27 +2158,76 @@ class KlineRepository:
|
||||
if generation_asset is not None
|
||||
else None
|
||||
)
|
||||
for date_df in df.partition_by("date"):
|
||||
dt = date_df["date"][0]
|
||||
ds = dt.isoformat() if hasattr(dt, "isoformat") else str(dt)
|
||||
out = base / f"date={ds}" / "part.parquet"
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._optimistic_upsert_partition(out, date_df, publication)
|
||||
|
||||
@staticmethod
|
||||
def _partition_fingerprint(path: Path) -> tuple[int, int] | None:
|
||||
"""分区文件的修改指纹 (mtime_ns, size); 不存在返回 None。"""
|
||||
try:
|
||||
st = path.stat()
|
||||
except FileNotFoundError:
|
||||
return None
|
||||
return (st.st_mtime_ns, st.st_size)
|
||||
|
||||
def _optimistic_upsert_partition(
|
||||
self,
|
||||
out: Path,
|
||||
incoming: pl.DataFrame,
|
||||
publication: EnrichedPublication | None,
|
||||
*,
|
||||
retries: int = 3,
|
||||
) -> None:
|
||||
"""单分区 merge-upsert: polars 读/合并/排序在 _write_lock 外, 锁内只做
|
||||
指纹校验 + 原子替换 + commit。
|
||||
|
||||
背景: polars 并发执行存在死锁风险 (见 app.polars_guard), 重活若在
|
||||
_write_lock 内悬死, 全局写锁被永久持有, 所有写路径排队冻结。乐观模式
|
||||
把读/算移出锁外; 锁内用指纹确认基底未被其他写入者改动, 失配则重试,
|
||||
重试耗尽退回锁内直读直写 (正确性优先, 牺牲隔离性)。
|
||||
"""
|
||||
def _merge(existing: pl.DataFrame) -> pl.DataFrame:
|
||||
if existing.is_empty():
|
||||
return incoming.sort(["symbol", "date"])
|
||||
return pl.concat([existing, incoming], how="diagonal_relaxed").unique(
|
||||
subset=["symbol", "date"], keep="last"
|
||||
).sort(["symbol", "date"])
|
||||
|
||||
for _ in range(retries):
|
||||
existing = pl.read_parquet(out) if out.exists() else pl.DataFrame()
|
||||
base_fp = self._partition_fingerprint(out)
|
||||
merged = _merge(existing)
|
||||
with self._write_lock:
|
||||
if self._partition_fingerprint(out) != base_fp:
|
||||
continue # 基底被并发写入者改过, 出锁重读重算
|
||||
self._write_partition_locked(out, merged, existing, publication)
|
||||
return
|
||||
# 乐观重试耗尽 (罕见: 高频并发写同一分区): 退回锁内全量模式保证正确性
|
||||
with self._write_lock:
|
||||
for date_df in df.partition_by("date"):
|
||||
dt = date_df["date"][0]
|
||||
ds = dt.isoformat() if hasattr(dt, "isoformat") else str(dt)
|
||||
out = base / f"date={ds}" / "part.parquet"
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
existing = pl.DataFrame()
|
||||
if out.exists():
|
||||
existing = pl.read_parquet(out)
|
||||
date_df = pl.concat([existing, date_df], how="diagonal_relaxed").unique(
|
||||
subset=["symbol", "date"], keep="last"
|
||||
)
|
||||
date_df = date_df.sort(["symbol", "date"])
|
||||
if not existing.is_empty() and existing.equals(date_df):
|
||||
continue
|
||||
if publication is None:
|
||||
self._atomic_write_parquet(date_df, out)
|
||||
else:
|
||||
publication.write_parquet(date_df, out)
|
||||
existing = pl.read_parquet(out) if out.exists() else pl.DataFrame()
|
||||
self._write_partition_locked(out, _merge(existing), existing, publication)
|
||||
|
||||
def _write_partition_locked(
|
||||
self,
|
||||
out: Path,
|
||||
merged: pl.DataFrame,
|
||||
existing: pl.DataFrame,
|
||||
publication: EnrichedPublication | None,
|
||||
) -> None:
|
||||
"""锁内的纯文件阶段: 无变化跳过; 否则原子替换 + 提交 generation。"""
|
||||
if not existing.is_empty() and existing.equals(merged):
|
||||
if publication is not None:
|
||||
publication.commit()
|
||||
publication.commit() # 未写入时为无害空提交
|
||||
return
|
||||
if publication is None:
|
||||
self._atomic_write_parquet(merged, out)
|
||||
else:
|
||||
publication.write_parquet(merged, out)
|
||||
publication.commit()
|
||||
|
||||
def merge_live_daily_asset(self, asset_type: str, df: pl.DataFrame) -> None:
|
||||
"""按 symbol 合并当天指定资产日K分区。用于少量自选实时,不覆盖全市场。"""
|
||||
@@ -2200,14 +2245,7 @@ class KlineRepository:
|
||||
ds = dt.isoformat() if hasattr(dt, "isoformat") else str(dt)
|
||||
out = base / f"date={ds}" / "part.parquet"
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
with self._write_lock:
|
||||
date_df = df.sort(["symbol", "date"])
|
||||
if out.exists():
|
||||
existing = pl.read_parquet(out)
|
||||
date_df = pl.concat([existing, date_df], how="diagonal_relaxed").unique(
|
||||
subset=["symbol", "date"], keep="last"
|
||||
)
|
||||
self._atomic_write_parquet(date_df.sort(["symbol", "date"]), out)
|
||||
self._optimistic_upsert_partition(out, df, None)
|
||||
|
||||
def _with_instrument_metadata(self, asset_type: str, df: pl.DataFrame) -> pl.DataFrame:
|
||||
"""补齐实时内存缓存所需的维表字段;这些字段不会写入 enriched 分区。"""
|
||||
@@ -2264,21 +2302,7 @@ class KlineRepository:
|
||||
if asset_type in {"stock", "etf"}
|
||||
else None
|
||||
)
|
||||
with self._write_lock:
|
||||
existing = pl.DataFrame()
|
||||
if out.exists():
|
||||
existing = pl.read_parquet(out)
|
||||
df_storage = pl.concat([existing, df_storage], how="diagonal_relaxed").unique(
|
||||
subset=["symbol", "date"], keep="last"
|
||||
)
|
||||
df_storage = df_storage.sort(["symbol"])
|
||||
if existing.is_empty() or not existing.equals(df_storage):
|
||||
if publication is None:
|
||||
self._atomic_write_parquet(df_storage, out)
|
||||
else:
|
||||
publication.write_parquet(df_storage, out)
|
||||
if publication is not None:
|
||||
publication.commit()
|
||||
self._optimistic_upsert_partition(out, df_storage, publication)
|
||||
|
||||
if asset_type == "stock":
|
||||
self._enriched_cache = merged_cache
|
||||
@@ -2312,8 +2336,10 @@ class KlineRepository:
|
||||
ds = dt.isoformat() if hasattr(dt, "isoformat") else str(dt)
|
||||
out = base / f"date={ds}" / "part.parquet"
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
# 覆写语义: 排序在锁外, 锁内只做原子替换。
|
||||
df_sorted = df.sort(["symbol", "date"])
|
||||
with self._write_lock:
|
||||
self._atomic_write_parquet(df.sort(["symbol", "date"]), out)
|
||||
self._atomic_write_parquet(df_sorted, out)
|
||||
|
||||
def flush_live_enriched(self, df: pl.DataFrame) -> None:
|
||||
"""覆写当天 kline_daily_enriched 分区 (实时 enriched 落盘, 非merge)。
|
||||
@@ -2349,15 +2375,11 @@ class KlineRepository:
|
||||
if asset_type in {"stock", "etf"}
|
||||
else None
|
||||
)
|
||||
# 覆写语义: 读旧内容只为跳过无变化的写, 读在锁外 (误判最多造成一次
|
||||
# 冗余覆写, 不影响正确性); 锁内只做替换 + commit。
|
||||
existing = pl.read_parquet(out) if out.exists() else pl.DataFrame()
|
||||
with self._write_lock:
|
||||
existing = pl.read_parquet(out) if out.exists() else pl.DataFrame()
|
||||
if existing.is_empty() or not existing.equals(df_storage):
|
||||
if publication is None:
|
||||
self._atomic_write_parquet(df_storage, out)
|
||||
else:
|
||||
publication.write_parquet(df_storage, out)
|
||||
if publication is not None:
|
||||
publication.commit()
|
||||
self._write_partition_locked(out, df_storage, existing, publication)
|
||||
|
||||
if asset_type == "stock":
|
||||
self._enriched_cache = cache_df
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
"""后端自愈看门狗。
|
||||
|
||||
2026-09-07 事故形态: polars 并发死锁把线程悬死在 collect 内部 (0 CPU 永久
|
||||
挂起), 其中持锁者让 _write_lock 永久被占, 所有请求线程排队冻结, 只能人工
|
||||
重启。并发闸 (app.polars_guard) 与写锁瘦身 (repository 乐观并发) 分别削减
|
||||
触发概率与扩散半径; 看门狗是最后一层兜底 —— 探测走与事故相同的共享资源
|
||||
路径 (collect 闸 + 全局写锁), 连续 N 次超时即判定进程已僵死, 主动退出交由
|
||||
supervisor / Docker restart / dev 脚本拉起, 把恢复时间从"人工发现"缩短到
|
||||
约一分钟。
|
||||
|
||||
误伤防护: 探测本身是毫秒级微型 collect + 1s 写锁试探, 阈值要求连续失败
|
||||
(默认 2 次 × 15s 超时), 高负载下"慢而未死"不会触发。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
|
||||
import polars as pl
|
||||
|
||||
from app.config import settings
|
||||
from app.polars_guard import guarded_collect
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def default_probe(write_lock: threading.Lock | None = None) -> None:
|
||||
"""探测关键共享资源: polars collect 闸 + 仓库全局写锁。
|
||||
|
||||
任一被悬死线程占住即超时 — 正是 2026-09-07 冻结事故中被毒化的两条路径。
|
||||
"""
|
||||
guarded_collect(pl.LazyFrame({"probe": [1]}).sum())
|
||||
if write_lock is not None:
|
||||
acquired = write_lock.acquire(timeout=1.0)
|
||||
if not acquired:
|
||||
raise TimeoutError("repository write lock unavailable")
|
||||
write_lock.release()
|
||||
|
||||
|
||||
class HealthWatchdog:
|
||||
"""周期探测; 连续 failure_threshold 次失败后调用 exit_cb(退出码)。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
probe: Callable[[], None],
|
||||
*,
|
||||
exit_cb: Callable[[int], None],
|
||||
interval_s: float | None = None,
|
||||
probe_timeout_s: float | None = None,
|
||||
failure_threshold: int | None = None,
|
||||
) -> None:
|
||||
self._probe = probe
|
||||
self._exit_cb = exit_cb
|
||||
self._interval_s = settings.watchdog_interval_s if interval_s is None else interval_s
|
||||
self._probe_timeout_s = (
|
||||
settings.watchdog_probe_timeout_s if probe_timeout_s is None else probe_timeout_s
|
||||
)
|
||||
self._failure_threshold = (
|
||||
settings.watchdog_failure_threshold if failure_threshold is None else failure_threshold
|
||||
)
|
||||
self._consecutive_failures = 0
|
||||
self._task: asyncio.Task | None = None
|
||||
|
||||
async def _loop(self) -> None:
|
||||
while True:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
asyncio.to_thread(self._probe), timeout=self._probe_timeout_s
|
||||
)
|
||||
self._consecutive_failures = 0
|
||||
except BaseException as exc: # 探测任何异常都算失败 (含 to_thread 超时)
|
||||
self._consecutive_failures += 1
|
||||
logger.error(
|
||||
"watchdog probe failed (%d/%d): %r",
|
||||
self._consecutive_failures,
|
||||
self._failure_threshold,
|
||||
exc,
|
||||
)
|
||||
if self._consecutive_failures >= self._failure_threshold:
|
||||
logger.critical(
|
||||
"watchdog: backend wedged (probe failed %d consecutive times); "
|
||||
"exiting for supervisor restart",
|
||||
self._consecutive_failures,
|
||||
)
|
||||
self._exit_cb(70)
|
||||
return
|
||||
await asyncio.sleep(self._interval_s)
|
||||
|
||||
def start(self) -> None:
|
||||
if self._task is None or self._task.done():
|
||||
self._task = asyncio.create_task(self._loop(), name="health-watchdog")
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._task is not None:
|
||||
self._task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await self._task
|
||||
self._task = None
|
||||
|
||||
|
||||
def start_watchdog(app_state, repo) -> HealthWatchdog | None:
|
||||
"""lifespan 启动钩子; 返回实例挂到 app.state.watchdog 便于关闭。"""
|
||||
if not settings.watchdog_enabled:
|
||||
return None
|
||||
write_lock = getattr(repo, "_write_lock", None)
|
||||
watchdog = HealthWatchdog(
|
||||
lambda: default_probe(write_lock),
|
||||
exit_cb=lambda code: os._exit(code),
|
||||
)
|
||||
watchdog.start()
|
||||
return watchdog
|
||||
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "tickflow-stock-panel-backend"
|
||||
version = "0.2.2"
|
||||
version = "0.2.3"
|
||||
description = "A 股选股 + 监控 + 回测面板 — TickFlow 适配"
|
||||
requires-python = ">=3.11"
|
||||
license = { text = "MIT" }
|
||||
@@ -13,7 +13,9 @@ dependencies = [
|
||||
"python-multipart>=0.0.6",
|
||||
"sse-starlette>=2.0",
|
||||
# Data
|
||||
"polars>=1.0",
|
||||
# 1.44 起含流式引擎 executor 线程调度修复; 上限锁定一个已验证的大版本,
|
||||
# 升级需重跑 scripts/stress_polars_concurrency.py 压测 (并发死锁回归)。
|
||||
"polars>=1.44,<1.45",
|
||||
"duckdb>=1.0",
|
||||
"pyarrow>=16.0",
|
||||
"pandas>=2.2", # 仅在 BacktestService 边界使用,见 §7.4 / ADR-19
|
||||
@@ -45,7 +47,7 @@ dependencies = [
|
||||
# machines without AVX2/FMA support.
|
||||
# Enable with: uv sync --extra legacy-cpu
|
||||
legacy-cpu = [
|
||||
"polars[rtcompat]>=1.0",
|
||||
"polars[rtcompat]>=1.44,<1.45",
|
||||
]
|
||||
|
||||
# vectorbt 还会引入绘图、交互组件等完整分析栈,仍保持为可选 extras。
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user