Compare commits
86
Commits
2a52ee5e83
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d53c4d3ca9 | ||
|
|
a395db892d | ||
|
|
e58367af27 | ||
|
|
c974415691 | ||
|
|
57d6082f91 | ||
|
|
a3055eff5d | ||
|
|
2a88ca6076 | ||
|
|
bb48c91853 | ||
|
|
e8fa67ad40 | ||
|
|
7e369d9680 | ||
|
|
633176a3d1 | ||
|
|
48a97c2a12 | ||
|
|
36fe018075 | ||
|
|
f13b34c59e | ||
|
|
2e90f3eeac | ||
|
|
40bd603b44 | ||
|
|
50a1030afa | ||
|
|
d1287799f7 | ||
|
|
9aaca12751 | ||
|
|
82240e383d | ||
|
|
23972e7063 | ||
|
|
7e15b7251e | ||
|
|
adaf97f6d7 | ||
|
|
861a4051ca | ||
|
|
ed54096331 | ||
|
|
7c268e43df | ||
|
|
67d3aa1349 | ||
|
|
0d05bfd187 | ||
|
|
93e32f4e63 | ||
|
|
9cc4bfccac | ||
|
|
03fb463216 | ||
|
|
37510c1b89 | ||
|
|
5bde8f9f5f | ||
|
|
e1a0a8aa38 | ||
|
|
63c61ded37 | ||
|
|
bfeac7aa4c | ||
|
|
995ed08548 | ||
|
|
8abfd6538c | ||
|
|
2836efc607 | ||
|
|
314bfc159f | ||
|
|
db147c4232 | ||
|
|
e5a23f176d | ||
|
|
e1ac23fa25 | ||
|
|
8b2f8ac35c | ||
|
|
9d25d466e5 | ||
|
|
692bdb3be5 | ||
|
|
ba52edc2d6 | ||
|
|
ef09d5b419 | ||
|
|
4fa2bb748e | ||
|
|
273aee2772 | ||
|
|
8f47b5b603 | ||
|
|
0ab9038570 | ||
|
|
0d3e123de3 | ||
|
|
c60dc78c88 | ||
|
|
75c5472c31 | ||
|
|
25a1d9531a | ||
|
|
f3586adb25 | ||
|
|
697ffc767b | ||
|
|
0ffd574f30 | ||
|
|
6c2f198261 | ||
|
|
a3fabf9ae9 | ||
|
|
db520d0430 | ||
|
|
24f98e90c6 | ||
|
|
442999f701 | ||
|
|
a77d3c13c3 | ||
|
|
1cf7c84897 | ||
|
|
55b684fbe1 | ||
|
|
195f5d41f4 | ||
|
|
02e42184be | ||
|
|
56254172b3 | ||
|
|
e2741a0236 | ||
|
|
778c4beb07 | ||
|
|
b78852f01b | ||
|
|
fc12ba89ad | ||
|
|
bbb5c1ea52 | ||
|
|
b8f67f99ae | ||
|
|
880f4c50fe | ||
|
|
c5349bf9da | ||
|
|
01d818e12a | ||
|
|
8f8b6d274f | ||
|
|
d9be75a98f | ||
|
|
0ea229d766 | ||
|
|
92627f5b6b | ||
|
|
e9f59d3cf8 | ||
|
|
2da234220a | ||
|
|
7a89d97c0b |
+13
-5
@@ -9,16 +9,24 @@
|
||||
TUSHARE_TOKEN=
|
||||
|
||||
# ---- 数据库连接 ----
|
||||
# 留空时使用默认 SQLite:<项目根>/data/quant.db(相对路径自动解析到项目根)
|
||||
# 默认库由 config.yaml database.mysql 决定(本机 MariaDB 127.0.0.1/qlib)。
|
||||
# 连接优先级:本文件 DATABASE_URL > config.yaml database.mysql > SQLite 兜底。
|
||||
DATABASE_URL=
|
||||
# SQLite 示例(显式指定):
|
||||
# DATABASE_URL=sqlite:///./data/quant.db
|
||||
# MySQL 示例(未来切换,仅需改此值,业务代码不变):
|
||||
# DATABASE_URL=mysql+pymysql://user:password@127.0.0.1:3306/quant
|
||||
# MySQL 示例(手动覆盖 config.yaml 的 mysql 段时):
|
||||
# DATABASE_URL=mysql+pymysql://user:password@127.0.0.1:3306/qlib
|
||||
|
||||
# ---- AI Agent / LLM(Phase 5 接入,先留空)----
|
||||
# MySQL 密码(config.yaml database.mysql.password_env 引用;host/port/db/user 在 config.yaml)
|
||||
MYSQL_PASSWORD=
|
||||
|
||||
# ---- AI Agent / LLM(Phase 5)----
|
||||
# 只需填写 API Key;URL 与模型名已在 config.yaml 的 agent.llm 中配置
|
||||
# (默认 dashscope 兼容端点 + qwen-plus)。如需覆盖 config.yaml 的值,
|
||||
# 再取消注释下面两行。
|
||||
LLM_API_KEY=
|
||||
LLM_BASE_URL=
|
||||
# LLM_BASE_URL=https://dashscope.aliyuncs.com/compatible-mode/v1
|
||||
# LLM_MODEL=qwen-plus
|
||||
|
||||
# ---- API 应用密钥(必须修改为随机值)----
|
||||
APP_SECRET_KEY=please-change-me
|
||||
|
||||
+11
@@ -17,6 +17,7 @@ build/
|
||||
# ================= 密钥 / 凭证 =================
|
||||
# 真实密钥只放根目录 .env(复制自 .env.example),严禁提交
|
||||
.env
|
||||
.env.local
|
||||
*.pem
|
||||
*.key
|
||||
|
||||
@@ -48,3 +49,13 @@ experiments/*
|
||||
.idea/
|
||||
.vscode/
|
||||
*.swp
|
||||
*.tsbuildinfo
|
||||
|
||||
# ================= archify 视觉验证副产物(重新生成即可) =================
|
||||
docs/diagrams/*.visual-check.*
|
||||
# 界面自检截图:临时证据,不入库(要看图直接打开磁盘上的文件)
|
||||
docs/screenshots/
|
||||
|
||||
# ================= 运行期(dev.sh 管理) =================
|
||||
.run/
|
||||
logs/
|
||||
|
||||
@@ -10,6 +10,49 @@
|
||||
|
||||
---
|
||||
|
||||
# 0. 网络下载与代理规则
|
||||
|
||||
当网络下载(git clone、pip / uv 安装、wget、数据下载、Docker pull 等)出现**困难或超时**时,
|
||||
统一使用本机局域网 HTTP 代理:
|
||||
|
||||
```text
|
||||
192.168.1.160:3128
|
||||
```
|
||||
|
||||
示例(仅在下载失败 / 超时时启用,不把代理写入代码或 Git):
|
||||
|
||||
```bash
|
||||
# git clone(仅对 GitHub 等外网 https 目标使用)
|
||||
git -c http.proxy=http://192.168.1.160:3128 -c https.proxy=http://192.168.1.160:3128 clone <url>
|
||||
|
||||
# pip / uv 安装
|
||||
export http_proxy=http://192.168.1.160:3128
|
||||
export https_proxy=http://192.168.1.160:3128
|
||||
uv pip install <pkg>
|
||||
|
||||
# wget / curl
|
||||
wget -e use_proxy=yes -e http_proxy=http://192.168.1.160:3128 <url>
|
||||
```
|
||||
|
||||
注意事项:
|
||||
|
||||
- 代理只用于**外网**下载;内网资源(如 ssh git@192.168.1.10、局域网服务)不要走代理。
|
||||
- 代理 IP 属于局域网配置,不写入 `.env` / `config.yaml` / 任何提交进 Git 的文件。
|
||||
- 直接下载成功时不要多此一举走代理。
|
||||
|
||||
### 0.1 数据库目标:只允许本机 MariaDB(硬约束)
|
||||
|
||||
- 数据库一律指向**本机 MariaDB 10.11**(`127.0.0.1:3306/qlib`)。
|
||||
- **禁止把 `192.168.1.10` 作为数据库目标**(该地址在本文档中仅作外网代理白名单/SSH 提及,
|
||||
不是数据库目标)。远端旧库只作为历史数据来源,已不再写入。
|
||||
- 该约束不是「靠自觉」:`app/core/config.py::assert_db_target_allowed` 在 `get_settings()`
|
||||
解析出 URL 后立刻校验,命中禁用主机**直接抛错**(应用起不来),不允许「起得来但连错库」。
|
||||
禁用列表默认 `192.168.1.10`,可用环境变量 `QLIB_FORBIDDEN_DB_HOSTS` 覆盖(空串=不限制)。
|
||||
- 理由:连错库属于最危险的静默错误 —— 不报错、界面正常,回测/归档/策略却写到了另一台机器上。
|
||||
- 新增脚本/文档/示例配置时,DB 主机一律写 `127.0.0.1`,禁止出现远端库地址作为**默认值**。
|
||||
|
||||
---
|
||||
|
||||
# 1. 项目定位
|
||||
|
||||
这是一个:
|
||||
@@ -774,6 +817,54 @@ factor_exposure
|
||||
|
||||
前端组件只依赖这个结构。
|
||||
|
||||
## 27.1 买卖理由:必须能回答「为什么买 / 为什么卖」,且数字来自引擎
|
||||
|
||||
用户看回测结果时的第一个问题不是「赚了多少」,而是「**为什么在这里买 / 卖**」。
|
||||
因此每个买卖点(**成交的与未成交的都要**)必须带结构化理由
|
||||
(`TradeReason{code, text, data}`,见 `backend/app/quant/trade_reasons.py`):
|
||||
|
||||
- **原因分类是封闭词表**:组合引擎(`combo_engine.py`)与单策略引擎(`local_engine.py`)
|
||||
共用同一套 code 与文案构造器,禁止各自手写措辞 —— 否则同一件事会出现两种说法。
|
||||
- **`data` 里只能是引擎当时的真实数字**(名次 / 候选数 / 综合分 / 各因子**原始值** /
|
||||
持有交易日 / 预算 / 涨停比值…)。前端**只展示不推算**;拿不到名次就写「未给出名次」,
|
||||
不拿旧名次或其他日期的数据冒充。
|
||||
- **把事实说准**:跌出 TopN ≠ 不在候选池(被股票池/条件过滤)≠ 全量换仓
|
||||
(策略每次调仓先清仓,被卖的股票可能仍排在前列)≠ Tmin 保护暂留 ≠ 超 Tmax 强制了结
|
||||
≠ 涨停/跌停/停牌/现金不足。宁可为一种情况新增一个 code,也不要套一个语义不符的旧 code。
|
||||
- **因子曲线**(`result.factor_curves`)= 当日**持仓按市值加权平均的原始值**
|
||||
(不做 z-score、不按方向取反,空仓日不落点、不插值、不用 0 填充),
|
||||
界面必须同时写出方向与单位(如「股息率 %,越高越好」),否则读者会误判曲线的含义。
|
||||
- **每条曲线都要能新页面放大**(`/charts/{归档id}?s=...`):放大页从**归档**读同一份数据
|
||||
(URL 可分享、口径不漂移);没有归档 id 时如实说明「未归档,无法放大」,不给坏链接。
|
||||
|
||||
## 27.2 长任务反馈:点了立刻有字、看得出在动、出事了能自救
|
||||
|
||||
用户原话:「点了回测没有任何反馈,不清楚是不是已经开始」。因此跑异步 Job 的页面必须:
|
||||
|
||||
- 提交**瞬间**就有反馈(提交中 → 排队中),不等到后端返回;
|
||||
- 显示**真实**作业号 / 阶段 / 逐秒自增的已用时间(用 `lib/jobs.ts` 的 `useJobRunner`
|
||||
+ `components/JobProgress.tsx`,不要各页手写一套);
|
||||
- 排队/运行中可**取消**;失败显示后端原文;成功给「打开归档 / 去对比」入口;
|
||||
- 反馈条出现时**自动滚入视野**,并带 `role="status" aria-live="polite"`;
|
||||
- 同步接口(如 `POST /api/signals`)**没有**作业号与阶段时,如实说明「同步请求、无阶段、
|
||||
不可取消」,**禁止**合成假作业号或假进度喂给反馈组件。
|
||||
|
||||
## 27.3 宽表:两端固定 + 按需提示,绝不"右侧被切掉"
|
||||
|
||||
用户原话:「experiments 页面右侧内容溢出了」。实测原因不是整页溢出,而是**卡片内的横向滚动**:
|
||||
在 macOS 上覆盖式滚动条不滚动就不显示,于是看起来就是内容被切掉、也没有滚动条可拉。因此宽表:
|
||||
|
||||
- 列数 ≥ 9 或列宽会随数据增长的表(实验列表 `.tbl--wide`、对比表 `.tbl--pin-first`、
|
||||
月度收益表 `.tbl--monthly`)必须**固定首列**;操作列在右端时用 `.tbl--wide` 固定右端,
|
||||
保证任何窗口宽度下「看的是哪一行」和「能点哪里」都在视野内;
|
||||
- 固定列必须有不透明底色(卡片是渐变,取近似纯色)并覆盖 `:hover` / `.row-active`,
|
||||
否则中间列会从下面透出来、整行高亮在两端断开;
|
||||
- 列宽按**实测单行内容宽度**定(用 CDP 量 `scrollWidth` 与单元格内容宽度),不要凭感觉写,
|
||||
并且**不要在 JSX 里写行内 `width`**——它会盖掉 CSS(曾让「操作」列多占 88px,整表放不下);
|
||||
- 「可横向滚动」提示只在实测 `scrollWidth > clientWidth` 时出现(`useTableScrollHint`),
|
||||
窗口够宽时不显示废话;
|
||||
- 门禁:`python3 scripts/verify_ui_alignment.py`(160 项)必须 0 失败,含 375px 与 1500px 两档。
|
||||
|
||||
---
|
||||
|
||||
# 28. AI Agent 约束
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
# qlib-platform
|
||||
|
||||
个人 A 股量化研究平台:**选股 · 因子研究 · 回测**,面向低频 / 中低频交易研究。
|
||||
个人 A 股量化研究平台:**股票筛选 · 因子研究 · 交易信号 · 选股回测**,面向低频 / 中低频交易研究。
|
||||
|
||||
> 核心量化引擎:[Qlib](https://github.com/microsoft/qlib) · 首选数据源:Tushare(新浪财经为备用)· 当前数据库:SQLite(未来 MySQL,业务层无感切换)
|
||||
> 核心量化引擎:[Qlib](https://github.com/microsoft/qlib) · 首选数据源:Tushare(新浪财经为备用)· 当前数据库:MySQL(业务层无感切换;SQLite 仅兜底)
|
||||
|
||||
## 定位
|
||||
|
||||
@@ -16,13 +16,17 @@
|
||||
详见 [AGENT.md](./AGENT.md)(开发约束)与 [docs/ARCHITECTURE.md](./docs/ARCHITECTURE.md)(架构文档),**改代码前必须阅读**。
|
||||
|
||||
```text
|
||||
Web 前端 (frontend/web)
|
||||
↓ REST / SSE
|
||||
FastAPI (backend)
|
||||
↓ Research Specification
|
||||
Application Service ── Quant Service ── Qlib Adapter ── Qlib
|
||||
↓ ↓
|
||||
Repository / DAO Parquet / SQLite
|
||||
Web 前端 (frontend/web:总览 / 股票池 / 股票筛选 / 因子研究 / 因子组合 / 交易信号 / 选股回测 / 实验 / 归档详情)
|
||||
↓ REST / SSE(异步 Job 状态机)
|
||||
FastAPI (backend:业务对象 API,见下「核心能力」)
|
||||
↓ Research Specification(统一研究契约)
|
||||
Application Service(SelectionService / SignalService / ResearchService / Strategy …)
|
||||
├── Selection Engine(条件选股 / 因子评分 / as_of 历史与当前一致)
|
||||
├── Signal Engine(BUY / WATCH / SELL + 理由)
|
||||
├── Portfolio Engine(等权;约束预留并如实标注)
|
||||
└── Quant Service ── Composite Engine ── Qlib Adapter ── Qlib
|
||||
↓ ↓
|
||||
Repository / DAO MySQL(默认)/ Parquet
|
||||
```
|
||||
|
||||
### 目录结构
|
||||
@@ -35,7 +39,8 @@ qlib/
|
||||
│ │ ├── application/ # 用例 / 应用服务(编排,不含框架细节)
|
||||
│ │ ├── domain/ # 领域实体 + Repository Protocol
|
||||
│ │ ├── infrastructure/ # SQLAlchemy / Alembic 等基础设施实现
|
||||
│ │ ├── quant/qlib_adapter/ # Qlib 适配层(业务禁止直接 import qlib)
|
||||
│ │ ├── quant/ # 因子 / composite / universe / selection / signal / portfolio 引擎
|
||||
│ │ │ └── qlib_adapter/ # Qlib 适配层(业务禁止直接 import qlib)
|
||||
│ │ ├── agent/ # AI Research Agent(Phase 5,仅 Tool 访问)
|
||||
│ │ └── core/ # 配置、通用组件
|
||||
│ └── tests/
|
||||
@@ -50,6 +55,10 @@ qlib/
|
||||
└── .env.example # 密钥模板(复制为 .env,勿提交)
|
||||
```
|
||||
|
||||
## 使用说明
|
||||
|
||||
完整使用文档(安装 / 配置 / 数据同步 / API / Agent / 常见问题)见 **[docs/USAGE.md](./docs/USAGE.md)**。
|
||||
|
||||
## 快速开始(后端)
|
||||
|
||||
前置:安装 [uv](https://docs.astral.sh/uv/)(`pip install uv` 或官方脚本)。
|
||||
@@ -60,10 +69,16 @@ cp .env.example .env
|
||||
|
||||
# 2. 安装依赖(自动使用 Python 3.12,见 backend/.python-version)
|
||||
cd backend
|
||||
uv sync
|
||||
uv sync # 含 pyqlib(GitHub 源码依赖,固定 commit)。若网络下载困难/超时,按 AGENT.md §0 设置代理 192.168.1.160:3128 后重试
|
||||
|
||||
# 3. 运行测试
|
||||
uv run pytest
|
||||
uv run pytest # 全量 388 条
|
||||
uv run ruff check app tests
|
||||
|
||||
# 3b. 端到端契约自检(真实提交回测 Job,验证页面↔后端字段不漂移)
|
||||
PYTHONPATH=. .venv/bin/python ../scripts/verify_strategy_workspace.py # 策略库/说明/名称/选股直通/归档链路(59 项)
|
||||
PYTHONPATH=. .venv/bin/python ../scripts/verify_backtest_page_contract.py # 回测结果结构契约
|
||||
python3 ../scripts/verify_ui_alignment.py # UI 对齐与控件一致性(140 项,需前端已启动)
|
||||
|
||||
# 4. 启动开发服务
|
||||
uv run uvicorn app.main:app --reload --port 8000
|
||||
@@ -71,21 +86,47 @@ uv run uvicorn app.main:app --reload --port 8000
|
||||
# 交互文档:http://127.0.0.1:8000/docs
|
||||
```
|
||||
|
||||
### 数据库迁移(Alembic,已就位)
|
||||
### 数据库(默认 MySQL)
|
||||
|
||||
- 连接在根目录 `config.yaml → database.mysql`(host/port/db/user/charset),密码放 `.env` 的
|
||||
`MYSQL_PASSWORD`;URL 优先级:`DATABASE_URL` 环境变量 > `database.mysql` > SQLite 兜底。
|
||||
- 结构变更(Model → Migration → Test):
|
||||
|
||||
```bash
|
||||
cd backend
|
||||
uv run alembic upgrade head # 首次运行会在 data/quant.db 建立版本表
|
||||
uv run alembic revision --autogenerate -m "add xxx table" # 修改 Model 后生成迁移
|
||||
uv run alembic upgrade head # 在当前配置指向的库上建表(默认 MySQL qlib)
|
||||
uv run alembic revision --autogenerate -m "add xxx table"
|
||||
```
|
||||
|
||||
## 开发阶段(对应架构文档 §20)
|
||||
## 核心能力(现状)
|
||||
|
||||
- **Phase 1 数据**:Tushare → 标准化 → SQLite(stock / daily / 复权因子 / 交易日历 / 财务指标),Parquet 导出
|
||||
- **Phase 2 Qlib**:Parquet → Qlib Dataset → 因子(Alpha158 / 自定义)→ LightGBM → 回测
|
||||
- **Phase 3 Web**:股票池 / 因子研究 / 选股 / 回测 / 结果可视化(Next.js + ECharts)
|
||||
- **Phase 4 Experiment**:所有研究自动可复现存档
|
||||
- **Phase 5 AI Agent**:自然语言 → Research Plan → 受控 Tool → Experiment
|
||||
| 能力 | 说明 / 入口 |
|
||||
|---|---|
|
||||
| 股票筛选 | `POST /api/selections`(A 条件/B 评分,`as_of` 历史/当前,可解释);`POST /api/selections/jobs` 全市场异步;Web `/selection` |
|
||||
| 因子目录 | `factor_definition` 落库;`GET /api/factors`;组合可保存复用 `POST /api/composites` |
|
||||
| 因子研究 | Job 异步 IC / RankIC / 分层测试(`/api/jobs`) |
|
||||
| 交易信号 | `POST /api/signals`:评分排名 + 趋势规则 → BUY/WATCH/SELL + 理由;Web `/signals` |
|
||||
| 选股回测 | `POST /api/backtests`(成本/涨跌停近似;与当前选股共用同一评分引擎,v2 §25 一致性) |
|
||||
| 定期调仓回测 | **两级截断(候选池 n → 持仓 x)+ 双周期(每 m 月择股 / 每 y 月调仓)+ 复权口径(none/qfq/hfq)+ 组合条件过滤 + 最低佣金**;输出净值/个股曲线并标注买卖点;Web `/backtest`(含高股息案例预设)、`scripts/run_dividend_case.py` |
|
||||
| 每日指标 | `daily_basic`(股息率 `dv_ratio`/`dv_ttm`、PE/PB/市值)+ `dividend_yield` 因子;`sync daily_basic` 回补 |
|
||||
| 退市股与时点股票池 | `sync basic --include-delisted` 入库已退市/暂停上市(`status`/`delist_date`)+ 定向 `sync daily` 补行情;`filter_stocks` 按 `as_of` 正确纳入/排除(实测 338 只) |
|
||||
| 时点 ST / 名称历史 | `sync namechange` 入库名称生效区间(14,213 行);`exclude_st` 在**每个择股日**按当时名称判定(回测与 `/api/selections` 同口径),消除「曾高股息后 ST」的股息陷阱隐藏偏差(实测 3.70pp) |
|
||||
| 策略库 | `strategy` 落库 + `/api/strategies` CRUD/`PUT` 原地更新/展开为回测 spec;**每个策略有后端推导的「一句话说明 + 计算公式 + 执行步骤 + 注意事项」**(`describe_strategy`,与引擎实执行规则同源);一键回测;Web `/strategies` |
|
||||
| Experiment / 归档 | **完整存档**:每次回测(异步 Job 与同步接口都算)把结果 + spec + `code_version` + `data_version` 数据快照指纹写入 `experiment` 表,个股收益曲线默认**全量保存**(超出体积预算才裁剪并显式标注);列表支持 `kind`/`q` 过滤且总数经 `X-Total-Count` 暴露(不再静默截断);`DELETE /api/experiments/{id}` + `python -m app.cli.prune_experiments --keep N`(默认 dry-run,删除即失去结果,建议先导出)治理体积;**实验对比**:勾选 2~3 个 → 归一化净值曲线叠加 + 指标差值表 + `config_snapshot` 参数 diff;**归档详情页 `/experiments/{id}`**(Server Component)只读复看完整结果,并明确回答「**选股条件**」与「**交易执行依据**」;归档页可**导出完整 JSON / 以此参数再跑 / 删除**(确认框写明后果);体积用 `app.cli.prune_experiments`(默认 dry-run)治理,删错的历史归档可用 `app.cli.restore_experiment_from_job` 从 Job 副本按原 id 重建 |
|
||||
| 图表基座 | **统一 TradingView Lightweight Charts 4.2.3**:`components/charts/LwChart.tsx`(折线/面积/柱状 + 买卖点标记 + tooltip + 可点击图例)与 `components/StockChart/CandleChartLW.tsx`(个股 K 线 + 成交量 + MA);ECharts 已完全下线(`package.json`、`pnpm-lock.yaml`、`node_modules`、文档与架构图标注均已清除)。实测个股页图表根节点是 Lightweight Charts 自身的 `div.tv-lightweight-charts`,页面 canvas 无一来自其它图表库 |
|
||||
| 代码即名称 | 任何出现股票代码的位置都成对显示股票名称且可点击进入个股页(基本信息 + 走势图);后端在 `SelectionCandidate/SymbolCurve/ActionRecord/RankedPick/Position/Trade` 上填充 `name`,前端 `GET /api/stocks/names` 一次性缓存兜底 |
|
||||
| AI Agent | `POST /api/agent/chat`,**14 个受控 Tool**(v3 §25 全清单) |
|
||||
|
||||
## 里程碑(详见 [ROADMAP.md](./docs/ROADMAP.md) 与 [docs/DEV_PLAN_v2.md](./docs/DEV_PLAN_v2.md))
|
||||
|
||||
- **M0–M5**:工程骨架 → 数据层 → 研究引擎 → Web → Experiment/Job → AI Agent
|
||||
- **M-DB**:SQLite 全量迁移 MySQL(config.yaml 配置化 + 迁移工具 + 一致性校验)
|
||||
- **M6**:选股系统主线(Universe + Selection A/B + 落库/API + 回测共用 + Web)
|
||||
- **M7**:因子层(因子定义入库 + Composite 模块化 + 行情口径显式化)
|
||||
- **M8**:Signal / Portfolio / Strategy / Agent 工具 / Web 做实(M8.4 Qlib 模型选股按需延后)
|
||||
- **V3**(见 DEV_PLAN_v3):Chart Service/K线与成交点、Signal↔Fill 区分、个股研究页、
|
||||
复权口径、Bar Replay、指数历史成分、因子相关性、组合单股上限、Job stage/取消、
|
||||
选股异步化、Agent 14 工具(B2/B3 停牌/ST/报表与 C3 模型选股/D3 Redis 按需延后)
|
||||
|
||||
## 约定速查
|
||||
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
"""LLM 客户端抽象(Phase 5)。
|
||||
|
||||
实现约定:真实 Key 来自 .env(LLM_API_KEY / LLM_BASE_URL / LLM_MODEL,见 config.yaml 引用)。
|
||||
测试注入 FakeLLMClient 走完整编排链路,不触网。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
import httpx
|
||||
|
||||
from app.agent.tools import tools_schema
|
||||
from app.agent.tools_impl import build_tools
|
||||
from app.core.config import get_settings
|
||||
|
||||
SYSTEM_TEMPLATE = """你是个人 A 股量化研究助手(Research Assistant),不是系统管理员。
|
||||
|
||||
你可以调用以下工具(每次只能输出一个动作):
|
||||
{tools}
|
||||
|
||||
输出规则:只输出一行 JSON,两种形态之一:
|
||||
1. 需要调用工具:{{"tool": "<工具名>", "args": {{...}}}}
|
||||
2. 给出结论:{{"final": "结论文本"}}
|
||||
|
||||
研究纪律(必须遵守):
|
||||
- 先提出假设 → 用工具做因子测试或回测 → 基于实验事实分析,再给结论
|
||||
- 不得仅凭单次回测高收益就宣布策略有效;要主动说明样本外、过拟合、
|
||||
look-ahead bias、交易成本、参数敏感性等风险(未做验证的项要明说「未验证」)
|
||||
- 全程只读:不得要求删除/修改数据或执行任意命令(你也没有这类工具)
|
||||
- 回答使用简体中文
|
||||
"""
|
||||
|
||||
|
||||
class LLMClient(Protocol):
|
||||
def chat(self, messages: list[dict]) -> str: ...
|
||||
|
||||
|
||||
class OpenAICompatibleClient:
|
||||
"""OpenAI 兼容 Chat Completions(qwen/dashscope、deepseek、openai 等均适用)。"""
|
||||
|
||||
def __init__(self, *, base_url: str, api_key: str, model: str, timeout: float = 60.0) -> None:
|
||||
self._base_url = base_url.rstrip("/")
|
||||
self._api_key = api_key
|
||||
self._model = model
|
||||
self._timeout = timeout
|
||||
|
||||
def chat(self, messages: list[dict]) -> str:
|
||||
with httpx.Client(timeout=self._timeout) as client:
|
||||
resp = client.post(
|
||||
f"{self._base_url}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {self._api_key}"},
|
||||
json={
|
||||
"model": self._model,
|
||||
"messages": messages,
|
||||
"temperature": 0.2,
|
||||
},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return data["choices"][0]["message"]["content"]
|
||||
|
||||
|
||||
def build_llm_from_settings():
|
||||
"""从配置构造 LLM;未配置 Key 时返回 None(调用方给出引导提示)。"""
|
||||
settings = get_settings()
|
||||
if not settings.llm_api_key:
|
||||
return None
|
||||
return OpenAICompatibleClient(
|
||||
base_url=settings.llm_base_url or "https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||
api_key=settings.llm_api_key,
|
||||
model=settings.llm_model,
|
||||
)
|
||||
|
||||
|
||||
def system_prompt() -> str:
|
||||
return SYSTEM_TEMPLATE.format(tools=tools_schema(build_tools()))
|
||||
@@ -0,0 +1,98 @@
|
||||
"""Agent 编排:自然语言 → 受控工具调用循环 → 结论(AGENT.md §29 假设-实验-分析循环)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from app.agent.llm import SYSTEM_TEMPLATE, LLMClient
|
||||
from app.agent.tools import Tool, tools_schema
|
||||
from app.agent.tools_impl import build_tools
|
||||
|
||||
MAX_TOOL_ROUNDS = 5
|
||||
_DECISION_PATTERN = re.compile(r"\{.*\}", re.DOTALL)
|
||||
|
||||
|
||||
def _parse_decision(text: str) -> dict[str, Any]:
|
||||
"""容忍 LLM 输出中的代码块/前后缀,提取首个 JSON 对象。"""
|
||||
match = _DECISION_PATTERN.search(text)
|
||||
if not match:
|
||||
raise ValueError(f"无法从模型输出中解析动作 JSON:{text[:200]}")
|
||||
try:
|
||||
return json.loads(match.group(0))
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"模型输出的 JSON 不合法:{text[:200]}") from exc
|
||||
|
||||
|
||||
class AgentService:
|
||||
"""受控研究 Agent:每轮让 LLM 决策 tool/final,执行工具并把结果回喂,直至 final。"""
|
||||
|
||||
def __init__(self, llm: LLMClient, tools: list[Tool] | None = None) -> None:
|
||||
self._llm = llm
|
||||
self._tools = {t.name: t for t in (tools or build_tools())}
|
||||
|
||||
def _by_name(self, name: str) -> Tool:
|
||||
tool = self._tools.get(name)
|
||||
if tool is None:
|
||||
raise ValueError(f"工具不存在:{name}(可用 {sorted(self._tools)})")
|
||||
return tool
|
||||
|
||||
def chat(self, message: str, max_rounds: int = MAX_TOOL_ROUNDS) -> dict:
|
||||
messages: list[dict] = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": SYSTEM_TEMPLATE.format(tools=tools_schema(list(self._tools.values()))),
|
||||
},
|
||||
{"role": "user", "content": message},
|
||||
]
|
||||
actions: list[dict] = []
|
||||
|
||||
for _round in range(max_rounds):
|
||||
try:
|
||||
decision = _parse_decision(self._llm.chat(messages))
|
||||
except ValueError as exc:
|
||||
# LLM 回复格式异常:回传错误要求重试一次结构化输出
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"输出格式错误,请只输出一行 JSON(tool 或 final):{exc}",
|
||||
}
|
||||
)
|
||||
continue
|
||||
if "final" in decision:
|
||||
return {"reply": str(decision["final"]), "actions": actions}
|
||||
tool_name = str(decision.get("tool", ""))
|
||||
args = decision.get("args") or {}
|
||||
if not isinstance(args, dict):
|
||||
args = {}
|
||||
tool = self._by_name(tool_name) # 白名单之外的调用直接报错
|
||||
output = tool.invoke(args)
|
||||
actions.append({"tool": tool_name, "args": args, "output": output[:1000]})
|
||||
messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": f'{{"tool": "{tool_name}", "args": {json.dumps(args, ensure_ascii=False)}}}',
|
||||
}
|
||||
)
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"工具 {tool_name} 返回:\n{output}\n请继续(如需再调用输出 tool,否则输出 final)。",
|
||||
}
|
||||
)
|
||||
|
||||
# 轮次耗尽:强制收尾
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": "已达工具调用轮次上限,请直接给出基于已有事实的最终结论(只输出 final JSON)。",
|
||||
}
|
||||
)
|
||||
try:
|
||||
decision = _parse_decision(self._llm.chat(messages))
|
||||
except ValueError:
|
||||
decision = {
|
||||
"final": "研究轮次耗尽且模型未给出结构化结论,请人工查看 actions 中的实验输出。"
|
||||
}
|
||||
return {"reply": str(decision.get("final", decision)), "actions": actions}
|
||||
@@ -0,0 +1,45 @@
|
||||
"""AI Research Agent(Phase 5)。
|
||||
|
||||
Agent 是 Research Assistant(AGENT.md §28/§29):
|
||||
- 只能调用本目录 tools 提供的**白名单受控工具**(只读研究能力)
|
||||
- 禁止:执行任意 shell / 修改删除数据 / 修改配置与凭证
|
||||
- 研究行为:提出假设 → 建立实验 → 运行测试 → 分析 → 下一步,禁止以单次高收益宣告策略有效
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
TOOL_CALL_TAG = "__tool__"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Tool:
|
||||
"""受控工具元数据(LLM 可见的 JSON Schema + 调用实现)。"""
|
||||
|
||||
name: str
|
||||
description: str
|
||||
parameters: dict[str, Any]
|
||||
handler: Callable[[dict[str, Any]], str] = field(repr=False)
|
||||
|
||||
def schema(self) -> dict[str, Any]:
|
||||
return {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"parameters": self.parameters,
|
||||
}
|
||||
|
||||
def invoke(self, args: dict[str, Any]) -> str:
|
||||
"""执行工具;任何异常都转为可读错误串(不向上抛,避免中断整轮对话)。"""
|
||||
try:
|
||||
result = self.handler(args)
|
||||
return result if isinstance(result, str) else json.dumps(result, ensure_ascii=False)
|
||||
except Exception as exc: # noqa: BLE001 —— 工具异常反馈给 LLM 而非崩溃
|
||||
return json.dumps({"error": f"{type(exc).__name__}: {exc}"}, ensure_ascii=False)
|
||||
|
||||
|
||||
def tools_schema(tools: list[Tool]) -> str:
|
||||
return json.dumps([t.schema() for t in tools], ensure_ascii=False, indent=1)
|
||||
@@ -1 +0,0 @@
|
||||
"""Agent 受控工具集:每个 Tool 是后端服务的只读 / 沙箱化入口。"""
|
||||
@@ -0,0 +1,668 @@
|
||||
"""受控工具集实现(AGENT.md §28):Agent 只能调用这里的白名单工具。
|
||||
|
||||
全部工具经 Job/Experiment 链路或只读查询执行:
|
||||
- 不提供 shell / 任意代码执行 / 修改配置与凭证 / 删除数据
|
||||
- 任何研究都会产出 Experiment 归档(可复现)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import date
|
||||
|
||||
from app.agent.tools import Tool
|
||||
from app.application.services.job_executor import default_factories, submit_and_run
|
||||
from app.application.services.selection_service import SelectionService
|
||||
from app.application.services.signal_service import SignalService
|
||||
from app.domain.entities.composite import CompositeComponent, CompositeDefinition
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
FactorTestReport,
|
||||
ResearchSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
from app.domain.entities.signal import SignalRules
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import (
|
||||
SqlAlchemyCompositeRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
|
||||
SqlAlchemyFactorRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl import (
|
||||
SqlAlchemySelectionRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
|
||||
SqlAlchemyStrategyRepository,
|
||||
)
|
||||
from app.quant.factors import FactorError, get_factor
|
||||
|
||||
|
||||
def _day(text: str) -> date:
|
||||
return date.fromisoformat(text)
|
||||
|
||||
|
||||
def _split_factor_list(raw: str) -> list[str]:
|
||||
"""按逗号切因子列表,但**不切参数化因子键里的逗号**。
|
||||
|
||||
参数化因子的名字把参数写全了(`momentum(window=90,direction=higher_is_better)`),
|
||||
直接 `.split(",")` 会把它劈成「momentum(window=90」和「direction=…):0.7」两段,
|
||||
模型与用户只会收到「因子不存在」这种看不懂的错。括号深度感知的切分让两种写法都能用:
|
||||
|
||||
momentum_60,volatility_60
|
||||
momentum(window=90,direction=lower_is_better),volatility_60
|
||||
"""
|
||||
out: list[str] = []
|
||||
depth = 0
|
||||
buf: list[str] = []
|
||||
for ch in raw:
|
||||
if ch == "(":
|
||||
depth += 1
|
||||
elif ch == ")":
|
||||
depth = max(0, depth - 1)
|
||||
if ch == "," and depth == 0:
|
||||
out.append("".join(buf).strip())
|
||||
buf = []
|
||||
else:
|
||||
buf.append(ch)
|
||||
out.append("".join(buf).strip())
|
||||
return [x for x in out if x]
|
||||
|
||||
|
||||
def _split_name_weight(part: str) -> tuple[str, str]:
|
||||
"""把 `name:weight` 按**括号外**的第一个冒号切开(参数化键里的 `=`/`,` 不受影响)。"""
|
||||
depth = 0
|
||||
for i, ch in enumerate(part):
|
||||
if ch == "(":
|
||||
depth += 1
|
||||
elif ch == ")":
|
||||
depth = max(0, depth - 1)
|
||||
elif ch == ":" and depth == 0:
|
||||
return part[:i].strip(), part[i + 1 :].strip()
|
||||
return part.strip(), ""
|
||||
|
||||
|
||||
def _pick(mapping: dict, key: str, default=None):
|
||||
val = mapping.get(key, default)
|
||||
if isinstance(val, str):
|
||||
val = val.strip()
|
||||
if val == "":
|
||||
return default
|
||||
return val
|
||||
|
||||
|
||||
def _job_result_json(job, *, session_factory, experiment_repo_factory) -> str | None:
|
||||
"""取 Job 的完整结果 JSON。
|
||||
|
||||
2026-09 起完整结果只在 experiment 存一份(`job.result_json` 为 None),故先按
|
||||
`job.experiment_id` 回读归档;归档不存在 / 老记录再回退 `job.result_json`。
|
||||
"""
|
||||
if job.experiment_id:
|
||||
try:
|
||||
with session_factory() as session:
|
||||
exp = experiment_repo_factory(session).get(job.experiment_id)
|
||||
if exp is not None:
|
||||
return exp.result_json
|
||||
except Exception: # noqa: BLE001 —— 回读失败则回退 job 副本,不阻断工具
|
||||
pass
|
||||
return job.result_json
|
||||
|
||||
|
||||
def build_tools(factories: dict | None = None) -> list[Tool]:
|
||||
facts = factories or default_factories()
|
||||
session_factory = facts["session_factory"]
|
||||
stock_repo_f = facts["stock_repo_factory"]
|
||||
daily_repo_f = facts["daily_repo_factory"]
|
||||
exp_repo_f = facts["experiment_repo_factory"]
|
||||
|
||||
def search_stocks(args: dict) -> str:
|
||||
q = str(_pick(args, "q", "") or "").upper()
|
||||
with session_factory() as session:
|
||||
stocks = stock_repo_f(session).list()
|
||||
rows = [
|
||||
s for s in stocks if (not q) or q in s.symbol.upper() or q in (s.name or "").upper()
|
||||
][:15]
|
||||
if not rows:
|
||||
return "未找到匹配股票"
|
||||
return "\n".join(
|
||||
f"{s.symbol} {s.name} 行业={s.industry or '-'} 上市={s.list_date}" for s in rows
|
||||
)
|
||||
|
||||
def get_market_data(args: dict) -> str:
|
||||
symbol = str(_pick(args, "symbol", "")).upper()
|
||||
start = _day(str(_pick(args, "start", "2024-01-01")))
|
||||
end = _day(str(_pick(args, "end", date.today().isoformat())))
|
||||
with session_factory() as session:
|
||||
bars = daily_repo_f(session).get_range(symbol, start, end)
|
||||
if not bars:
|
||||
return f"{symbol} 在 {start}~{end} 无日线数据(可能未同步)"
|
||||
head, tail = bars[0], bars[-1]
|
||||
last = "\n".join(f"{b.trade_date} close={b.close}" for b in bars[-8:])
|
||||
change = float(tail.close) / float(head.close) - 1 if head.close and tail.close else None
|
||||
return (
|
||||
f"{symbol} {start}~{end} 共 {len(bars)} 根日线;"
|
||||
f"区间 {head.trade_date}→{tail.trade_date} 收盘 {head.close}→{tail.close}"
|
||||
f"(涨跌 {change * 100:.2f}% 若数据完整);最近 8 根:\n{last}"
|
||||
)
|
||||
|
||||
def _run_spec(spec: ResearchSpec, desc: str) -> str:
|
||||
job = submit_and_run(spec, factories=facts)
|
||||
if job.status != "success":
|
||||
return f"{desc} 执行失败:{job.error}"
|
||||
result_json = _job_result_json(
|
||||
job, session_factory=session_factory, experiment_repo_factory=exp_repo_f
|
||||
)
|
||||
if spec.type == "backtest":
|
||||
result = BacktestResult.model_validate_json(result_json or "{}")
|
||||
s = result.summary
|
||||
return (
|
||||
f"回测完成(Experiment {job.experiment_id},代码版本 {_code_version(job, exp_repo_f)})。"
|
||||
f"总收益 {s.total_return_pct:.2f}%,年化 {s.annual_return_pct:.2f}%,"
|
||||
f"Sharpe {s.sharpe:.2f},最大回撤 {s.max_drawdown_pct:.2f}%,"
|
||||
f"交易 {s.total_trades} 笔,平均换手 {s.avg_turnover_pct:.1f}%。"
|
||||
f"未建模约束 {len(result.unimplemented)} 项(成本/涨跌停近似见实验详情)。"
|
||||
)
|
||||
report = FactorTestReport.model_validate_json(result_json or "{}")
|
||||
qs = ", ".join(f"Q{q.quantile + 1}: {q.return_pct:.2f}%" for q in report.quantile_returns)
|
||||
return (
|
||||
f"因子测试完成(Experiment {job.experiment_id})。IC {report.ic_mean:.4f},"
|
||||
f"RankIC {report.rank_ic_mean:.4f},ICIR {report.icir:.2f},正收益占比 "
|
||||
f"{report.positive_ratio_pct:.1f}%,样本 {report.sample_days} 日;分层未来收益 {qs}。"
|
||||
f"注意:单因子测试不代表策略有效,需结合稳健性分析。"
|
||||
)
|
||||
|
||||
def test_factor(args: dict) -> str:
|
||||
name = str(_pick(args, "name", ""))
|
||||
start = _day(str(_pick(args, "start", "2024-01-01")))
|
||||
end = _day(str(_pick(args, "end", "2024-12-31")))
|
||||
spec = ResearchSpec(
|
||||
type="factor_test",
|
||||
universe={"exclude_st": True, "min_listing_days": 0},
|
||||
factors=[{"name": name, "weight": 1.0}],
|
||||
selection={"top_n": 10},
|
||||
rebalance="monthly",
|
||||
period=(start, end),
|
||||
)
|
||||
return _run_spec(spec, f"因子 {name} 测试")
|
||||
|
||||
def run_backtest(args: dict) -> str:
|
||||
factor_names = _split_factor_list(str(_pick(args, "factors", "momentum_60")))
|
||||
top_n = int(_pick(args, "top_n", 5) or 5)
|
||||
rebalance = str(_pick(args, "rebalance", "monthly"))
|
||||
exclude_st = bool(_pick(args, "exclude_st", True))
|
||||
start = _day(str(_pick(args, "start", "2024-01-01")))
|
||||
end = _day(str(_pick(args, "end", "2024-12-31")))
|
||||
spec = ResearchSpec(
|
||||
type="backtest",
|
||||
universe={"exclude_st": exclude_st, "min_listing_days": 0},
|
||||
factors=[{"name": n, "weight": 1.0} for n in factor_names],
|
||||
selection={"top_n": top_n},
|
||||
rebalance=rebalance,
|
||||
period=(start, end),
|
||||
)
|
||||
return _run_spec(spec, "回测")
|
||||
|
||||
def get_experiment(args: dict) -> str:
|
||||
exp_id = str(_pick(args, "experiment_id", "")).upper()
|
||||
with session_factory() as session:
|
||||
exp = exp_repo_f(session).get(exp_id)
|
||||
if exp is None:
|
||||
return f"Experiment {exp_id} 不存在(可用列表:GET /api/experiments)"
|
||||
spec = json.loads(exp.spec_json)
|
||||
return (
|
||||
f"Experiment {exp.id} [{exp.kind}] 因子={[f['name'] for f in spec.get('factors', [])]} "
|
||||
f"区间={spec.get('period')} 调仓={spec.get('rebalance')};摘要:{exp.summary_text or '-'} "
|
||||
f"代码版本={exp.code_version or '-'} 创建={exp.created_at}"
|
||||
)
|
||||
|
||||
def compare_experiments(args: dict) -> str:
|
||||
ids = [
|
||||
x.strip().upper()
|
||||
for x in str(_pick(args, "experiment_ids", "")).split(",")
|
||||
if x.strip()
|
||||
]
|
||||
if not ids:
|
||||
return "请提供 experiment_ids(逗号分隔)"
|
||||
with session_factory() as session:
|
||||
repo = exp_repo_f(session)
|
||||
rows = [(i, repo.get(i)) for i in ids]
|
||||
out = []
|
||||
for exp_id, exp in rows:
|
||||
if exp is None:
|
||||
out.append(f"{exp_id}: 不存在")
|
||||
else:
|
||||
spec = json.loads(exp.spec_json)
|
||||
out.append(
|
||||
f"{exp.id}: 因子={[f['name'] for f in spec.get('factors', [])]} "
|
||||
f"区间={spec.get('period')} → {exp.summary_text or '-'}"
|
||||
)
|
||||
return "\n".join(out)
|
||||
|
||||
def _scope_symbols(raw: str | None) -> list[str]:
|
||||
"""白名单(可选):避免全市场长任务拖垮同步对话(全市场可用 Web 页异步)。"""
|
||||
if not raw:
|
||||
return []
|
||||
return [x.strip().upper() for x in raw.split(",") if x.strip()][:60]
|
||||
|
||||
def screen_stocks(args: dict) -> str:
|
||||
factors = _split_factor_list(str(_pick(args, "factors", "momentum_60")))
|
||||
top_n = int(_pick(args, "top_n", 10) or 10)
|
||||
as_of = _day(str(_pick(args, "as_of", date.today().isoformat())))
|
||||
symbols = _scope_symbols(str(_pick(args, "symbols", "") or ""))
|
||||
if not symbols:
|
||||
return (
|
||||
"为避免全市场长任务(>1 分钟),请传 symbols 白名单(≤60,逗号分隔)"
|
||||
"或使用 Web 选股页执行全市场选股。"
|
||||
)
|
||||
query = SelectionQuery(
|
||||
universe=UniverseSpec(
|
||||
exclude_st=bool(_pick(args, "exclude_st", True)),
|
||||
min_listing_days=0,
|
||||
symbols=symbols,
|
||||
),
|
||||
factors=[{"name": f, "weight": 1.0} for f in factors],
|
||||
top_n=top_n,
|
||||
as_of=as_of,
|
||||
)
|
||||
with session_factory() as session:
|
||||
service = SelectionService(
|
||||
stock_repo_f(session), daily_repo_f(session)
|
||||
)
|
||||
result = service.select(query)
|
||||
if not result.candidates:
|
||||
return (
|
||||
f"{as_of} 无候选(范围 {result.statistics.universe_size} 只,"
|
||||
f"可评分 {result.statistics.evaluated})。如需白名单可传 symbols(≤60)。"
|
||||
)
|
||||
lines = [f"as_of={result.as_of_date} 选出 Top{len(result.candidates)}:"]
|
||||
for c in result.candidates:
|
||||
vals = ", ".join(f"{k}={v:.4f}" for k, v in c.factor_values.items())
|
||||
lines.append(f" #{c.rank} {c.symbol} score={c.score:.4f}({vals})")
|
||||
lines.append("入选理由见 explain_selection(selection_id)。")
|
||||
return "\n".join(lines)
|
||||
|
||||
def explain_selection(args: dict) -> str:
|
||||
sel_id = str(_pick(args, "selection_id", "")).upper()
|
||||
with session_factory() as session:
|
||||
repo = SqlAlchemySelectionRepository(session)
|
||||
result = repo.get(sel_id)
|
||||
if result is None:
|
||||
return f"选股记录 {sel_id} 不存在(先通过 Web 选股页或 screen_stocks 生成)"
|
||||
out = [f"选股 {sel_id} as_of={result.as_of_date}({result.method},选出 {len(result.candidates)} 只)"]
|
||||
for c in result.candidates[:10]:
|
||||
reasons = "; ".join(c.selection_reason[:3])
|
||||
out.append(f" #{c.rank} {c.symbol} score={c.score:.4f} — {reasons}")
|
||||
return "\n".join(out)
|
||||
|
||||
def generate_signals(args: dict) -> str:
|
||||
factors = _split_factor_list(str(_pick(args, "factors", "momentum_60")))
|
||||
as_of = _day(str(_pick(args, "as_of", date.today().isoformat())))
|
||||
symbols = _scope_symbols(str(_pick(args, "symbols", "") or ""))
|
||||
query = SelectionQuery(
|
||||
universe=UniverseSpec(
|
||||
exclude_st=bool(_pick(args, "exclude_st", True)),
|
||||
min_listing_days=0,
|
||||
symbols=symbols,
|
||||
),
|
||||
factors=[{"name": f, "weight": 1.0} for f in factors],
|
||||
top_n=int(_pick(args, "top_n", 50) or 50),
|
||||
as_of=as_of,
|
||||
)
|
||||
rules = SignalRules(
|
||||
buy_rank_threshold=int(_pick(args, "buy_rank", 20) or 20),
|
||||
sell_rank_threshold=int(_pick(args, "sell_rank", 50) or 50),
|
||||
)
|
||||
with session_factory() as session:
|
||||
res = SignalService(stock_repo_f(session), daily_repo_f(session)).signal(query, rules)
|
||||
out = [
|
||||
f"信号 as_of={res.as_of_date}: BUY {res.statistics.buy} / WATCH {res.statistics.watch} / "
|
||||
f"SELL {res.statistics.sell}(前 8 条)"
|
||||
]
|
||||
for e in res.events[:8]:
|
||||
out.append(f" {e.signal_type} {e.symbol} score={e.score:.4f} — {e.trigger_reason[0] if e.trigger_reason else ''}")
|
||||
return "\n".join(out)
|
||||
|
||||
def create_strategy(args: dict) -> str:
|
||||
name = str(_pick(args, "name", ""))
|
||||
if not name:
|
||||
return "请提供 name"
|
||||
factors = [
|
||||
{"name": x, "weight": 1.0}
|
||||
for x in _split_factor_list(str(_pick(args, "factors", "momentum_60")))
|
||||
]
|
||||
if not factors:
|
||||
return "请提供至少一个 factors(逗号分隔)"
|
||||
description = str(_pick(args, "description", "") or "")
|
||||
# 选股策略只存「选股条件组合」:股票池 + 因子(+ 可选条件)。
|
||||
# top_n / rebalance 等回测执行参数已移到「回测组合」,Agent 不再在此指定。
|
||||
st = SelectionStrategy(
|
||||
name=name,
|
||||
description=description,
|
||||
universe=UniverseSpec(
|
||||
exclude_st=bool(_pick(args, "exclude_st", True)), min_listing_days=0
|
||||
),
|
||||
factors=factors,
|
||||
)
|
||||
from app.application.services.job_executor import new_id
|
||||
|
||||
with session_factory() as session:
|
||||
saved = SqlAlchemyStrategyRepository(session).save(
|
||||
st.model_copy(update={"id": new_id("STG")})
|
||||
)
|
||||
session.commit()
|
||||
return f"策略已保存:{saved.id} {saved.name}(factors={[f.name for f in saved.factors]})"
|
||||
|
||||
def inspect_factor(args: dict) -> str:
|
||||
name = str(_pick(args, "name", ""))
|
||||
# 先问引擎:目录里有没有这行是「管理」问题,引擎算不算得出来才是「能不能用」。
|
||||
# 参数化因子(momentum(window=90,direction=…))经常还没进目录就被引用,也能算。
|
||||
try:
|
||||
defn, _fn = get_factor(name)
|
||||
except FactorError as exc:
|
||||
return f"因子不可用:{exc}"
|
||||
with session_factory() as session:
|
||||
row = SqlAlchemyFactorRepository(session).get(name)
|
||||
params = ",".join(f"{k}={v}" for k, v in defn.params.items())
|
||||
head = f"{defn.label}({defn.name})" if defn.label else defn.name
|
||||
return (
|
||||
f"{head}:{defn.description}\n公式:{defn.formula}\n方向:"
|
||||
f"{'越高越好' if defn.direction == 'higher_is_better' else '越低越好'}"
|
||||
f"(lookback {defn.lookback},输入 {defn.requires})\n"
|
||||
f"参数:{params or '(无:内置实例名固定口径)'}\n"
|
||||
f"来源:{'代码注册表内置' if defn.source == 'builtin' else '目录里的参数化实例'}"
|
||||
f"{';在目录中已停用(仍可被引用)' if row is not None and not row.enabled else ''}\n"
|
||||
f"简介:{defn.brief}"
|
||||
)
|
||||
|
||||
def create_composite_factor(args: dict) -> str:
|
||||
name = str(_pick(args, "name", ""))
|
||||
raw = str(_pick(args, "factors", ""))
|
||||
if not name or not raw:
|
||||
return (
|
||||
"请提供 name 与 factors(格式:momentum_60:0.7,volatility_60:0.3;"
|
||||
"参数化因子写成 momentum(window=90,direction=lower_is_better):0.7)"
|
||||
)
|
||||
comps: list[CompositeComponent] = []
|
||||
for part in _split_factor_list(raw):
|
||||
fname, weight_text = _split_name_weight(part)
|
||||
weight = float(weight_text) if weight_text else 1.0
|
||||
if not fname:
|
||||
continue
|
||||
try:
|
||||
defn, _fn = get_factor(fname)
|
||||
except FactorError as exc:
|
||||
return f"无法创建:{exc}"
|
||||
comps.append(CompositeComponent(name=fname, weight=weight, direction=defn.direction))
|
||||
if not comps:
|
||||
return "未解析到任何因子组件"
|
||||
from app.application.services.job_executor import new_id
|
||||
|
||||
cf = CompositeDefinition(
|
||||
name=name, description=str(_pick(args, "description", "") or ""), components=comps
|
||||
)
|
||||
with session_factory() as session:
|
||||
saved = SqlAlchemyCompositeRepository(session).save(
|
||||
cf.model_copy(update={"id": new_id("CF")})
|
||||
)
|
||||
session.commit()
|
||||
return (
|
||||
f"组合已保存:{saved.id} {saved.name}("
|
||||
+ ", ".join(f"{c.name}:{c.weight}" for c in saved.components)
|
||||
+ ")"
|
||||
)
|
||||
|
||||
def get_backtest_result(args: dict) -> str:
|
||||
exp_id = str(_pick(args, "experiment_id", "")).upper()
|
||||
with session_factory() as session:
|
||||
exp = exp_repo_f(session).get(exp_id)
|
||||
if exp is None:
|
||||
return f"Experiment {exp_id} 不存在"
|
||||
try:
|
||||
result = BacktestResult.model_validate_json(exp.result_json)
|
||||
except Exception: # noqa: BLE001
|
||||
return f"{exp_id} 不是回测结果"
|
||||
sm = result.summary
|
||||
return (
|
||||
f"回测 {exp_id} {sm.start}~{sm.end}:总收益 {sm.total_return_pct:.2f}%,"
|
||||
f"年化 {sm.annual_return_pct:.2f}%,Sharpe {sm.sharpe:.2f},"
|
||||
f"最大回撤 {sm.max_drawdown_pct:.2f}%,期末 {sm.final_equity:,.0f} 元;"
|
||||
f"交易 {sm.total_trades} 笔胜率 {sm.win_rate_pct:.1f}%;"
|
||||
f"选股记录 {len(result.selection_history)} / 信号 {len(result.signal_history)} / "
|
||||
f"成交 {len(result.fills)};未建模 {len(result.unimplemented)} 项"
|
||||
)
|
||||
|
||||
def create_experiment(args: dict) -> str:
|
||||
"""把成功 Job 兜底归档为 Experiment(研究工具已自动归档;本工具用于补档)。"""
|
||||
job_id = str(_pick(args, "job_id", "")).upper()
|
||||
if not job_id:
|
||||
return "请提供 job_id"
|
||||
from app.application.services.job_executor import new_id
|
||||
from app.domain.entities.research import ExperimentRecord
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
|
||||
with session_factory() as session:
|
||||
job = SqlAlchemyJobRepository(session).get(job_id)
|
||||
if job is None:
|
||||
return f"Job {job_id} 不存在"
|
||||
# 顺序要紧:新形态记录结果只在 experiment 侧(job.result_json 为 None),
|
||||
# 先判「已归档」,否则成功 Job 会被误判为「无结果可归档」
|
||||
if job.experiment_id:
|
||||
return f"Job {job_id} 已归档为 {job.experiment_id}"
|
||||
if job.status != "success" or not job.result_json:
|
||||
return f"Job {job_id} 未成功(无结果可归档)"
|
||||
# 走到这里必然是老形态记录(结果仍在 job 侧,无 experiment 关联)
|
||||
result_json = job.result_json
|
||||
exp_repo = exp_repo_f(session)
|
||||
summary = None
|
||||
try:
|
||||
if job.kind == "backtest":
|
||||
r = BacktestResult.model_validate_json(result_json)
|
||||
summary = (
|
||||
f"总收益 {r.summary.total_return_pct:.2f}% · 年化 "
|
||||
f"{r.summary.annual_return_pct:.2f}% · 回撤 {r.summary.max_drawdown_pct:.2f}%"
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
exp = ExperimentRecord(
|
||||
id=new_id("EXP"),
|
||||
kind=job.kind,
|
||||
spec_json=job.spec_json,
|
||||
result_json=result_json,
|
||||
summary_text=summary,
|
||||
job_id=job.id,
|
||||
created_at=job.created_at,
|
||||
)
|
||||
exp_repo.save(exp)
|
||||
job.experiment_id = exp.id
|
||||
SqlAlchemyJobRepository(session).update(job)
|
||||
session.commit()
|
||||
return f"已归档:{exp.id}(Job {job_id} → Experiment)"
|
||||
|
||||
return [
|
||||
Tool(
|
||||
"search_stocks",
|
||||
"按代码或名称搜索股票,返回基础信息(只读)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"q": {"type": "string", "description": "代码或名称关键字"}},
|
||||
},
|
||||
search_stocks,
|
||||
),
|
||||
Tool(
|
||||
"get_market_data",
|
||||
"读取一只股票一段区间的日线行情摘要(只读,不复权)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"symbol": {"type": "string", "description": "如 600519.SH"},
|
||||
"start": {"type": "string", "description": "YYYY-MM-DD"},
|
||||
"end": {"type": "string", "description": "YYYY-MM-DD"},
|
||||
},
|
||||
"required": ["symbol"],
|
||||
},
|
||||
get_market_data,
|
||||
),
|
||||
Tool(
|
||||
"test_factor",
|
||||
"对单个因子做 IC/RankIC/分层测试并归档 Experiment",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "因子名(momentum_60 / volatility_20 等)",
|
||||
},
|
||||
"start": {"type": "string"},
|
||||
"end": {"type": "string"},
|
||||
},
|
||||
"required": ["name"],
|
||||
},
|
||||
test_factor,
|
||||
),
|
||||
Tool(
|
||||
"run_backtest",
|
||||
"运行 TopK 低频回测并归档 Experiment(成本/涨跌停近似建模)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"factors": {"type": "string", "description": "逗号分隔的因子名"},
|
||||
"top_n": {"type": "integer"},
|
||||
"rebalance": {"type": "string", "enum": ["monthly", "weekly"]},
|
||||
"exclude_st": {"type": "boolean"},
|
||||
"start": {"type": "string"},
|
||||
"end": {"type": "string"},
|
||||
},
|
||||
},
|
||||
run_backtest,
|
||||
),
|
||||
Tool(
|
||||
"get_experiment",
|
||||
"读取已归档实验的摘要",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"experiment_id": {"type": "string"}},
|
||||
"required": ["experiment_id"],
|
||||
},
|
||||
get_experiment,
|
||||
),
|
||||
Tool(
|
||||
"compare_experiments",
|
||||
"对比多个实验(因子/区间/收益摘要)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"experiment_ids": {"type": "string"}},
|
||||
"required": ["experiment_ids"],
|
||||
},
|
||||
compare_experiments,
|
||||
),
|
||||
Tool(
|
||||
"screen_stocks",
|
||||
"按因子评分筛选股票(TopN;传 symbols 白名单避免全市场长任务)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"factors": {"type": "string", "description": "逗号分隔因子名"},
|
||||
"top_n": {"type": "integer"},
|
||||
"as_of": {"type": "string", "description": "YYYY-MM-DD"},
|
||||
"symbols": {"type": "string", "description": "逗号分隔白名单(可选,≤60)"},
|
||||
},
|
||||
},
|
||||
screen_stocks,
|
||||
),
|
||||
Tool(
|
||||
"explain_selection",
|
||||
"解释一次选股结果:为什么选这些股票(含因子值与理由)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"selection_id": {"type": "string"}},
|
||||
"required": ["selection_id"],
|
||||
},
|
||||
explain_selection,
|
||||
),
|
||||
Tool(
|
||||
"generate_signals",
|
||||
"基于选股评分+趋势生成 BUY/WATCH/SELL 信号",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"factors": {"type": "string"},
|
||||
"as_of": {"type": "string"},
|
||||
"symbols": {"type": "string", "description": "白名单(可选)"},
|
||||
"buy_rank": {"type": "integer"},
|
||||
"sell_rank": {"type": "integer"},
|
||||
},
|
||||
},
|
||||
generate_signals,
|
||||
),
|
||||
Tool(
|
||||
"create_strategy",
|
||||
"创建/保存一个命名策略(可随后展开为回测)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"description": {"type": "string"},
|
||||
"factors": {"type": "string"},
|
||||
"top_n": {"type": "integer"},
|
||||
"rebalance": {"type": "string", "enum": ["monthly", "weekly"]},
|
||||
},
|
||||
"required": ["name", "factors"],
|
||||
},
|
||||
create_strategy,
|
||||
),
|
||||
Tool(
|
||||
"inspect_factor",
|
||||
"查看因子目录元数据(公式/方向/lookback/输入列)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"name": {"type": "string"}},
|
||||
"required": ["name"],
|
||||
},
|
||||
inspect_factor,
|
||||
),
|
||||
Tool(
|
||||
"create_composite_factor",
|
||||
"创建并保存多因子组合(factors 格式:momentum_60:0.7,volatility_60:0.3)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"factors": {"type": "string"},
|
||||
"description": {"type": "string"},
|
||||
},
|
||||
"required": ["name", "factors"],
|
||||
},
|
||||
create_composite_factor,
|
||||
),
|
||||
Tool(
|
||||
"get_backtest_result",
|
||||
"读取回测 Experiment 的详细结果(收益/回撤/交易/意图与成交统计)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"experiment_id": {"type": "string"}},
|
||||
"required": ["experiment_id"],
|
||||
},
|
||||
get_backtest_result,
|
||||
),
|
||||
Tool(
|
||||
"create_experiment",
|
||||
"把成功 Job 兜底归档为 Experiment(补档;研究工具已自动归档)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"job_id": {"type": "string"}},
|
||||
"required": ["job_id"],
|
||||
},
|
||||
create_experiment,
|
||||
),
|
||||
]
|
||||
|
||||
def _code_version(job, exp_repo_f) -> str:
|
||||
try:
|
||||
with default_factories()["session_factory"]() as session:
|
||||
exp = exp_repo_f(session).get(job.experiment_id or "")
|
||||
return exp.code_version or "-" if exp else "-"
|
||||
except Exception: # noqa: BLE001
|
||||
return "-"
|
||||
@@ -0,0 +1,40 @@
|
||||
"""AI Research Agent API(Phase 5)。
|
||||
|
||||
POST /api/agent/chat {message} → {reply, actions:[{tool,args,output}]}
|
||||
未配置 LLM Key 时返回 400 引导(不会崩溃)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.llm import LLMClient, build_llm_from_settings
|
||||
from app.agent.service import AgentService
|
||||
|
||||
router = APIRouter(prefix="/agent", tags=["agent"])
|
||||
|
||||
|
||||
class AgentChatRequest(BaseModel):
|
||||
message: str = Field(min_length=1, max_length=2000)
|
||||
|
||||
|
||||
def _llm_or_raise() -> LLMClient:
|
||||
llm = build_llm_from_settings()
|
||||
if llm is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="未配置 LLM:请在根目录 .env 中设置 LLM_API_KEY(可选 LLM_BASE_URL / LLM_MODEL),"
|
||||
"参考 .env.example 与 AGENT.md §33",
|
||||
)
|
||||
return llm
|
||||
|
||||
|
||||
@router.post("/chat", summary="与 AI 研究助手对话(受控工具)")
|
||||
def agent_chat(
|
||||
body: AgentChatRequest,
|
||||
llm: Annotated[LLMClient, Depends(_llm_or_raise)],
|
||||
) -> dict:
|
||||
return AgentService(llm).chat(body.message)
|
||||
@@ -0,0 +1,130 @@
|
||||
"""Chart API(v3 §20.2):统一可视化数据只读接口。
|
||||
|
||||
- GET /api/stocks/{symbol}/chart?start&end&adjust K 线 + 量 + 指标(显示层折算)
|
||||
- GET /api/stocks/{symbol}/signals 该股历史信号(markers)
|
||||
- GET /api/stocks/{symbol}/selections 该股历史选股命中(markers)
|
||||
- GET /api/backtests/{experiment_id}/stocks/{symbol}/chart 回测个股:K 线 + 实际成交 fills
|
||||
- GET /api/backtests/{experiment_id}/trades|positions 回测成交/持仓展开
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Query
|
||||
|
||||
from app.api.deps import (
|
||||
ChartServiceDep,
|
||||
ExperimentRepoDep,
|
||||
SelectionRepoDep,
|
||||
SignalRepoDep,
|
||||
)
|
||||
from app.domain.entities.chart import ChartResult, EventMarker, SelectionHit
|
||||
from app.domain.entities.research import BacktestResult
|
||||
from app.domain.entities.signal import SignalHit
|
||||
|
||||
router = APIRouter(tags=["charts"])
|
||||
|
||||
_AdjustQuery = Annotated[str, Query(pattern="^(none|qfq|hfq)$")]
|
||||
_StartQuery = Annotated[date | None, Query(description="开始日期(默认 2024-01-01)")]
|
||||
_EndQuery = Annotated[date | None, Query(description="结束日期(默认今天)")]
|
||||
|
||||
|
||||
def _signal_markers(hits: list[SignalHit]) -> list[EventMarker]:
|
||||
kind_map = {"BUY": "signal_buy", "SELL": "signal_sell", "WATCH": "signal_watch"}
|
||||
return [
|
||||
EventMarker(
|
||||
time=h.signal_date,
|
||||
kind=kind_map.get(h.signal_type, "signal_watch"),
|
||||
symbol="",
|
||||
price=h.price,
|
||||
score=h.score,
|
||||
text=h.trigger_reason,
|
||||
ref_id=h.signal_id,
|
||||
)
|
||||
for h in hits
|
||||
]
|
||||
|
||||
|
||||
def _selection_markers(hits: list[SelectionHit]) -> list[EventMarker]:
|
||||
return [
|
||||
EventMarker(
|
||||
time=h.as_of,
|
||||
kind="selection",
|
||||
symbol=h.symbol,
|
||||
score=h.score,
|
||||
text=[f"rank #{h.rank}"] + h.selection_reason,
|
||||
ref_id=h.selection_id,
|
||||
)
|
||||
for h in hits
|
||||
]
|
||||
|
||||
|
||||
@router.get("/stocks/{symbol}/chart", response_model=ChartResult, summary="个股 K 线图数据")
|
||||
def stock_chart(
|
||||
symbol: str,
|
||||
service: ChartServiceDep,
|
||||
start: _StartQuery = None,
|
||||
end: _EndQuery = None,
|
||||
adjust: _AdjustQuery = "none",
|
||||
) -> ChartResult:
|
||||
start = start or date(2024, 1, 1)
|
||||
end = end or date.today()
|
||||
if service.stock(symbol) is None:
|
||||
raise HTTPException(status_code=404, detail=f"未找到股票 {symbol}")
|
||||
return service.stock_chart(symbol, start, end, adjust)
|
||||
|
||||
|
||||
@router.get("/stocks/{symbol}/signals", response_model=list[EventMarker], summary="该股历史信号")
|
||||
def symbol_signals(symbol: str, signal_repo: SignalRepoDep) -> list[EventMarker]:
|
||||
return _signal_markers(signal_repo.list_by_symbol(symbol))
|
||||
|
||||
|
||||
@router.get("/stocks/{symbol}/selections", response_model=list[EventMarker], summary="该股历史选股命中")
|
||||
def symbol_selections(symbol: str, selection_repo: SelectionRepoDep) -> list[EventMarker]:
|
||||
return _selection_markers(selection_repo.list_by_symbol(symbol))
|
||||
|
||||
|
||||
@router.get(
|
||||
"/backtests/{experiment_id}/stocks/{symbol}/chart",
|
||||
response_model=ChartResult,
|
||||
summary="回测个股图(K 线 + 实际成交 fills)",
|
||||
)
|
||||
def backtest_stock_chart(
|
||||
experiment_id: str,
|
||||
symbol: str,
|
||||
service: ChartServiceDep,
|
||||
experiment_repo: ExperimentRepoDep,
|
||||
start: _StartQuery = None,
|
||||
end: _EndQuery = None,
|
||||
adjust: _AdjustQuery = "none",
|
||||
) -> ChartResult:
|
||||
start = start or date(2024, 1, 1)
|
||||
end = end or date.today()
|
||||
exp = experiment_repo.get(experiment_id)
|
||||
if exp is None:
|
||||
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
|
||||
try:
|
||||
result = BacktestResult.model_validate_json(exp.result_json)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
raise HTTPException(status_code=400, detail=f"{experiment_id} 不是 backtest 结果") from exc
|
||||
return service.backtest_stock_chart(result, symbol, start, end, adjust)
|
||||
|
||||
|
||||
@router.get("/backtests/{experiment_id}/trades", summary="回测成交明细")
|
||||
def backtest_trades(experiment_id: str, experiment_repo: ExperimentRepoDep):
|
||||
exp = experiment_repo.get(experiment_id)
|
||||
if exp is None:
|
||||
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
|
||||
result = BacktestResult.model_validate_json(exp.result_json)
|
||||
return result.trades
|
||||
|
||||
|
||||
@router.get("/backtests/{experiment_id}/positions", summary="回测持仓明细")
|
||||
def backtest_positions(experiment_id: str, experiment_repo: ExperimentRepoDep):
|
||||
exp = experiment_repo.get(experiment_id)
|
||||
if exp is None:
|
||||
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
|
||||
result = BacktestResult.model_validate_json(exp.result_json)
|
||||
return result.positions
|
||||
@@ -0,0 +1,142 @@
|
||||
"""回测组合 API:/api/combos(CRUD + 运行)。
|
||||
|
||||
POST /api/combos 保存组合(name 唯一)
|
||||
GET /api/combos 列表
|
||||
GET /api/combos/{id} 详情
|
||||
PUT /api/combos/{id} 原地更新
|
||||
DELETE /api/combos/{id} 删除
|
||||
POST /api/combos/{id}/run 提交已保存组合为异步 Job(kind=combo)
|
||||
POST /api/combos/run 提交临时组合(不保存)为异步 Job
|
||||
|
||||
运行时:按 combo.strategy_ids 取齐选股策略 + 读公共配置 → ComboService.run → 归档。
|
||||
费率/复权来自公共配置,快照进归档 config_snapshot(可复现,AGENT.md §21)。
|
||||
|
||||
依赖注入说明:建 Job 与策略校验都走 FastAPI 注入的 session/repo(而非直接 SessionLocal),
|
||||
这样测试用 dependency_overrides 替换数据库时也能命中同一份库,行为一致可测。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, HTTPException
|
||||
|
||||
from app.api.deps import ComboRepoDep, DbSession, StrategyRepoDep
|
||||
from app.application.services.job_executor import new_id, run_job_background
|
||||
from app.domain.entities.combo import BacktestCombo
|
||||
from app.domain.entities.research import JobRecord, JobStatus
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/combos", tags=["combos"])
|
||||
|
||||
|
||||
def _ensure_strategies_exist(combo: BacktestCombo, strategy_repo) -> None:
|
||||
"""引用的选股策略必须都存在;缺任何一个即 400(提前失败,不等后台 Job 才暴露)。"""
|
||||
for sid in combo.strategy_ids:
|
||||
if strategy_repo.get(sid) is None:
|
||||
raise HTTPException(status_code=400, detail=f"组合引用的选股策略 {sid} 不存在")
|
||||
|
||||
|
||||
@router.post("", response_model=BacktestCombo, summary="保存回测组合")
|
||||
def create_combo(
|
||||
combo: BacktestCombo,
|
||||
repo: ComboRepoDep,
|
||||
session: DbSession,
|
||||
) -> BacktestCombo:
|
||||
try:
|
||||
saved = repo.save(combo.model_copy(update={"id": new_id("CMB")}))
|
||||
session.commit()
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
return repo.get(saved.id) or saved
|
||||
|
||||
|
||||
@router.get("", response_model=list[BacktestCombo], summary="回测组合列表")
|
||||
def list_combos(repo: ComboRepoDep) -> list[BacktestCombo]:
|
||||
return repo.list()
|
||||
|
||||
|
||||
@router.get("/{combo_id}", response_model=BacktestCombo, summary="读取回测组合")
|
||||
def get_combo(combo_id: str, repo: ComboRepoDep) -> BacktestCombo:
|
||||
row = repo.get(combo_id)
|
||||
if row is None:
|
||||
raise HTTPException(status_code=404, detail=f"回测组合 {combo_id} 不存在")
|
||||
return row
|
||||
|
||||
|
||||
@router.put("/{combo_id}", response_model=BacktestCombo, summary="原地更新回测组合")
|
||||
def update_combo(
|
||||
combo_id: str,
|
||||
combo: BacktestCombo,
|
||||
repo: ComboRepoDep,
|
||||
session: DbSession,
|
||||
) -> BacktestCombo:
|
||||
existing = repo.get(combo_id)
|
||||
if existing is None:
|
||||
raise HTTPException(status_code=404, detail=f"回测组合 {combo_id} 不存在")
|
||||
payload = combo.model_copy(update={"id": combo_id, "created_at": existing.created_at})
|
||||
try:
|
||||
saved = repo.save(payload)
|
||||
session.commit()
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
return repo.get(saved.id) or saved
|
||||
|
||||
|
||||
@router.delete("/{combo_id}", summary="删除回测组合")
|
||||
def delete_combo(combo_id: str, repo: ComboRepoDep, session: DbSession) -> dict:
|
||||
if not repo.delete(combo_id):
|
||||
raise HTTPException(status_code=404, detail=f"回测组合 {combo_id} 不存在")
|
||||
session.commit()
|
||||
return {"deleted": combo_id}
|
||||
|
||||
|
||||
def _submit_combo_job(
|
||||
combo: BacktestCombo,
|
||||
*,
|
||||
strategy_repo,
|
||||
session,
|
||||
background: BackgroundTasks,
|
||||
) -> dict:
|
||||
"""校验策略 → 建 Job(kind=combo)→ 入后台执行。
|
||||
|
||||
spec_json 存 BacktestCombo JSON;执行端(job_executor)识别 kind="combo",
|
||||
再按 strategy_ids 取策略 + 读公共配置后调 ComboService。Job 表只存组合本身,
|
||||
策略/配置的「当时快照」由 ComboService 写进归档 config_snapshot(可复现)。
|
||||
校验与建 Job 都用注入的 session/repo,保证与测试覆写一致。
|
||||
"""
|
||||
_ensure_strategies_exist(combo, strategy_repo)
|
||||
job = JobRecord(
|
||||
id=new_id("JOB"), kind="combo", status=JobStatus.QUEUED,
|
||||
spec_json=combo.model_dump_json(), created_at=datetime.now(),
|
||||
)
|
||||
SqlAlchemyJobRepository(session).create(job)
|
||||
session.commit()
|
||||
background.add_task(run_job_background, job.id)
|
||||
return {"job_id": job.id, "status": job.status}
|
||||
|
||||
|
||||
@router.post("/{combo_id}/run", summary="运行已保存的回测组合(异步 Job)")
|
||||
def run_saved_combo(
|
||||
combo_id: str,
|
||||
repo: ComboRepoDep,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
background: BackgroundTasks,
|
||||
) -> dict:
|
||||
combo = repo.get(combo_id)
|
||||
if combo is None:
|
||||
raise HTTPException(status_code=404, detail=f"回测组合 {combo_id} 不存在")
|
||||
return _submit_combo_job(combo, strategy_repo=strategy_repo, session=session, background=background)
|
||||
|
||||
|
||||
@router.post("/run", summary="运行临时回测组合(不保存,异步 Job)")
|
||||
def run_adhoc_combo(
|
||||
combo: BacktestCombo,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
background: BackgroundTasks,
|
||||
) -> dict:
|
||||
return _submit_combo_job(combo, strategy_repo=strategy_repo, session=session, background=background)
|
||||
@@ -0,0 +1,69 @@
|
||||
"""因子组合 API(M7.2b):/api/composites CRUD。
|
||||
|
||||
组合 = 可复用因子权重集;计算时方向取因子注册表,落库冗余快照。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
|
||||
from app.api.deps import CompositeRepoDep, DbSession
|
||||
from app.application.services.job_executor import new_id
|
||||
from app.domain.entities.composite import CompositeComponent, CompositeDefinition
|
||||
from app.quant.factors import FactorError, get_factor
|
||||
|
||||
router = APIRouter(prefix="/composites", tags=["composites"])
|
||||
|
||||
|
||||
def _fill_direction(definition: CompositeDefinition) -> CompositeDefinition:
|
||||
"""以因子注册表元数据补齐/校正组件 direction(登记但不可计算的因子报错)。"""
|
||||
out = []
|
||||
for c in definition.components:
|
||||
try:
|
||||
defn, _fn = get_factor(c.name)
|
||||
except FactorError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
out.append(CompositeComponent(name=c.name, weight=c.weight, direction=defn.direction))
|
||||
return definition.model_copy(update={"components": out})
|
||||
|
||||
|
||||
@router.post("", summary="保存因子组合", response_model=CompositeDefinition)
|
||||
def create_composite(
|
||||
definition: CompositeDefinition,
|
||||
composite_repo: CompositeRepoDep,
|
||||
session: DbSession,
|
||||
) -> CompositeDefinition:
|
||||
prepared = _fill_direction(definition.model_copy(update={"id": ""}))
|
||||
try:
|
||||
saved = composite_repo.save(
|
||||
prepared.model_copy(update={"id": new_id("CF")})
|
||||
)
|
||||
session.commit()
|
||||
return saved
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.get("", summary="因子组合列表", response_model=list[CompositeDefinition])
|
||||
def list_composites(composite_repo: CompositeRepoDep) -> list[CompositeDefinition]:
|
||||
return composite_repo.list()
|
||||
|
||||
|
||||
@router.get("/{composite_id}", summary="读取因子组合", response_model=CompositeDefinition)
|
||||
def get_composite(composite_id: str, composite_repo: CompositeRepoDep) -> CompositeDefinition:
|
||||
row = composite_repo.get(composite_id)
|
||||
if row is None:
|
||||
raise HTTPException(status_code=404, detail=f"组合 {composite_id} 不存在")
|
||||
return row
|
||||
|
||||
|
||||
@router.delete("/{composite_id}", summary="删除因子组合")
|
||||
def delete_composite(
|
||||
composite_id: str,
|
||||
composite_repo: CompositeRepoDep,
|
||||
session: DbSession,
|
||||
) -> dict:
|
||||
if not composite_repo.delete(composite_id):
|
||||
raise HTTPException(status_code=404, detail=f"组合 {composite_id} 不存在")
|
||||
session.commit()
|
||||
return {"deleted": composite_id}
|
||||
@@ -0,0 +1,177 @@
|
||||
"""字段库 API:/api/condition-fields(2026-10)。
|
||||
|
||||
GET /api/condition-fields 字段库列表(首次读取自动 seed 内置字段)
|
||||
GET /api/condition-fields/available 引擎支持但尚未进库的字段(「新增」可选项)
|
||||
POST /api/condition-fields 新增自定义字段(必须指向引擎真能算的字段)
|
||||
PUT /api/condition-fields/{name} 改中文名/含义/分组/单位/启用状态
|
||||
DELETE /api/condition-fields/{name} 删除自定义字段(内置字段只能停用)
|
||||
|
||||
设计要点(AGENT.md §24 不假装支持):字段能不能算由 `quant.condition_fields` 对着引擎域
|
||||
判定;库里登记不出来的字段会被拒绝(422),否则用户会建出「永远选不出股票」的空策略。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from app.api.deps import ConditionFieldRepoDep, DbSession
|
||||
from app.application.services import condition_field_catalog as svc
|
||||
from app.domain.entities.condition_field import ConditionField
|
||||
from app.quant.condition_fields import FieldDef, get_field, unit_options
|
||||
|
||||
router = APIRouter(prefix="/condition-fields", tags=["condition-fields"])
|
||||
|
||||
|
||||
class UnitOption(BaseModel):
|
||||
"""可选**界面单位** + 它到**基准单位**的换算系数(提交前 ×factor,回显时 ÷factor)。
|
||||
|
||||
``factor`` 由注册表给出(如 总市值:万元=1、亿元=10000),前端据此换算。
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
unit: str
|
||||
factor: float
|
||||
|
||||
|
||||
class ConditionFieldOut(ConditionField):
|
||||
"""字段库响应:目录字段 + **从注册表派生**的单位信息。
|
||||
|
||||
``base_unit`` / ``units`` 不落库:它们是引擎口径(注册表)的投影,若存进表里就会
|
||||
和代码漂移。``unit`` 才是库里存的那一项(当前界面单位,必须落在 ``units`` 里)。
|
||||
"""
|
||||
|
||||
base_unit: str = ""
|
||||
units: list[UnitOption] = Field(default_factory=list)
|
||||
|
||||
@classmethod
|
||||
def from_item(cls, item: ConditionField) -> ConditionFieldOut:
|
||||
d = get_field(item.name)
|
||||
return cls(
|
||||
**item.model_dump(exclude={"ops"}), # ops 是 computed_field,不参与构造
|
||||
base_unit=(d.unit if d else item.unit),
|
||||
units=[UnitOption(unit=u, factor=f) for u, f in unit_options(item.name)],
|
||||
)
|
||||
|
||||
|
||||
class FieldOption(BaseModel):
|
||||
"""「可新增字段」的建议项:来自代码注册表,附带默认中文名与含义。"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
name: str
|
||||
label: str
|
||||
description: str
|
||||
kind: str
|
||||
group_name: str
|
||||
unit: str = ""
|
||||
units: list[UnitOption] = Field(default_factory=list)
|
||||
|
||||
@classmethod
|
||||
def from_def(cls, d: FieldDef) -> FieldOption:
|
||||
return cls(
|
||||
name=d.name,
|
||||
label=d.label,
|
||||
description=d.description,
|
||||
kind=d.kind,
|
||||
group_name=d.group_name,
|
||||
unit=d.unit,
|
||||
units=[UnitOption(unit=u, factor=f) for u, f in d.unit_options],
|
||||
)
|
||||
|
||||
|
||||
class ConditionFieldCreate(BaseModel):
|
||||
"""新增请求。kind 不收:类型是引擎事实,由服务端按注册表填。"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
name: str = Field(min_length=1, max_length=64)
|
||||
label: str = Field(default="", max_length=64)
|
||||
description: str = Field(default="", max_length=500)
|
||||
group_name: str = Field(default="", max_length=32)
|
||||
unit: str = Field(default="", max_length=16)
|
||||
enabled: bool = True
|
||||
|
||||
|
||||
class ConditionFieldUpdate(BaseModel):
|
||||
"""编辑请求。name/kind/source 有意不可改(name 是引擎字段名,改了就换字段了)。"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
label: str | None = Field(default=None, max_length=64)
|
||||
description: str | None = Field(default=None, max_length=500)
|
||||
group_name: str | None = Field(default=None, max_length=32)
|
||||
unit: str | None = Field(default=None, max_length=16)
|
||||
enabled: bool | None = None
|
||||
|
||||
|
||||
@router.get("", response_model=list[ConditionFieldOut], summary="字段库列表")
|
||||
def list_condition_fields(
|
||||
repo: ConditionFieldRepoDep, session: DbSession, include_disabled: bool = True
|
||||
) -> list[ConditionFieldOut]:
|
||||
"""读字段库;缺失的内置字段当场补齐(幂等,稳态零写入)。
|
||||
|
||||
`include_disabled=false` 供条件编辑器使用(只列启用项);字段库管理页用默认值
|
||||
(列出全部,含停用项,否则用户没法把停用的字段再打开)。
|
||||
"""
|
||||
return [ConditionFieldOut.from_item(f) for f in svc.list_fields(repo, session, include_disabled=include_disabled)]
|
||||
|
||||
|
||||
@router.get("/available", response_model=list[FieldOption], summary="可新增的字段")
|
||||
def list_available_fields(repo: ConditionFieldRepoDep, session: DbSession) -> list[FieldOption]:
|
||||
return [FieldOption.from_def(d) for d in svc.list_available(repo, session)]
|
||||
|
||||
|
||||
@router.post("", response_model=ConditionFieldOut, summary="新增自定义字段")
|
||||
def create_condition_field(
|
||||
body: ConditionFieldCreate, repo: ConditionFieldRepoDep, session: DbSession
|
||||
) -> ConditionFieldOut:
|
||||
try:
|
||||
saved = svc.create_field(
|
||||
repo,
|
||||
session,
|
||||
name=body.name,
|
||||
label=body.label,
|
||||
description=body.description,
|
||||
group_name=body.group_name,
|
||||
unit=body.unit,
|
||||
enabled=body.enabled,
|
||||
)
|
||||
except ValueError as exc:
|
||||
# 422:语义是「引擎算不出来 / 已在库里 / 单位不在可选范围」,属于请求内容不可处理
|
||||
raise HTTPException(status_code=422, detail=str(exc)) from exc
|
||||
return ConditionFieldOut.from_item(saved)
|
||||
|
||||
|
||||
@router.put("/{name}", response_model=ConditionFieldOut, summary="编辑字段(中文名/含义/单位/启用)")
|
||||
def update_condition_field(
|
||||
name: str, body: ConditionFieldUpdate, repo: ConditionFieldRepoDep, session: DbSession
|
||||
) -> ConditionFieldOut:
|
||||
try:
|
||||
saved = svc.update_field(
|
||||
repo,
|
||||
session,
|
||||
name,
|
||||
label=body.label,
|
||||
description=body.description,
|
||||
group_name=body.group_name,
|
||||
unit=body.unit,
|
||||
enabled=body.enabled,
|
||||
)
|
||||
except KeyError as exc:
|
||||
raise HTTPException(status_code=404, detail=f"字段 {name} 不存在") from exc
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=422, detail=str(exc)) from exc
|
||||
return ConditionFieldOut.from_item(saved)
|
||||
|
||||
|
||||
@router.delete("/{name}", summary="删除自定义字段(内置只能停用)")
|
||||
def delete_condition_field(name: str, repo: ConditionFieldRepoDep, session: DbSession) -> dict:
|
||||
try:
|
||||
svc.delete_field(repo, session, name)
|
||||
except KeyError as exc:
|
||||
raise HTTPException(status_code=404, detail=f"字段 {name} 不存在") from exc
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
return {"deleted": name}
|
||||
@@ -0,0 +1,30 @@
|
||||
"""公共配置 API:/api/config(全局唯一一份费率/滑点/复权口径)。
|
||||
|
||||
GET /api/config 读取(未配置过返回带默认值的实例)
|
||||
PUT /api/config 更新(upsert 单例)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api.deps import DbSession, GlobalConfigRepoDep
|
||||
from app.domain.entities.combo import GlobalConfig
|
||||
|
||||
router = APIRouter(prefix="/config", tags=["config"])
|
||||
|
||||
|
||||
@router.get("", response_model=GlobalConfig, summary="读取公共配置")
|
||||
def get_config(repo: GlobalConfigRepoDep) -> GlobalConfig:
|
||||
return repo.get()
|
||||
|
||||
|
||||
@router.put("", response_model=GlobalConfig, summary="更新公共配置")
|
||||
def update_config(
|
||||
config: GlobalConfig,
|
||||
repo: GlobalConfigRepoDep,
|
||||
session: DbSession,
|
||||
) -> GlobalConfig:
|
||||
saved = repo.save(config)
|
||||
session.commit()
|
||||
return saved
|
||||
@@ -0,0 +1,240 @@
|
||||
"""API 依赖注入:Repository / 研究服务的装配点(composition root 的一部分)。
|
||||
|
||||
路由层统一使用 Annotated 注入(FastAPI 推荐写法,配合 ruff B008 无冲突)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import Depends
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.application.services.chart_service import ChartService
|
||||
from app.application.services.replay_service import ReplayService
|
||||
from app.application.services.selection_service import SelectionService
|
||||
from app.application.services.signal_service import SignalService
|
||||
from app.domain.repositories.combo import ComboRepository, GlobalConfigRepository
|
||||
from app.domain.repositories.composite import CompositeRepository
|
||||
from app.domain.repositories.condition_field import ConditionFieldRepository
|
||||
from app.domain.repositories.factor import FactorRepository
|
||||
from app.domain.repositories.index import IndexConstituentRepository
|
||||
from app.domain.repositories.jobs import ExperimentRepository, JobRepository
|
||||
from app.domain.repositories.market import (
|
||||
AdjustFactorRepository,
|
||||
DailyBarRepository,
|
||||
DailyBasicRepository,
|
||||
FinancialRepository,
|
||||
StockNameHistoryRepository,
|
||||
StockRepository,
|
||||
)
|
||||
from app.domain.repositories.selection import SelectionRepository
|
||||
from app.domain.repositories.signal import SignalRepository
|
||||
from app.domain.repositories.strategy import StrategyRepository
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.combo_impl import (
|
||||
SqlAlchemyComboRepository,
|
||||
SqlAlchemyGlobalConfigRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import (
|
||||
SqlAlchemyCompositeRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.condition_field_impl import (
|
||||
SqlAlchemyConditionFieldRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
|
||||
SqlAlchemyFactorRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.index_impl import (
|
||||
SqlAlchemyIndexConstituentRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyAdjustFactorRepository,
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyDailyBasicRepository,
|
||||
SqlAlchemyFinancialRepository,
|
||||
SqlAlchemyStockNameHistoryRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl import (
|
||||
SqlAlchemySelectionRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.signal_impl import (
|
||||
SqlAlchemySignalRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
|
||||
SqlAlchemyStrategyRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.session import get_session
|
||||
from app.quant.engine import LocalEngine, QuantEngine
|
||||
from app.quant.service import ResearchService
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_session)]
|
||||
|
||||
|
||||
def _stock_repo_factory(session: DbSession) -> StockRepository:
|
||||
return SqlAlchemyStockRepository(session)
|
||||
|
||||
|
||||
def _daily_repo_factory(session: DbSession) -> DailyBarRepository:
|
||||
return SqlAlchemyDailyBarRepository(session)
|
||||
|
||||
|
||||
def _financial_repo_factory(session: DbSession) -> FinancialRepository:
|
||||
return SqlAlchemyFinancialRepository(session)
|
||||
|
||||
|
||||
def _adjust_repo_factory(session: DbSession) -> AdjustFactorRepository:
|
||||
return SqlAlchemyAdjustFactorRepository(session)
|
||||
|
||||
|
||||
def _daily_basic_repo_factory(session: DbSession) -> DailyBasicRepository:
|
||||
return SqlAlchemyDailyBasicRepository(session)
|
||||
|
||||
|
||||
def _chart_service_factory(
|
||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||
adj_repo: Annotated[AdjustFactorRepository, Depends(_adjust_repo_factory)],
|
||||
) -> ChartService:
|
||||
return ChartService(stock_repo, daily_repo, adj_repo)
|
||||
|
||||
|
||||
def _name_repo_factory(session: DbSession) -> StockNameHistoryRepository:
|
||||
"""名称变更历史仓储(StockNameHistoryRepository 实现)。"""
|
||||
return SqlAlchemyStockNameHistoryRepository(session)
|
||||
|
||||
|
||||
def _index_repo_factory(session: DbSession) -> IndexConstituentRepository:
|
||||
return SqlAlchemyIndexConstituentRepository(session)
|
||||
|
||||
|
||||
def _engine_factory() -> QuantEngine:
|
||||
return LocalEngine()
|
||||
|
||||
|
||||
def _service_factory(
|
||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||
engine: Annotated[QuantEngine, Depends(_engine_factory)],
|
||||
index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)],
|
||||
basic_repo: Annotated[DailyBasicRepository, Depends(_daily_basic_repo_factory)],
|
||||
financial_repo: Annotated[FinancialRepository, Depends(_financial_repo_factory)],
|
||||
name_repo: Annotated[
|
||||
StockNameHistoryRepository, Depends(_name_repo_factory)
|
||||
] = None,
|
||||
) -> ResearchService:
|
||||
return ResearchService(
|
||||
stock_repo,
|
||||
daily_repo,
|
||||
engine,
|
||||
index_repo,
|
||||
basic_repo=basic_repo,
|
||||
financial_repo=financial_repo,
|
||||
name_repo=name_repo,
|
||||
)
|
||||
|
||||
|
||||
def _replay_service_factory(
|
||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||
index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)],
|
||||
name_repo: Annotated[StockNameHistoryRepository, Depends(_name_repo_factory)],
|
||||
) -> ReplayService:
|
||||
return ReplayService(stock_repo, daily_repo, index_repo, name_repo=name_repo)
|
||||
|
||||
|
||||
def _signal_service_factory(
|
||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||
index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)],
|
||||
name_repo: Annotated[StockNameHistoryRepository, Depends(_name_repo_factory)],
|
||||
) -> SignalService:
|
||||
return SignalService(stock_repo, daily_repo, index_repo, name_repo=name_repo)
|
||||
|
||||
|
||||
def _selection_service_factory(
|
||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||
financial_repo: Annotated[FinancialRepository, Depends(_financial_repo_factory)],
|
||||
index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)],
|
||||
basic_repo: Annotated[DailyBasicRepository, Depends(_daily_basic_repo_factory)],
|
||||
name_repo: Annotated[StockNameHistoryRepository, Depends(_name_repo_factory)],
|
||||
) -> SelectionService:
|
||||
return SelectionService(
|
||||
stock_repo,
|
||||
daily_repo,
|
||||
financial_repo,
|
||||
index_repo,
|
||||
basic_repo=basic_repo,
|
||||
name_repo=name_repo,
|
||||
)
|
||||
|
||||
|
||||
def _selection_repo_factory(session: DbSession) -> SelectionRepository:
|
||||
return SqlAlchemySelectionRepository(session)
|
||||
|
||||
|
||||
def _factor_repo_factory(session: DbSession) -> FactorRepository:
|
||||
return SqlAlchemyFactorRepository(session)
|
||||
|
||||
|
||||
def _condition_field_repo_factory(session: DbSession) -> ConditionFieldRepository:
|
||||
return SqlAlchemyConditionFieldRepository(session)
|
||||
|
||||
|
||||
def _composite_repo_factory(session: DbSession) -> CompositeRepository:
|
||||
return SqlAlchemyCompositeRepository(session)
|
||||
|
||||
|
||||
def _signal_repo_factory(session: DbSession) -> SignalRepository:
|
||||
return SqlAlchemySignalRepository(session)
|
||||
|
||||
|
||||
def _strategy_repo_factory(session: DbSession) -> StrategyRepository:
|
||||
return SqlAlchemyStrategyRepository(session)
|
||||
|
||||
|
||||
def _global_config_repo_factory(session: DbSession) -> GlobalConfigRepository:
|
||||
return SqlAlchemyGlobalConfigRepository(session)
|
||||
|
||||
|
||||
def _combo_repo_factory(session: DbSession) -> ComboRepository:
|
||||
return SqlAlchemyComboRepository(session)
|
||||
|
||||
|
||||
StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
|
||||
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
|
||||
EngineDep = Annotated[QuantEngine, Depends(_engine_factory)]
|
||||
ResearchServiceDep = Annotated[ResearchService, Depends(_service_factory)]
|
||||
SelectionServiceDep = Annotated[SelectionService, Depends(_selection_service_factory)]
|
||||
SelectionRepoDep = Annotated[SelectionRepository, Depends(_selection_repo_factory)]
|
||||
FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)]
|
||||
ConditionFieldRepoDep = Annotated[ConditionFieldRepository, Depends(_condition_field_repo_factory)]
|
||||
CompositeRepoDep = Annotated[CompositeRepository, Depends(_composite_repo_factory)]
|
||||
SignalRepoDep = Annotated[SignalRepository, Depends(_signal_repo_factory)]
|
||||
SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)]
|
||||
ReplayServiceDep = Annotated[ReplayService, Depends(_replay_service_factory)]
|
||||
ChartServiceDep = Annotated[ChartService, Depends(_chart_service_factory)]
|
||||
StrategyRepoDep = Annotated[StrategyRepository, Depends(_strategy_repo_factory)]
|
||||
GlobalConfigRepoDep = Annotated[GlobalConfigRepository, Depends(_global_config_repo_factory)]
|
||||
ComboRepoDep = Annotated[ComboRepository, Depends(_combo_repo_factory)]
|
||||
|
||||
|
||||
def _job_repo_factory(session: DbSession):
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
|
||||
return SqlAlchemyJobRepository(session)
|
||||
|
||||
|
||||
def _experiment_repo_factory(session: DbSession):
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyExperimentRepository,
|
||||
)
|
||||
|
||||
return SqlAlchemyExperimentRepository(session)
|
||||
|
||||
|
||||
JobRepoDep = Annotated[JobRepository, Depends(_job_repo_factory)]
|
||||
ExperimentRepoDep = Annotated[ExperimentRepository, Depends(_experiment_repo_factory)]
|
||||
@@ -0,0 +1,177 @@
|
||||
"""Experiment API(Phase 4):列表 / 详情 / 删除(单个 + 批量)/ 一键复跑。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, HTTPException, Query, Response
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.api.deps import DbSession, ExperimentRepoDep, JobRepoDep
|
||||
from app.application.services.job_executor import new_id, run_job_background
|
||||
from app.domain.entities.research import (
|
||||
ExperimentRecord,
|
||||
ExperimentSummary,
|
||||
JobRecord,
|
||||
JobStatus,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/experiments", tags=["experiments"])
|
||||
|
||||
# 列表默认返回条数与上限:原实现硬编码 limit=50 且不暴露总数,>50 条时更老的
|
||||
# 归档静默不可见(AGENT.md §7);现在默认 200、上限 1000,并用 X-Total-Count
|
||||
# 暴露过滤后的真实总数,客户端可据此翻页(offset)。
|
||||
LIST_DEFAULT_LIMIT = 200
|
||||
LIST_MAX_LIMIT = 1000
|
||||
|
||||
# 单次批量删除的 id 上限:够覆盖列表一页(默认 200 条),又不至于让一次请求
|
||||
# 把整张表拖进内存。超了直接 422(附上限值),不静默截断成前 N 个。
|
||||
BULK_DELETE_MAX_IDS = 200
|
||||
|
||||
|
||||
class BulkDeleteRequest(BaseModel):
|
||||
"""批量删除请求体(**必填** id 列表,1~200 个)。"""
|
||||
|
||||
ids: list[str] = Field(
|
||||
min_length=1,
|
||||
max_length=BULK_DELETE_MAX_IDS,
|
||||
description=f"要删除的归档 id,1~{BULK_DELETE_MAX_IDS} 个(重复 id 自动去重)",
|
||||
)
|
||||
|
||||
|
||||
def _experiment_meta(exp: ExperimentRecord | ExperimentSummary) -> dict:
|
||||
"""列表项视图(body 形状与旧版一致,仅**新增** data_version / job_id / result_bytes)。
|
||||
|
||||
`exp` 可以是完整实体(详情路径)或 `ExperimentSummary`(列表路径,不含
|
||||
result_json):两条路径都只读元数据字段,无需把大字段拉回来算体积。
|
||||
"""
|
||||
spec = json.loads(exp.spec_json)
|
||||
result_bytes = getattr(exp, "result_bytes", None)
|
||||
if result_bytes is None: # 完整实体:result_json 已在内存,len() 零成本
|
||||
result_bytes = len(exp.result_json or "")
|
||||
return {
|
||||
"id": exp.id,
|
||||
"kind": exp.kind,
|
||||
"factors": [f["name"] for f in spec.get("factors", [])],
|
||||
"period": spec.get("period"),
|
||||
"rebalance": spec.get("rebalance"),
|
||||
"top_n": spec.get("selection", {}).get("top_n"),
|
||||
"summary_text": exp.summary_text,
|
||||
"code_version": exp.code_version,
|
||||
"data_version": exp.data_version,
|
||||
"job_id": exp.job_id,
|
||||
"result_bytes": int(result_bytes),
|
||||
"created_at": exp.created_at,
|
||||
}
|
||||
|
||||
|
||||
def _experiment_full(exp: ExperimentRecord) -> dict:
|
||||
from app.api.jobs import _decode_result
|
||||
|
||||
return {
|
||||
**_experiment_meta(exp),
|
||||
"spec": json.loads(exp.spec_json),
|
||||
"result": _decode_result(exp.kind, exp.result_json),
|
||||
}
|
||||
|
||||
|
||||
@router.get("", summary="Experiment 列表(过滤 + 分页,X-Total-Count 给总数)")
|
||||
def list_experiments(
|
||||
experiment_repo: ExperimentRepoDep,
|
||||
response: Response,
|
||||
kind: Annotated[str | None, Query(description="按类型精确过滤(backtest/factor_test/selection)")] = None,
|
||||
q: Annotated[
|
||||
str | None, Query(description="大小写不敏感模糊匹配 id / 因子名 / summary_text")
|
||||
] = None,
|
||||
limit: Annotated[int, Query(ge=1, le=LIST_MAX_LIMIT)] = LIST_DEFAULT_LIMIT,
|
||||
offset: Annotated[int, Query(ge=0)] = 0,
|
||||
) -> list[dict]:
|
||||
"""归档列表。
|
||||
|
||||
- 过滤与分页在 SQL 层完成(仓储 `list_filtered`),不把全表拉回内存;
|
||||
- 响应头 `X-Total-Count` = **过滤后**的归档总数(不受 limit/offset 影响),
|
||||
客户端据此判断是否被截断并翻页(AGENT.md §7:不静默截断)。
|
||||
"""
|
||||
rows = experiment_repo.list_filtered(kind=kind, q=q, limit=limit, offset=offset)
|
||||
total = experiment_repo.count_filtered(kind=kind, q=q)
|
||||
response.headers["X-Total-Count"] = str(total)
|
||||
return [_experiment_meta(e) for e in rows]
|
||||
|
||||
|
||||
@router.post("/bulk-delete", summary=f"批量删除归档(1~{BULK_DELETE_MAX_IDS} 个,逐个回报)")
|
||||
def bulk_delete_experiments(
|
||||
body: BulkDeleteRequest,
|
||||
session: DbSession,
|
||||
experiment_repo: ExperimentRepoDep,
|
||||
) -> dict:
|
||||
"""批量删除归档,返回 `{"deleted": [...], "missing": [...], "count": n}`。
|
||||
|
||||
语义与单个删除完全一致(只删 experiment 行,job 历史保留),另外:
|
||||
- **重复 id 先按出现顺序去重**(同一个 id 报两次没有意义,也不该算两次成功);
|
||||
- **不静默跳过**:库里没有的 id 单独放进 `missing`,让界面能如实说
|
||||
「删了 3 个,2 个没找到(可能已被别处删掉)」——把缺失当成功会更难排查;
|
||||
- 一次请求内的 id 上限 `BULK_DELETE_MAX_IDS`,超了 Pydantic 直接 422 并带上限值。
|
||||
"""
|
||||
ids = list(dict.fromkeys(body.ids))
|
||||
deleted: list[str] = []
|
||||
missing: list[str] = []
|
||||
for experiment_id in ids:
|
||||
(deleted if experiment_repo.delete(experiment_id) else missing).append(experiment_id)
|
||||
session.commit()
|
||||
return {"deleted": deleted, "missing": missing, "count": len(deleted)}
|
||||
|
||||
|
||||
@router.get("/{experiment_id}", summary="Experiment 详情(含完整结果)")
|
||||
def get_experiment(experiment_id: str, experiment_repo: ExperimentRepoDep) -> dict:
|
||||
exp = experiment_repo.get(experiment_id)
|
||||
if exp is None:
|
||||
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
|
||||
return _experiment_full(exp)
|
||||
|
||||
|
||||
@router.delete("/{experiment_id}", summary="删除 Experiment 归档(不影响关联 Job 记录)")
|
||||
def delete_experiment(
|
||||
experiment_id: str,
|
||||
session: DbSession,
|
||||
experiment_repo: ExperimentRepoDep,
|
||||
) -> dict:
|
||||
"""删除归档本身,返回 `{"deleted": "<id>"}`。
|
||||
|
||||
语义(写清楚避免误解):
|
||||
- **只删 experiment 行**。关联的 job 记录是「执行历史」,一律保留,
|
||||
`GET /api/jobs/{id}` 仍可查到该 Job 的状态、阶段与错误信息。
|
||||
- 注意:完整结果现在只存归档一份(见 experiment_archive / job_executor),
|
||||
因此删除归档后 `GET /api/jobs/{id}` 的 `result` 会是 null,并附带
|
||||
`result_unavailable_reason` 说明归档已被删除(如实暴露,不静默给空结果)。
|
||||
- 删除不可恢复;如需长期保留结果,请勿删除对应归档。
|
||||
"""
|
||||
if not experiment_repo.delete(experiment_id):
|
||||
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
|
||||
session.commit()
|
||||
return {"deleted": experiment_id}
|
||||
|
||||
|
||||
@router.post("/{experiment_id}/rerun", summary="一键复跑(AGENT §21:历史实验可重放)")
|
||||
def rerun_experiment(
|
||||
experiment_id: str,
|
||||
background: BackgroundTasks,
|
||||
session: DbSession,
|
||||
experiment_repo: ExperimentRepoDep,
|
||||
job_repo: JobRepoDep,
|
||||
) -> dict:
|
||||
exp = experiment_repo.get(experiment_id)
|
||||
if exp is None:
|
||||
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
|
||||
job = JobRecord(
|
||||
id=new_id("JOB"),
|
||||
kind=exp.kind,
|
||||
spec_json=exp.spec_json,
|
||||
status=JobStatus.QUEUED,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
job_repo.create(job)
|
||||
session.commit()
|
||||
background.add_task(run_job_background, job.id)
|
||||
return {"job_id": job.id, "status": job.status, "origin_experiment": exp.id}
|
||||
@@ -0,0 +1,157 @@
|
||||
"""因子目录 API:/api/factors(M7.1 起读 DB,2026-10 支持参数化实例)。
|
||||
|
||||
## 目录的三条规则(详见 application/services/factor_catalog.py)
|
||||
|
||||
① 注册表有、库里没有 → 补齐(历史 bug:表非空后新因子永远进不了目录)。
|
||||
② 能算出来的行,口径字段按代码改回 —— 目录不允许与引擎口径不一致(手改会被纠正)。
|
||||
③ 库里多出来的行保留(只补不删),标 `resolvable` 告知是否算得出来。
|
||||
|
||||
## 参数化(本文件新增的部分)
|
||||
|
||||
- **暴露**:每个因子返回 `template` / `params` / `param_specs`(可编辑参数与允许范围)/
|
||||
`label`(中文名含参数)/ `source` / `resolvable` / `enabled`,界面据此渲染参数表与表单。
|
||||
- **新建**:`POST /api/factors` 传 `{template, params}` → 生成参数化实例
|
||||
`momentum(window=90,direction=higher_is_better)`。参数写在名字里,所以它**冻结**了自己的
|
||||
口径:以后无论谁再改参数,既有策略/归档按各自名字里的参数计算,不会变义。
|
||||
- **停用**:`PATCH /api/factors` 传 `{name, enabled}`。名字里有括号/等号/逗号,放路径里
|
||||
会被各种代理折腾,所以放在 body 里。
|
||||
- 越界参数、未知模板、重复的参数组合 → **422**,并在 detail 里说明允许范围。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.api.deps import DbSession, FactorRepoDep
|
||||
from app.application.services.factor_catalog import (
|
||||
create_parameterized_factor,
|
||||
set_factor_enabled,
|
||||
sync_registry_factors,
|
||||
)
|
||||
from app.domain.entities.factor import FactorDefinition, FactorParam
|
||||
from app.quant.factors import FactorError, list_templates, resolve_factor
|
||||
|
||||
router = APIRouter(prefix="/factors", tags=["factors"])
|
||||
|
||||
|
||||
class FactorTemplateOut(BaseModel):
|
||||
"""模板(算法家族):可编辑参数 + 默认值 + 口径说明,供「新建参数化因子」表单。"""
|
||||
|
||||
name: str
|
||||
label: str
|
||||
description: str
|
||||
formula: str
|
||||
brief: str
|
||||
requires: list[str] = Field(default_factory=list)
|
||||
frequency: str = "daily"
|
||||
direction_default: str = "higher_is_better"
|
||||
param_specs: list[FactorParam] = Field(default_factory=list)
|
||||
defaults: dict[str, Any] = Field(default_factory=dict)
|
||||
instances: list[str] = Field(default_factory=list) # 该模板已有的内置实例名
|
||||
|
||||
|
||||
class FactorCreate(BaseModel):
|
||||
"""新建参数化因子:模板 + 参数(缺省项取模板默认值)。"""
|
||||
|
||||
template: str
|
||||
params: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class FactorPatch(BaseModel):
|
||||
"""开关因子(name 放 body:名字里有括号/等号/逗号,不适合放路径)。"""
|
||||
|
||||
name: str
|
||||
enabled: bool
|
||||
|
||||
|
||||
def _template_out(tpl) -> FactorTemplateOut:
|
||||
return FactorTemplateOut(
|
||||
name=tpl.name,
|
||||
label=tpl.label,
|
||||
description=tpl.description,
|
||||
formula=tpl.formula,
|
||||
brief=tpl.brief,
|
||||
requires=list(tpl.requires),
|
||||
frequency=tpl.frequency,
|
||||
direction_default=tpl.direction_default,
|
||||
param_specs=[FactorParam.from_spec(s) for s in tpl.specs()],
|
||||
defaults=tpl.defaults(),
|
||||
instances=[name for name, _params in tpl.instances],
|
||||
)
|
||||
|
||||
|
||||
def _enrich(row: FactorDefinition) -> FactorDefinition:
|
||||
"""把库里的行补成「引擎口径的投影」:能算出来的行一律以引擎为准。
|
||||
|
||||
- 算得出来 → description/formula/brief/frequency/lookback/direction/requires/
|
||||
template/params/param_specs/label/source 全部取自引擎(目录永不撒谎);
|
||||
`enabled` 仍是库里的人配值。
|
||||
- 算不出来(历史手工登记行)→ 原样返回并标 `resolvable=False`,界面显示为不可用。
|
||||
"""
|
||||
try:
|
||||
defn, _fn = resolve_factor(row.name)
|
||||
except FactorError:
|
||||
return row.model_copy(update={"resolvable": False, "label": row.name})
|
||||
return FactorDefinition.from_factor_def(defn, enabled=row.enabled).model_copy(
|
||||
update={"created_at": row.created_at, "version": row.version}
|
||||
)
|
||||
|
||||
|
||||
@router.get("", summary="因子目录(注册表投影 + 参数化实例,含可编辑参数)")
|
||||
def list_factor_catalog(factor_repo: FactorRepoDep, session: DbSession) -> list[FactorDefinition]:
|
||||
"""读目录(先做幂等同步,稳态零写入),每行补上参数化视图。
|
||||
|
||||
前端据此渲染:因子名 / 中文名含参数 / 窗口 · 方向等真实参数 / 是否可当过滤条件。
|
||||
"""
|
||||
sync_registry_factors(factor_repo, session)
|
||||
return [_enrich(row) for row in factor_repo.list()]
|
||||
|
||||
|
||||
@router.get("/templates", summary="因子模板:可编辑参数与允许范围")
|
||||
def list_factor_templates() -> list[FactorTemplateOut]:
|
||||
"""全部模板(动量 / 波动率 / 量比 / 乖离 / 反转 / 接近新高 / 股息率…)。
|
||||
|
||||
每个模板给出 `param_specs`(参数名、类型、允许范围/枚举、默认值、说明)——
|
||||
界面据此渲染受控表单:**参数只在给定范围内选/填**,越界在 API 层就被拒。
|
||||
"""
|
||||
return [_template_out(tpl) for tpl in list_templates()]
|
||||
|
||||
|
||||
@router.post("", status_code=201, summary="新建参数化因子(模板 + 参数)")
|
||||
def create_factor(
|
||||
payload: FactorCreate,
|
||||
factor_repo: FactorRepoDep,
|
||||
session: DbSession,
|
||||
) -> FactorDefinition:
|
||||
"""从模板派生一个新的参数化因子实例;参数写进名字,因此口径被冻结。
|
||||
|
||||
重复的参数组合不会重复创建(409 语义由 422 承载并给出已有名字,前端直接提示即可)。
|
||||
"""
|
||||
try:
|
||||
row = create_parameterized_factor(
|
||||
factor_repo, session, template=payload.template, params=payload.params
|
||||
)
|
||||
except FactorError as exc:
|
||||
raise HTTPException(status_code=422, detail=str(exc)) from exc
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=422, detail=str(exc)) from exc
|
||||
return _enrich(row)
|
||||
|
||||
|
||||
@router.patch("", summary="启用/停用因子(只影响能否被选中)")
|
||||
def patch_factor(
|
||||
payload: FactorPatch,
|
||||
factor_repo: FactorRepoDep,
|
||||
session: DbSession,
|
||||
) -> FactorDefinition:
|
||||
"""停用只把因子从下拉里拿掉:既有策略/归档仍按名字解析(历史不变义)。"""
|
||||
try:
|
||||
row = set_factor_enabled(factor_repo, session, name=payload.name, enabled=payload.enabled)
|
||||
except LookupError as exc:
|
||||
raise HTTPException(status_code=404, detail=str(exc)) from exc
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=422, detail=str(exc)) from exc
|
||||
return _enrich(row)
|
||||
@@ -0,0 +1,156 @@
|
||||
"""异步研究 Job API(Phase 4):提交 / 查询 / SSE 进度。
|
||||
|
||||
POST /api/jobs 创建 Job(BackgroundTasks 后台执行),立即返回 job_id
|
||||
GET /api/jobs/{id} 状态 + 结果(成功时内嵌 result)
|
||||
GET /api/jobs/{id}/events SSE 进度(queued→running→success|failed)
|
||||
|
||||
执行模式见 config job.mode:subprocess 时研究任务在独立子进程跑(内存隔离),
|
||||
API worker 不被重任务拖垮(内存优化专项)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, HTTPException, Query
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from app.api.deps import DbSession, ExperimentRepoDep, JobRepoDep
|
||||
from app.application.services.job_executor import new_id, run_job_background, terminate_active
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
FactorTestReport,
|
||||
JobRecord,
|
||||
JobStatus,
|
||||
ResearchSpec,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.session import SessionLocal
|
||||
|
||||
router = APIRouter(prefix="/jobs", tags=["jobs"])
|
||||
|
||||
|
||||
def _decode_result(kind: str, result_json: str | None):
|
||||
from app.domain.entities.selection import SelectionResult
|
||||
|
||||
if result_json is None:
|
||||
return None
|
||||
if kind == "selection":
|
||||
return SelectionResult.model_validate_json(result_json)
|
||||
# combo 回测的归档 kind 记为 "backtest",但 Job.kind 仍是 "combo" ——
|
||||
# 其结果同样是 BacktestResult,按 backtest 解码(否则会被当成因子测试而校验失败)。
|
||||
model = BacktestResult if kind in ("backtest", "combo") else FactorTestReport
|
||||
return model.model_validate_json(result_json)
|
||||
|
||||
|
||||
def _job_view(job: JobRecord, experiment_repo=None) -> dict:
|
||||
"""Job 视图:`result` 契约不变(成功时内嵌**完整**结果)。
|
||||
|
||||
结果来源(2026-09 起完整结果只在 experiment 存一份,job.result_json 不再重复写):
|
||||
1. `job.experiment_id` 有值且能读到归档 → 解码 experiment.result_json;
|
||||
2. 否则回退解码 `job.result_json`(老记录 / 归档被删除前的历史数据);
|
||||
3. 归档被删除且 job 侧无副本 → `result=None`,并给出
|
||||
`result_unavailable_reason` 如实说明原因(AGENT §7:不静默给空结果)。
|
||||
"""
|
||||
view = job.model_dump()
|
||||
view.pop("result_json", None)
|
||||
view["spec"] = json.loads(job.spec_json)
|
||||
|
||||
result = None
|
||||
source = None
|
||||
if job.experiment_id and experiment_repo is not None:
|
||||
exp = experiment_repo.get(job.experiment_id)
|
||||
if exp is not None:
|
||||
result = _decode_result(job.kind, exp.result_json)
|
||||
source = "experiment"
|
||||
else:
|
||||
view["result_unavailable_reason"] = (
|
||||
f"归档 {job.experiment_id} 已不存在(可能已被删除);"
|
||||
"完整结果仅存于归档,Job 记录本身不再保存结果副本"
|
||||
)
|
||||
if result is None and job.result_json:
|
||||
result = _decode_result(job.kind, job.result_json)
|
||||
source = "job"
|
||||
view["result"] = result
|
||||
view["result_source"] = source
|
||||
return view
|
||||
|
||||
|
||||
@router.post("", summary="创建异步研究 Job")
|
||||
def create_job(
|
||||
spec: ResearchSpec,
|
||||
background: BackgroundTasks,
|
||||
session: DbSession,
|
||||
job_repo: JobRepoDep,
|
||||
) -> dict:
|
||||
job = JobRecord(
|
||||
id=new_id("JOB"),
|
||||
kind=spec.type,
|
||||
spec_json=spec.model_dump_json(),
|
||||
status=JobStatus.QUEUED,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
job_repo.create(job)
|
||||
session.commit()
|
||||
background.add_task(run_job_background, job.id)
|
||||
return {"job_id": job.id, "status": job.status}
|
||||
|
||||
|
||||
@router.get("", summary="Job 列表")
|
||||
def list_jobs(
|
||||
job_repo: JobRepoDep,
|
||||
experiment_repo: ExperimentRepoDep,
|
||||
kind: Annotated[str | None, Query(description="按类型过滤(backtest/factor_test)")] = None,
|
||||
limit: Annotated[int, Query(ge=1, le=200)] = 20,
|
||||
) -> list[dict]:
|
||||
return [
|
||||
_job_view(j, experiment_repo) for j in job_repo.list_recent(kind=kind, limit=limit)
|
||||
]
|
||||
|
||||
|
||||
@router.post("/{job_id}/cancel", summary="取消 Job(queued/running)")
|
||||
def cancel_job(job_id: str, session: DbSession, job_repo: JobRepoDep) -> dict:
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail=f"Job {job_id} 不存在")
|
||||
if job.status not in (JobStatus.QUEUED, JobStatus.RUNNING):
|
||||
return {"job_id": job_id, "status": job.status, "cancelled": False}
|
||||
job.status = JobStatus.CANCELLED
|
||||
job.stage = None
|
||||
job_repo.update(job)
|
||||
session.commit()
|
||||
terminate_active(job_id) # 终止研究子进程(若有);父进程兜底已跳过 CANCELLED
|
||||
return {"job_id": job_id, "status": JobStatus.CANCELLED, "cancelled": True}
|
||||
|
||||
|
||||
@router.get("/{job_id}", summary="查询 Job 状态与结果")
|
||||
def get_job(job_id: str, job_repo: JobRepoDep, experiment_repo: ExperimentRepoDep) -> dict:
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail=f"Job {job_id} 不存在")
|
||||
return _job_view(job, experiment_repo)
|
||||
|
||||
|
||||
@router.get("/{job_id}/events", summary="Job 进度 SSE")
|
||||
async def job_events(job_id: str) -> StreamingResponse:
|
||||
async def gen():
|
||||
while True:
|
||||
with SessionLocal() as session:
|
||||
job = SqlAlchemyJobRepository(session).get(job_id)
|
||||
if job is None:
|
||||
yield "event: error\ndata: job not found\n\n"
|
||||
return
|
||||
payload = json.dumps(
|
||||
{"job_id": job.id, "status": job.status, "stage": job.stage}, ensure_ascii=False
|
||||
)
|
||||
yield f"data: {payload}\n\n"
|
||||
if job.status in (JobStatus.SUCCESS, JobStatus.FAILED, JobStatus.CANCELLED):
|
||||
return
|
||||
await asyncio.sleep(0.4)
|
||||
|
||||
return StreamingResponse(gen(), media_type="text/event-stream")
|
||||
@@ -0,0 +1,33 @@
|
||||
"""Bar Replay API(M9-6):POST /api/replays —— 线性重放选股/信号时间线。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.api.deps import ReplayServiceDep
|
||||
from app.domain.entities.replay import ReplayResult
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
from app.domain.entities.signal import SignalRules
|
||||
|
||||
router = APIRouter(prefix="/replays", tags=["replays"])
|
||||
|
||||
|
||||
class ReplayRequest(BaseModel):
|
||||
query: SelectionQuery
|
||||
rules: SignalRules = SignalRules()
|
||||
start: date
|
||||
end: date
|
||||
top_n: int = Field(default=5, ge=1, le=20)
|
||||
|
||||
|
||||
@router.post("", response_model=ReplayResult, summary="线性重放(as_of 逐日,仅用当时数据)")
|
||||
def run_replay(req: ReplayRequest, service: ReplayServiceDep) -> ReplayResult:
|
||||
if req.start >= req.end:
|
||||
raise HTTPException(status_code=400, detail="start 必须早于 end")
|
||||
try:
|
||||
return service.replay(req.query, req.rules, req.start, req.end, req.top_n)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
@@ -0,0 +1,165 @@
|
||||
"""研究执行 API:/api/factor-tests 与 /api/backtests。
|
||||
|
||||
Phase 3 为同步执行(样本有限);Phase 4 将改为 Job + SSE 异步(接口契约不变)。
|
||||
|
||||
**归档(2026-09 补齐)**:同步端点此前只把结果塞进进程内存 `_LAST_*`,重启即丢,
|
||||
完全不落库 —— 与「研究可复现」(AGENT.md §21)矛盾。现改为:返回结果前调用
|
||||
`archive_experiment` 落库(复用与异步 Job 完全相同的归档实现),并通过响应头
|
||||
`X-Experiment-Id` 暴露归档 id。body 形状保持不变(前端与既有测试依赖它)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import date
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Response
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.api.deps import DbSession, ExperimentRepoDep, ResearchServiceDep
|
||||
from app.application.services.experiment_archive import archive_experiment
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
FactorCorrelationReport,
|
||||
FactorSpec,
|
||||
FactorTestReport,
|
||||
ResearchSpec,
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
from app.quant.factors import FactorError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(tags=["research"])
|
||||
|
||||
# 内存中的最近结果(仅作「最近一次」便捷读取;持久化归档见 archive_experiment)
|
||||
_LAST_BACKTEST: dict[str, BacktestResult] = {}
|
||||
_LAST_FACTOR_TEST: dict[str, FactorTestReport] = {}
|
||||
|
||||
|
||||
def _archive_or_expose(
|
||||
response: Response,
|
||||
*,
|
||||
session,
|
||||
kind: str,
|
||||
spec_json: str,
|
||||
result,
|
||||
experiment_repo,
|
||||
) -> None:
|
||||
"""归档同步端点的计算结果,并把归档结果如实反映到响应头。
|
||||
|
||||
取舍(AGENT.md §7「不静默」+ §24):研究结果是用户真实等待数十秒得到的产出,
|
||||
归档是副产物。若数据库故障导致归档失败,直接抛 500 会把**已经算出来的可用结果**
|
||||
一并丢掉;故此处捕获异常、接口仍然 200 返回完整结果,同时:
|
||||
- `logger.warning` 落盘(服务端可观测);
|
||||
- 响应头 `X-Archive-Error: <Type>: <msg>` 如实暴露失败原因(客户端可判读)。
|
||||
即「结果不丢 + 失败不静默」两者兼顾;成功时给 `X-Experiment-Id`。
|
||||
"""
|
||||
try:
|
||||
experiment = archive_experiment(
|
||||
session=session,
|
||||
kind=kind,
|
||||
spec_json=spec_json,
|
||||
result=result,
|
||||
experiment_repo=experiment_repo,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 —— 归档失败不得吞掉计算结果
|
||||
logger.warning("同步端点归档失败(kind=%s):%s: %s", kind, type(exc).__name__, exc)
|
||||
# HTTP 头只能承载 latin-1:中文错误信息降级为 ASCII 转义,避免编码异常掩盖原因
|
||||
reason = f"{type(exc).__name__}: {exc}"[:180]
|
||||
response.headers["X-Archive-Error"] = reason.encode("ascii", "backslashreplace").decode(
|
||||
"ascii"
|
||||
)
|
||||
return
|
||||
response.headers["X-Experiment-Id"] = experiment.id
|
||||
|
||||
|
||||
@router.post("/factor-tests", response_model=FactorTestReport, summary="运行单因子测试(同步)")
|
||||
def run_factor_test(
|
||||
spec: ResearchSpec,
|
||||
service: ResearchServiceDep,
|
||||
session: DbSession,
|
||||
experiment_repo: ExperimentRepoDep,
|
||||
response: Response,
|
||||
) -> FactorTestReport:
|
||||
try:
|
||||
report = service.run_factor_test(spec)
|
||||
except (ValueError, FactorError) as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
_LAST_FACTOR_TEST["default"] = report
|
||||
_archive_or_expose(
|
||||
response,
|
||||
session=session,
|
||||
kind="factor_test",
|
||||
spec_json=spec.model_dump_json(),
|
||||
result=report,
|
||||
experiment_repo=experiment_repo,
|
||||
)
|
||||
return report
|
||||
|
||||
|
||||
@router.post("/backtests", response_model=BacktestResult, summary="运行回测(同步)")
|
||||
def run_backtest(
|
||||
spec: ResearchSpec,
|
||||
service: ResearchServiceDep,
|
||||
session: DbSession,
|
||||
experiment_repo: ExperimentRepoDep,
|
||||
response: Response,
|
||||
) -> BacktestResult:
|
||||
try:
|
||||
result = service.run_backtest(spec)
|
||||
except (ValueError, FactorError) as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
_LAST_BACKTEST["default"] = result
|
||||
_archive_or_expose(
|
||||
response,
|
||||
session=session,
|
||||
kind="backtest",
|
||||
spec_json=spec.model_dump_json(),
|
||||
result=result,
|
||||
experiment_repo=experiment_repo,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/backtests/last", response_model=BacktestResult, summary="最近一次回测结果")
|
||||
def last_backtest() -> BacktestResult:
|
||||
if "default" not in _LAST_BACKTEST:
|
||||
raise HTTPException(status_code=404, detail="尚无回测结果,请先 POST /api/backtests")
|
||||
return _LAST_BACKTEST["default"]
|
||||
|
||||
|
||||
class FactorCorrelationRequest(BaseModel):
|
||||
"""因子相关性分析请求(v3 §12)。"""
|
||||
|
||||
universe: UniverseSpec = UniverseSpec()
|
||||
factors: list[FactorSpec] = Field(min_length=2)
|
||||
period: tuple[date, date]
|
||||
price_adjustment: str = Field(default="none", pattern="^(none|qfq)$")
|
||||
|
||||
|
||||
@router.post(
|
||||
"/factor-correlations",
|
||||
response_model=FactorCorrelationReport,
|
||||
summary="多因子两两相关(横截面 Spearman)",
|
||||
)
|
||||
def run_factor_correlation(
|
||||
req: FactorCorrelationRequest,
|
||||
service: ResearchServiceDep,
|
||||
) -> FactorCorrelationReport:
|
||||
if req.period[0] >= req.period[1]:
|
||||
raise HTTPException(status_code=400, detail="period 必须满足 start < end")
|
||||
spec = ResearchSpec(
|
||||
type="factor_test",
|
||||
universe=req.universe,
|
||||
price_adjustment=req.price_adjustment,
|
||||
factors=req.factors,
|
||||
selection=SelectionSpec(top_n=10),
|
||||
rebalance="monthly",
|
||||
period=req.period,
|
||||
)
|
||||
try:
|
||||
return service.run_factor_correlation(spec)
|
||||
except (ValueError, FactorError) as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
@@ -1,14 +1,46 @@
|
||||
"""API 路由聚合。
|
||||
|
||||
后续业务路由按 AGENT.md §17 面向业务对象挂载:
|
||||
/api/stocks /api/universes /api/factors /api/strategies /api/backtests /api/experiments /api/jobs
|
||||
业务路由面向业务对象(AGENT.md §17):/api/stocks /api/factors
|
||||
/api/factor-tests /api/backtests /api/experiments(Phase4) /api/jobs(Phase4) /api/agent(Phase5)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api import health
|
||||
from app.api import (
|
||||
agent,
|
||||
charts,
|
||||
combos,
|
||||
composites,
|
||||
condition_fields,
|
||||
config,
|
||||
experiments,
|
||||
factors,
|
||||
health,
|
||||
jobs,
|
||||
replays,
|
||||
research,
|
||||
selections,
|
||||
signals,
|
||||
stocks,
|
||||
strategies,
|
||||
)
|
||||
|
||||
api_router = APIRouter()
|
||||
api_router.include_router(health.router)
|
||||
api_router.include_router(stocks.router)
|
||||
api_router.include_router(factors.router)
|
||||
api_router.include_router(condition_fields.router)
|
||||
api_router.include_router(composites.router)
|
||||
api_router.include_router(research.router)
|
||||
api_router.include_router(selections.router)
|
||||
api_router.include_router(charts.router)
|
||||
api_router.include_router(replays.router)
|
||||
api_router.include_router(signals.router)
|
||||
api_router.include_router(strategies.router)
|
||||
api_router.include_router(config.router)
|
||||
api_router.include_router(combos.router)
|
||||
api_router.include_router(jobs.router)
|
||||
api_router.include_router(experiments.router)
|
||||
api_router.include_router(agent.router)
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
"""选股 API(M6.3):提交/查询选股,结果落库可复现。
|
||||
|
||||
POST /api/selections 同步执行一次选股并落库 → {selection_id, result}
|
||||
GET /api/selections/{id} 读回某次选股完整结果
|
||||
GET /api/selections 历史选股元数据(可过滤 as_of/method)
|
||||
|
||||
同步执行:单日全市场因子评分/条件计算量轻(秒级);未来若超时再迁 Job。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, HTTPException, Query
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.api.deps import (
|
||||
DbSession,
|
||||
JobRepoDep,
|
||||
SelectionRepoDep,
|
||||
SelectionServiceDep,
|
||||
)
|
||||
from app.application.services.job_executor import new_id, run_job_background
|
||||
from app.domain.entities.research import JobRecord, JobStatus
|
||||
from app.domain.entities.selection import SelectionMeta, SelectionQuery, SelectionResult
|
||||
|
||||
router = APIRouter(prefix="/selections", tags=["selections"])
|
||||
|
||||
|
||||
class SelectionRun(BaseModel):
|
||||
selection_id: str
|
||||
result: SelectionResult
|
||||
|
||||
|
||||
@router.post("/jobs", summary="提交全市场/长任务选股为异步 Job")
|
||||
def submit_selection_job(
|
||||
query: SelectionQuery,
|
||||
background: BackgroundTasks,
|
||||
session: DbSession,
|
||||
job_repo: JobRepoDep,
|
||||
) -> dict:
|
||||
job = JobRecord(
|
||||
id=new_id("JOB"),
|
||||
kind="selection",
|
||||
spec_json=query.model_dump_json(),
|
||||
status=JobStatus.QUEUED,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
job_repo.create(job)
|
||||
session.commit()
|
||||
background.add_task(run_job_background, job.id)
|
||||
return {"job_id": job.id, "status": job.status}
|
||||
|
||||
|
||||
@router.post("", response_model=SelectionRun, summary="执行一次选股(同步)并落库")
|
||||
def run_selection(
|
||||
query: SelectionQuery,
|
||||
service: SelectionServiceDep,
|
||||
selection_repo: SelectionRepoDep,
|
||||
session: DbSession,
|
||||
) -> SelectionRun:
|
||||
result = service.select(query)
|
||||
selection_id = new_id("SEL")
|
||||
selection_repo.save(selection_id, result)
|
||||
session.commit()
|
||||
return SelectionRun(selection_id=selection_id, result=result)
|
||||
|
||||
|
||||
@router.get("/{selection_id}", response_model=SelectionResult, summary="读回一次选股结果")
|
||||
def get_selection(
|
||||
selection_id: str,
|
||||
selection_repo: SelectionRepoDep,
|
||||
) -> SelectionResult:
|
||||
result = selection_repo.get(selection_id)
|
||||
if result is None:
|
||||
raise HTTPException(status_code=404, detail=f"选股记录 {selection_id} 不存在")
|
||||
return result
|
||||
|
||||
|
||||
_AsOfQuery = Annotated[date | None, Query(description="按选股时点过滤")]
|
||||
_MethodQuery = Annotated[str | None, Query(pattern="^(score|condition)$")]
|
||||
_LimitQuery = Annotated[int, Query(ge=1, le=200)]
|
||||
|
||||
|
||||
@router.get("", response_model=list[SelectionMeta], summary="历史选股元数据列表")
|
||||
def list_selections(
|
||||
selection_repo: SelectionRepoDep,
|
||||
as_of: _AsOfQuery = None,
|
||||
method: _MethodQuery = None,
|
||||
limit: _LimitQuery = 20,
|
||||
) -> list[SelectionMeta]:
|
||||
return selection_repo.list_recent(as_of=as_of, method=method, limit=limit)
|
||||
@@ -0,0 +1,65 @@
|
||||
"""交易信号 API(M8.1):提交/查询信号(落库可复现)。
|
||||
|
||||
POST /api/signals body: {query: SelectionQuery, rules?: SignalRules}
|
||||
GET /api/signals/{id} 读回某次信号
|
||||
GET /api/signals 历史信号元数据(可过滤 as_of)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Query
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.api.deps import DbSession, SignalRepoDep, SignalServiceDep
|
||||
from app.application.services.job_executor import new_id
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
from app.domain.entities.signal import SignalMeta, SignalResult, SignalRules
|
||||
|
||||
router = APIRouter(prefix="/signals", tags=["signals"])
|
||||
|
||||
_AsOfQuery = Annotated[date | None, Query(description="按信号时点过滤")]
|
||||
_LimitQuery = Annotated[int, Query(ge=1, le=200)]
|
||||
|
||||
|
||||
class SignalRequest(BaseModel):
|
||||
query: SelectionQuery
|
||||
rules: SignalRules = SignalRules()
|
||||
|
||||
|
||||
class SignalRun(BaseModel):
|
||||
signal_id: str
|
||||
result: SignalResult
|
||||
|
||||
|
||||
@router.post("", response_model=SignalRun, summary="生成一次交易信号(同步)并落库")
|
||||
def run_signal(
|
||||
req: SignalRequest,
|
||||
service: SignalServiceDep,
|
||||
signal_repo: SignalRepoDep,
|
||||
session: DbSession,
|
||||
) -> SignalRun:
|
||||
result = service.signal(req.query, req.rules)
|
||||
signal_id = new_id("SIG")
|
||||
signal_repo.save(signal_id, result)
|
||||
session.commit()
|
||||
return SignalRun(signal_id=signal_id, result=result)
|
||||
|
||||
|
||||
@router.get("/{signal_id}", response_model=SignalResult, summary="读回一次信号结果")
|
||||
def get_signal(signal_id: str, signal_repo: SignalRepoDep) -> SignalResult:
|
||||
result = signal_repo.get(signal_id)
|
||||
if result is None:
|
||||
raise HTTPException(status_code=404, detail=f"信号记录 {signal_id} 不存在")
|
||||
return result
|
||||
|
||||
|
||||
@router.get("", response_model=list[SignalMeta], summary="历史信号元数据列表")
|
||||
def list_signals(
|
||||
signal_repo: SignalRepoDep,
|
||||
as_of: _AsOfQuery = None,
|
||||
limit: _LimitQuery = 20,
|
||||
) -> list[SignalMeta]:
|
||||
return signal_repo.list_recent(as_of=as_of, limit=limit)
|
||||
@@ -0,0 +1,45 @@
|
||||
"""股票查询 API:/api/stocks。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
|
||||
from app.api.deps import StockRepoDep
|
||||
from app.domain.entities.market import Stock
|
||||
|
||||
router = APIRouter(prefix="/stocks", tags=["stocks"])
|
||||
|
||||
|
||||
@router.get("", response_model=list[Stock], summary="股票列表")
|
||||
def list_stocks(
|
||||
repo: StockRepoDep,
|
||||
q: str | None = None,
|
||||
limit: int = 50,
|
||||
offset: int = 0,
|
||||
) -> list[Stock]:
|
||||
if limit > 500:
|
||||
limit = 500
|
||||
stocks = repo.list()
|
||||
if q:
|
||||
needle = q.upper()
|
||||
stocks = [s for s in stocks if needle in s.symbol or needle in s.name.upper()]
|
||||
return stocks[offset : offset + limit]
|
||||
|
||||
|
||||
@router.get("/names", response_model=dict[str, str], summary="全市场股票名称映射")
|
||||
def list_stock_names(repo: StockRepoDep) -> dict[str, str]:
|
||||
"""`{symbol: name}` 全市场名称缓存(前端启动时一次性拉取,避免表格 N+1 查询)。
|
||||
|
||||
契约固定为 dict(不是数组):前端已有调用方按 dict 形状消费。
|
||||
路由**必须**声明在 `/{symbol}` 之前,否则会被路径参数吞掉("names" 被当作代码)。
|
||||
名称缺失(空字符串)的标的直接略过 —— 前端按「名称未知」渲染,不塞空串假装有名称。
|
||||
"""
|
||||
return {s.symbol: s.name for s in repo.list() if s.name}
|
||||
|
||||
|
||||
@router.get("/{symbol}", response_model=Stock, summary="按代码查询")
|
||||
def get_stock(symbol: str, repo: StockRepoDep) -> Stock:
|
||||
stock = repo.get_by_symbol(symbol)
|
||||
if stock is None:
|
||||
raise HTTPException(status_code=404, detail=f"未找到股票 {symbol}")
|
||||
return stock
|
||||
@@ -0,0 +1,156 @@
|
||||
"""选股策略 API:/api/strategies CRUD + 说明/公式生成。
|
||||
|
||||
2026-09 重构:策略库只存「选股条件组合」(股票池+因子+条件),不再持有回测执行参数;
|
||||
回测改由「回测组合」(/api/combos)驱动,故旧的 /{id}/expand(→ResearchSpec)已移除。
|
||||
|
||||
POST /api/strategies 保存策略(name 唯一;description 为空时自动补全)
|
||||
POST /api/strategies/describe body: ResearchSpec → StrategyDoc(见下方说明,非策略库路径)
|
||||
GET /api/strategies 列表
|
||||
GET /api/strategies/{id}
|
||||
PUT /api/strategies/{id} 原地更新(不新建、不刷新 created_at)
|
||||
DELETE /api/strategies/{id}
|
||||
GET /api/strategies/{id}/describe → StrategyDoc
|
||||
|
||||
关于两个 describe 端点(不是历史遗留,各有明确用途,别合并):
|
||||
- `GET /{id}/describe` → 入参是已保存的 **SelectionStrategy**(策略库「看说明/公式」用);
|
||||
- `POST /describe` → 入参是 **ResearchSpec**,**给归档页**用:`/experiments/{id}` 要按当时
|
||||
归档的旧 ResearchSpec 快照(单策略回测路径,含 selection/rebalance/costs)复述口径。
|
||||
该路径仍然存在(`POST /api/backtests` 是底层 escape hatch),所以这里必须继续支持。
|
||||
回测页本身已不再调用它(组合回测走 ComboRunSpec + 归档页组合卡片)。
|
||||
|
||||
路由顺序注意:`/describe` 这类**字面量路径**一律声明在 `/{strategy_id}` 之前 ——
|
||||
否则会被路径参数吞掉(AGENT.md §17 的既有教训,/api/stocks/names 同源问题)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
|
||||
from app.api.deps import DbSession, StrategyRepoDep
|
||||
from app.application.services.job_executor import new_id
|
||||
from app.domain.entities.research import ResearchSpec
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.quant.strategy_doc import StrategyDoc, describe_strategy
|
||||
|
||||
router = APIRouter(prefix="/strategies", tags=["strategies"])
|
||||
|
||||
|
||||
|
||||
# strategy.description 列宽(StrategyModel.description = String(300))。
|
||||
# 自动补全的说明必须落在列宽内,否则 MySQL 严格模式会直接报 Data too long(SQLite 不拦,
|
||||
# 所以只在测试库上跑是发现不了的)。超长时按字符截断并加省略号 —— 显式标记有截断,
|
||||
# 不做「悄悄改短」;完整说明始终可由 POST /describe 重新生成。
|
||||
_DESCRIPTION_MAX_CHARS = 300
|
||||
|
||||
|
||||
def _ensure_description(definition: SelectionStrategy) -> SelectionStrategy:
|
||||
"""说明为空/纯空白时,用 `describe_strategy(...).summary` 补全(需求:策略必须有说明)。
|
||||
|
||||
说明由 spec **真实推导**(AGENT.md §24:不许编造),只在空值时补、不覆盖显式说明。
|
||||
放在 API 层是因为这是「保存契约」的准入补全;Agent 的 create_strategy 工具走仓储
|
||||
直写(description 非必填),因此不受影响(AGENT.md §28 工具链路保持可用)。
|
||||
"""
|
||||
if definition.description.strip():
|
||||
return definition
|
||||
summary = describe_strategy(definition).summary
|
||||
if len(summary) > _DESCRIPTION_MAX_CHARS:
|
||||
summary = summary[: _DESCRIPTION_MAX_CHARS - 1] + "…"
|
||||
return definition.model_copy(update={"description": summary})
|
||||
|
||||
|
||||
@router.post("", response_model=SelectionStrategy, summary="保存选股策略")
|
||||
def create_strategy(
|
||||
definition: SelectionStrategy,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
) -> SelectionStrategy:
|
||||
try:
|
||||
saved = strategy_repo.save(
|
||||
_ensure_description(definition).model_copy(update={"id": new_id("STG")})
|
||||
)
|
||||
session.commit()
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
# 回读持久化后的实体:仓储 save() 返回的是入参(created_at 为空),
|
||||
# 直接返回会让 POST 响应缺创建时间、与 GET/列表不一致(前端展示依赖该字段)。
|
||||
return strategy_repo.get(saved.id) or saved
|
||||
|
||||
|
||||
@router.post("/describe", response_model=StrategyDoc, summary="按 ResearchSpec 生成策略说明与公式")
|
||||
def describe_research_spec(spec: ResearchSpec) -> StrategyDoc:
|
||||
"""按 **ResearchSpec** 生成说明/公式 —— 服务于归档页,不是策略库路径。
|
||||
|
||||
调用方是 `/experiments/{id}`:它按归档里冻结的 ResearchSpec 快照(单策略回测)复述
|
||||
「选股条件 + 交易执行依据」。纯函数实现(app.quant.strategy_doc),无 IO/DB,因此
|
||||
不依赖任何保存状态,历史归档随时可复述。策略库自身的说明走 `GET /{id}/describe`。
|
||||
"""
|
||||
return describe_strategy(spec)
|
||||
|
||||
|
||||
@router.get("", response_model=list[SelectionStrategy], summary="选股策略列表")
|
||||
def list_strategies(strategy_repo: StrategyRepoDep) -> list[SelectionStrategy]:
|
||||
return strategy_repo.list()
|
||||
|
||||
|
||||
@router.get("/{strategy_id}", response_model=SelectionStrategy, summary="读取选股策略")
|
||||
def get_strategy(strategy_id: str, strategy_repo: StrategyRepoDep) -> SelectionStrategy:
|
||||
row = strategy_repo.get(strategy_id)
|
||||
if row is None:
|
||||
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
|
||||
return row
|
||||
|
||||
|
||||
@router.get("/{strategy_id}/describe", response_model=StrategyDoc, summary="生成策略说明与公式")
|
||||
def describe_saved_strategy(
|
||||
strategy_id: str, strategy_repo: StrategyRepoDep
|
||||
) -> StrategyDoc:
|
||||
"""已保存策略的说明/公式(404 语义与 GET /{strategy_id} 一致)。
|
||||
|
||||
策略定义不含回测区间,说明里的区间为占位文本(`warnings` 中已如实标注),
|
||||
"""
|
||||
row = strategy_repo.get(strategy_id)
|
||||
if row is None:
|
||||
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
|
||||
return describe_strategy(row)
|
||||
|
||||
|
||||
@router.put("/{strategy_id}", response_model=SelectionStrategy, summary="原地更新选股策略")
|
||||
def update_strategy(
|
||||
strategy_id: str,
|
||||
definition: SelectionStrategy,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
) -> SelectionStrategy:
|
||||
"""原地更新(策略库「编辑」用):id 以**路径**为准,created_at 沿用库中已有值。
|
||||
|
||||
为什么必须显式带上 created_at:仓储 `save()` 只在 created_at 为空时才写 now()
|
||||
(既有行不会覆盖该列),但返回值是**传入的实体**;若这里不带,响应里的创建时间
|
||||
就会变成 None,而策略库按创建时间展示 —— 一改就丢时间同样会误导前端。
|
||||
改名撞车由仓储 `save()` 抛 ValueError(策略名已存在:X),这里转 400。
|
||||
"""
|
||||
existing = strategy_repo.get(strategy_id)
|
||||
if existing is None:
|
||||
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
|
||||
payload = definition.model_copy(
|
||||
update={"id": strategy_id, "created_at": existing.created_at}
|
||||
)
|
||||
try:
|
||||
saved = strategy_repo.save(_ensure_description(payload))
|
||||
session.commit()
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
# 与 create 一致:回读持久化实体,保证响应 == GET 读回(含 description/created_at)
|
||||
return strategy_repo.get(saved.id) or saved
|
||||
|
||||
|
||||
@router.delete("/{strategy_id}", summary="删除策略")
|
||||
def delete_strategy(
|
||||
strategy_id: str,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
) -> dict:
|
||||
if not strategy_repo.delete(strategy_id):
|
||||
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
|
||||
session.commit()
|
||||
return {"deleted": strategy_id}
|
||||
|
||||
@@ -0,0 +1,221 @@
|
||||
"""Chart Service(v3 §20.1)—— 只聚合与坐标整理,不重算研究结果。
|
||||
|
||||
- K 线/量/指标:基于主口径(adjust=none)行情,按请求 adjust 在**显示层**折算 qfq/hfq
|
||||
(绝不回写研究数据;研究执行仍用 price_adjustment 指定口径)
|
||||
- 标记:selections/signals 来自各自落库历史(by-symbol);backtest fills 来自 Experiment
|
||||
内 BacktestResult 的 trades/positions
|
||||
- 口径纪律(v3 §20.5):显示价 basis 与回测执行价 basis 分别记录;不一致时对 marker 做
|
||||
与 K 线相同的坐标换算,保证成交点贴图
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
from app.domain.entities.chart import (
|
||||
OHLC,
|
||||
ChartMetadata,
|
||||
ChartResult,
|
||||
EventMarker,
|
||||
SeriesPoint,
|
||||
VolumePoint,
|
||||
)
|
||||
from app.domain.entities.market import AdjustFactor, DailyBar, Stock
|
||||
from app.domain.entities.research import BacktestResult
|
||||
from app.domain.repositories.market import (
|
||||
AdjustFactorRepository,
|
||||
DailyBarRepository,
|
||||
StockRepository,
|
||||
)
|
||||
|
||||
_MA_WINDOWS = (20, 60)
|
||||
|
||||
|
||||
def _main_bars(bars: list[DailyBar]) -> list[DailyBar]:
|
||||
"""仅取主口径行(adjust=none),丢弃新浪 qfq 兜底行 —— 显示层一律自 none 折算。"""
|
||||
return [b for b in bars if b.adjust == "none"]
|
||||
|
||||
|
||||
def _factor_multipliers(
|
||||
adj_repo: AdjustFactorRepository,
|
||||
symbol: str,
|
||||
start: date,
|
||||
end: date,
|
||||
mode: str,
|
||||
) -> dict[date, float]:
|
||||
"""返回 {trade_date: 显示折算系数};mode=none → 空。
|
||||
|
||||
基准(v3 §20.5):qfq 以**该股最新因子**(截至今天,而非图表区间末)归一,
|
||||
保证历史区间随最新除权平移正确;hfq 直接用累积因子。
|
||||
"""
|
||||
if mode == "none":
|
||||
return {}
|
||||
factors: list[AdjustFactor] = adj_repo.get_range(symbol, date(1990, 1, 1), date.today())
|
||||
if not factors:
|
||||
return {}
|
||||
by_day = {f.trade_date: float(f.factor) for f in factors}
|
||||
latest = max(by_day.values())
|
||||
out: dict[date, float] = {}
|
||||
for day, f in by_day.items():
|
||||
out[day] = f / latest if mode == "qfq" else f
|
||||
return out
|
||||
|
||||
|
||||
class ChartService:
|
||||
def __init__(
|
||||
self,
|
||||
stock_repo: StockRepository,
|
||||
daily_repo: DailyBarRepository,
|
||||
adj_repo: AdjustFactorRepository,
|
||||
) -> None:
|
||||
self._stock_repo = stock_repo
|
||||
self._daily_repo = daily_repo
|
||||
self._adj_repo = adj_repo
|
||||
|
||||
def stock(self, symbol: str) -> Stock | None:
|
||||
return self._stock_repo.get_by_symbol(symbol)
|
||||
|
||||
def stock_chart(
|
||||
self,
|
||||
symbol: str,
|
||||
start: date,
|
||||
end: date,
|
||||
adjust: str = "none",
|
||||
execution_price_basis: str | None = None,
|
||||
extra_markers: list[EventMarker] | None = None,
|
||||
) -> ChartResult:
|
||||
"""基础个股 K 线图(可叠加 fills 等外部标记)。"""
|
||||
stock = self.stock(symbol)
|
||||
name = stock.name if stock else ""
|
||||
raw = _main_bars(self._daily_repo.get_range(symbol, start, end))
|
||||
mult = _factor_multipliers(self._adj_repo, symbol, start, end, adjust)
|
||||
bars: list[OHLC] = []
|
||||
volume: list[VolumePoint] = []
|
||||
for b in raw:
|
||||
m = mult.get(b.trade_date, 1.0)
|
||||
bars.append(
|
||||
OHLC(
|
||||
time=b.trade_date,
|
||||
open=_v(b.open, m),
|
||||
high=_v(b.high, m),
|
||||
low=_v(b.low, m),
|
||||
close=_v(b.close, m),
|
||||
)
|
||||
)
|
||||
volume.append(VolumePoint(time=b.trade_date, value=_v(b.volume, 1.0)))
|
||||
|
||||
markers = _convert_markers(extra_markers or [], mult)
|
||||
indicators = _ma_indicators(bars)
|
||||
return ChartResult(
|
||||
metadata=ChartMetadata(
|
||||
symbol=symbol,
|
||||
name=name,
|
||||
adjust_mode=adjust,
|
||||
execution_price_basis=execution_price_basis,
|
||||
start=start,
|
||||
end=end,
|
||||
bar_count=len(bars),
|
||||
indicator_windows=list(_MA_WINDOWS),
|
||||
),
|
||||
bars=bars,
|
||||
volume=volume,
|
||||
indicators=indicators,
|
||||
fills=[m for m in markers if m.kind.startswith("fill_")],
|
||||
signals=[m for m in markers if m.kind.startswith("signal_")],
|
||||
selections=[m for m in markers if m.kind == "selection"],
|
||||
holding_periods=_holding_periods(bars),
|
||||
)
|
||||
|
||||
# ---- 由已存历史构造标记(不重算) ----
|
||||
|
||||
def backtest_stock_chart(
|
||||
self,
|
||||
result: BacktestResult,
|
||||
symbol: str,
|
||||
start: date,
|
||||
end: date,
|
||||
adjust: str = "none",
|
||||
) -> ChartResult:
|
||||
"""回测个股视图:K 线 + 选股意图/未成交信号/实际成交三类标记(v3 §20.3)。"""
|
||||
basis = (result.config_snapshot or {}).get("price_adjustment", "none")
|
||||
markers = _result_to_markers(result, symbol)
|
||||
return self.stock_chart(symbol, start, end, adjust, execution_price_basis=basis,
|
||||
extra_markers=markers)
|
||||
|
||||
|
||||
def _result_to_markers(result: BacktestResult, symbol: str) -> list[EventMarker]:
|
||||
"""由回测 history 生成个股标记:fills(成交)/ signals(未成交意图)/ selections(选股)。"""
|
||||
markers: list[EventMarker] = []
|
||||
# 实际成交(fills)与未成交信号(signal_history 中 filled=False)
|
||||
for a in result.signal_history:
|
||||
if a.symbol != symbol:
|
||||
continue
|
||||
if a.filled:
|
||||
kind = "fill_buy" if a.signal == "BUY" else "fill_sell"
|
||||
text = [f"{'买入' if a.signal=='BUY' else '卖出'} @ {a.price:.2f}(basis={_basis_of(result)})"]
|
||||
markers.append(
|
||||
EventMarker(time=a.date, kind=kind, symbol=symbol, price=a.price, text=text)
|
||||
)
|
||||
else:
|
||||
kind = "signal_buy" if a.signal == "BUY" else "signal_sell"
|
||||
text = [a.reject_reason or f"{a.signal} 未成交"]
|
||||
markers.append(EventMarker(time=a.date, kind=kind, symbol=symbol, price=a.price, text=text))
|
||||
# 选股意图(selection_history 中该 symbol 的命中)
|
||||
for pk in result.selection_history:
|
||||
if pk.symbol != symbol:
|
||||
continue
|
||||
markers.append(
|
||||
EventMarker(
|
||||
time=pk.date,
|
||||
kind="selection",
|
||||
symbol=symbol,
|
||||
score=pk.score,
|
||||
text=[f"选股意图 rank #{pk.rank}"],
|
||||
)
|
||||
)
|
||||
return markers
|
||||
|
||||
|
||||
def _basis_of(result: BacktestResult) -> str:
|
||||
return (result.config_snapshot or {}).get("price_adjustment", "none")
|
||||
|
||||
|
||||
def _convert_markers(markers: list[EventMarker], mult: dict[date, float]) -> list[EventMarker]:
|
||||
"""显示口径与执行价 basis 不一致时,把 marker 价格折算到 K 线坐标系(v3 §20.5)。"""
|
||||
out: list[EventMarker] = []
|
||||
for m in markers:
|
||||
if m.price is not None and mult:
|
||||
k = mult.get(m.time)
|
||||
if k is not None:
|
||||
m = m.model_copy(update={"price": round(m.price * k, 4)})
|
||||
out.append(m)
|
||||
return out
|
||||
|
||||
|
||||
def _holding_periods(bars: list[OHLC]) -> list[dict]:
|
||||
"""v1 空实现占位(持仓区间渲染 v3 §20.4 后续细化)。"""
|
||||
return []
|
||||
|
||||
|
||||
def _ma_indicators(bars: list[OHLC]) -> dict[str, list[SeriesPoint]]:
|
||||
import statistics
|
||||
|
||||
closes = [b.close for b in bars]
|
||||
out: dict[str, list[SeriesPoint]] = {}
|
||||
for w in _MA_WINDOWS:
|
||||
series: list[SeriesPoint] = []
|
||||
for i, b in enumerate(bars):
|
||||
if i + 1 < w:
|
||||
continue
|
||||
window = closes[i + 1 - w : i + 1]
|
||||
if all(v is not None for v in window):
|
||||
series.append(SeriesPoint(time=b.time, value=round(statistics.fmean(window), 4)))
|
||||
out[f"ma{w}"] = series
|
||||
return out
|
||||
|
||||
|
||||
def _v(v, m: float) -> float | None:
|
||||
if v is None:
|
||||
return None
|
||||
f = float(v)
|
||||
return round(f * m, 4)
|
||||
@@ -0,0 +1,263 @@
|
||||
"""回测组合服务:把「组合 + 选股策略 + 公共配置」解析并执行成 BacktestResult。
|
||||
|
||||
职责(应用层用例,AGENT.md §16/§17):
|
||||
- 装配行情数据(复用 ResearchService 的 load_daily_df / universe 过滤 / 名称回填);
|
||||
- 为每个选股策略构造「as_of → 合格股票集」闭包(复用 selection 求值器,保证与
|
||||
`/api/selections` 同口径,v2 §25);
|
||||
- 调 combo_engine.run_combo_backtest(多策略 Borda + 持仓区间 + 日/周/月);
|
||||
- 把可复现的 ComboRunSpec 写进结果 config_snapshot(已在引擎内完成)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, timedelta
|
||||
from typing import Any
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from app.domain.entities.combo import (
|
||||
BacktestCombo,
|
||||
ComboRunSpec,
|
||||
GlobalConfig,
|
||||
SelectionStrategyRef,
|
||||
)
|
||||
from app.domain.entities.research import BacktestResult, UniverseSpec
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.quant.combo_engine import run_combo_backtest
|
||||
from app.quant.selection import build_condition_fields, eligible_symbols
|
||||
from app.quant.service import _fill_names, load_daily_df, split_factor_columns
|
||||
from app.quant.universe import filter_stocks, names_as_of, resolve_members
|
||||
|
||||
|
||||
class ComboService:
|
||||
"""回测组合用例入口。依赖注入各 Repository + 引擎无关的数据装配函数。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stock_repo,
|
||||
daily_repo,
|
||||
*,
|
||||
index_repo=None,
|
||||
basic_repo=None,
|
||||
financial_repo=None,
|
||||
name_repo=None,
|
||||
) -> None:
|
||||
self._stock_repo = stock_repo
|
||||
self._daily_repo = daily_repo
|
||||
self._index_repo = index_repo
|
||||
self._basic_repo = basic_repo
|
||||
self._financial_repo = financial_repo
|
||||
self._name_repo = name_repo
|
||||
self._last_stocks: list = []
|
||||
|
||||
def run(
|
||||
self,
|
||||
combo: BacktestCombo,
|
||||
strategies: list[SelectionStrategy],
|
||||
config: GlobalConfig,
|
||||
on_stage=None,
|
||||
) -> BacktestResult:
|
||||
if not strategies:
|
||||
raise ValueError("回测组合至少需要引用一个选股策略")
|
||||
# 校验引用的策略 id 与传入一致(防御性:调用方应已按 combo.strategy_ids 取齐)
|
||||
given = {s.id for s in strategies}
|
||||
missing = [sid for sid in combo.strategy_ids if sid not in given]
|
||||
if missing:
|
||||
raise ValueError(f"组合引用的选股策略未提供:{missing}")
|
||||
|
||||
_stage(on_stage, "data_loading")
|
||||
daily = self._load_daily(combo, strategies, config)
|
||||
_stage(on_stage, "backtesting")
|
||||
eligibility_fns = [self._build_eligibility(s, daily) for s in strategies]
|
||||
refs = [_to_ref(s) for s in strategies]
|
||||
result = run_combo_backtest(
|
||||
combo=combo,
|
||||
strategies=refs,
|
||||
costs=config.to_cost_spec(),
|
||||
price_adjustment=config.price_adjustment,
|
||||
daily=daily,
|
||||
eligibility_fns=eligibility_fns,
|
||||
)
|
||||
# 把价格口径写进 config_snapshot(与 ResearchService._annotate_price_basis 同口径),
|
||||
# 否则归档页/结果头读不到 adjust_mode,会误显示「不复权」(组合实际用的是公共配置的复权)。
|
||||
# 注意:不能覆盖整个 config_snapshot —— 引擎已把 ComboRunSpec 固化在里面(可复现依据)。
|
||||
mode = config.price_adjustment
|
||||
result.config_snapshot["price_basis"] = {
|
||||
"adjust_mode": mode,
|
||||
"price_basis": "adjust_factor" if mode != "none" else "raw_close",
|
||||
"execution_price_basis": "close_adj" if mode != "none" else "close_raw",
|
||||
}
|
||||
_stage(on_stage, "analysis")
|
||||
return _fill_names(result, self._last_stocks)
|
||||
|
||||
# ---- 数据装配(与 ResearchService 同口径,复用底层函数) ----
|
||||
|
||||
def _merged_universe(self, strategies: list[SelectionStrategy]) -> UniverseSpec:
|
||||
"""合并各策略的股票池口径用于「装哪些股票的行情」。
|
||||
|
||||
取并集语义:symbols 白名单取并集;exclude_st / min_listing_days 取**最宽松**
|
||||
(任一策略不剔 ST 则不剔,min_listing_days 取最小)—— 因为最终选股由各策略
|
||||
自己的 eligibility 闭包再过滤,这里只为「行情装配覆盖足够多的股票」。
|
||||
index_code 不一致时无法合并 → 报错(同一组合里混用不同指数成分没有明确语义)。
|
||||
"""
|
||||
indices = {s.universe.index_code for s in strategies if s.universe.index_code}
|
||||
if len(indices) > 1:
|
||||
raise ValueError(
|
||||
f"组合内各选股策略的指数成分不一致({sorted(indices)}),无法合并股票池;"
|
||||
"请统一指数或改用 symbols 白名单"
|
||||
)
|
||||
symbols: set[str] = set()
|
||||
for s in strategies:
|
||||
symbols.update(s.universe.symbols)
|
||||
return UniverseSpec(
|
||||
market=strategies[0].universe.market,
|
||||
exclude_st=all(s.universe.exclude_st for s in strategies),
|
||||
exclude_suspended=all(s.universe.exclude_suspended for s in strategies),
|
||||
min_listing_days=min(s.universe.min_listing_days for s in strategies),
|
||||
index_code=indices.pop() if indices else None,
|
||||
symbols=sorted(symbols),
|
||||
)
|
||||
|
||||
def _load_daily(
|
||||
self, combo: BacktestCombo, strategies: list[SelectionStrategy], config: GlobalConfig
|
||||
) -> pd.DataFrame:
|
||||
start, end = combo.period
|
||||
data_start = start - timedelta(days=300) # 因子 warmup 余量
|
||||
all_stocks = self._stock_repo.list()
|
||||
merged = self._merged_universe(strategies)
|
||||
name_at, _applied = names_as_of(all_stocks, start, self._name_repo)
|
||||
stocks = filter_stocks(
|
||||
all_stocks, merged, as_of=start,
|
||||
members=resolve_members(self._index_repo, merged, start),
|
||||
name_at=name_at,
|
||||
)
|
||||
self._last_stocks = stocks
|
||||
# 所需列 = 所有策略因子 + 所有策略条件引用列 + close
|
||||
needed = {"close"}
|
||||
for s in strategies:
|
||||
from app.domain.entities.research import ResearchSpec
|
||||
from app.quant.engine import factor_required_columns
|
||||
|
||||
# 借用既有列裁剪逻辑:构造一个临时 spec 只为算 required_columns
|
||||
tmp = ResearchSpec(
|
||||
type="backtest", universe=s.universe, factors=s.factors,
|
||||
conditions=s.conditions, period=combo.period,
|
||||
)
|
||||
needed |= factor_required_columns(tmp)
|
||||
bar_cols, basic_cols = split_factor_columns(needed)
|
||||
symbols = [st.symbol for st in stocks]
|
||||
daily = load_daily_df(
|
||||
self._daily_repo, symbols, data_start, end, sorted(bar_cols),
|
||||
adjust="none", price_adjust=config.price_adjustment,
|
||||
)
|
||||
if basic_cols:
|
||||
daily = self._attach_basic(daily, symbols, data_start, end, sorted(basic_cols))
|
||||
return daily
|
||||
|
||||
def _attach_basic(self, daily, symbols, start, end, columns) -> pd.DataFrame:
|
||||
from app.quant.service import load_basic_df, merge_basic_into_daily
|
||||
|
||||
if self._basic_repo is None:
|
||||
raise ValueError(
|
||||
f"选股策略条件/因子需要每日指标列 {columns}(daily_basic),但未注入 DailyBasicRepository"
|
||||
)
|
||||
basic = load_basic_df(self._basic_repo, symbols, start, end, columns)
|
||||
if basic.empty:
|
||||
raise ValueError(
|
||||
f"daily_basic 在 {start}~{end} 无数据,无法计算需要 {columns} 的因子/条件"
|
||||
)
|
||||
return merge_basic_into_daily(daily, basic)
|
||||
|
||||
def _build_eligibility(self, strategy: SelectionStrategy, daily: pd.DataFrame):
|
||||
"""单策略的「as_of → 合格股票集」闭包(与 ResearchService._build_eligibility 同口径)。"""
|
||||
if not self._last_stocks:
|
||||
if not strategy.conditions and not strategy.universe.exclude_st:
|
||||
return None
|
||||
raise ValueError("universe 过滤结果为空,无法构造选股条件求值器")
|
||||
statics = {s.symbol: s.model_dump() for s in self._last_stocks}
|
||||
candidates = sorted(statics)
|
||||
st_fn = self._build_st_filter(strategy, candidates)
|
||||
if not strategy.conditions:
|
||||
if st_fn is None:
|
||||
return None
|
||||
allowed: dict[date, set[str]] = {}
|
||||
|
||||
def _st_only(as_of: date) -> set[str]:
|
||||
if as_of not in allowed:
|
||||
allowed[as_of] = set(candidates) - st_fn(as_of)
|
||||
return allowed[as_of]
|
||||
|
||||
return _st_only
|
||||
|
||||
uses_fundamental = any(
|
||||
f.startswith("fundamental.")
|
||||
for c in strategy.conditions
|
||||
for f in (c.field, c.ref or "")
|
||||
)
|
||||
cache: dict[date, set[str]] = {}
|
||||
|
||||
def _fn(as_of: date) -> set[str]:
|
||||
if as_of in cache:
|
||||
return cache[as_of]
|
||||
financial = self._load_financial(candidates, as_of) if uses_fundamental else {}
|
||||
fields = build_condition_fields(daily, strategy.conditions, pd.Timestamp(as_of))
|
||||
if not fields:
|
||||
cache[as_of] = set()
|
||||
return cache[as_of]
|
||||
passed = set(eligible_symbols(candidates, strategy.conditions, statics, fields, financial))
|
||||
if st_fn is not None:
|
||||
passed -= st_fn(as_of)
|
||||
cache[as_of] = passed
|
||||
return cache[as_of]
|
||||
|
||||
return _fn
|
||||
|
||||
def _build_st_filter(self, strategy: SelectionStrategy, candidates: list[str]):
|
||||
if not strategy.universe.exclude_st or self._name_repo is None:
|
||||
return None
|
||||
cache: dict[date, set[str]] = {}
|
||||
|
||||
def _fn(as_of: date) -> set[str]:
|
||||
if as_of not in cache:
|
||||
name_at, applied = names_as_of(self._last_stocks, as_of, self._name_repo)
|
||||
if not applied[0]:
|
||||
cache[as_of] = set()
|
||||
else:
|
||||
st_syms: set[str] = set()
|
||||
for st in self._last_stocks:
|
||||
nm = (name_at or {}).get(st.symbol) or st.name
|
||||
if nm and "ST" in nm.upper():
|
||||
st_syms.add(st.symbol)
|
||||
cache[as_of] = st_syms
|
||||
return cache[as_of]
|
||||
|
||||
return _fn
|
||||
|
||||
def _load_financial(self, symbols: list[str], as_of: date) -> dict[str, Any]:
|
||||
if self._financial_repo is None:
|
||||
raise ValueError("条件引用了 fundamental.* 字段,但未注入 FinancialRepository")
|
||||
getter = getattr(self._financial_repo, "list_announced_many", None)
|
||||
rows = list(getter(symbols, as_of)) if getter else []
|
||||
out: dict[str, Any] = {}
|
||||
for r in rows:
|
||||
out[r.symbol] = r
|
||||
return out
|
||||
|
||||
|
||||
def _to_ref(s: SelectionStrategy) -> SelectionStrategyRef:
|
||||
return SelectionStrategyRef(
|
||||
id=s.id,
|
||||
name=s.name,
|
||||
universe=s.universe.model_dump(),
|
||||
factors=[f.model_dump() for f in s.factors],
|
||||
conditions=[c.model_dump() for c in s.conditions],
|
||||
)
|
||||
|
||||
|
||||
def _stage(cb, name: str) -> None:
|
||||
if cb is not None:
|
||||
cb(name)
|
||||
|
||||
|
||||
# 让 ComboRunSpec 在模块导入时完成前向引用重建(entities/combo.py 末尾已 rebuild,此处兜底)
|
||||
ComboRunSpec.model_rebuild()
|
||||
@@ -0,0 +1,198 @@
|
||||
"""字段库用例:目录 seed + 校验 + 增删改(2026-10)。
|
||||
|
||||
与 factor_catalog 同一套规矩:**DB 是目录契约源,代码注册表是可用性的唯一事实来源**。
|
||||
读取时把「注册表有、库里没有」的内置字段补进去(只补不删、不覆盖用户改过的文案)。
|
||||
|
||||
单位的两层含义(2026-10 补)见 `quant/condition_fields.py`:``unit`` 存的是**界面单位**
|
||||
(输入/显示用,可从注册表给的阶梯里选),引擎始终按**基准单位**存储与比较,
|
||||
换算是提交/回显时按系数做的 —— 所以改单位不会让任何历史策略变义。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from app.domain.entities.condition_field import ConditionField
|
||||
from app.domain.repositories.condition_field import ConditionFieldRepository
|
||||
from app.quant.condition_fields import (
|
||||
GROUP_ORDER,
|
||||
FieldDef,
|
||||
available_fields,
|
||||
curated_fields,
|
||||
get_field,
|
||||
reason_unsupported,
|
||||
unit_allowed,
|
||||
unit_options,
|
||||
)
|
||||
|
||||
_SORT_BASE = 100
|
||||
|
||||
|
||||
def _check_unit(name: str, unit: str) -> str:
|
||||
"""界面单位必须落在注册表给的阶梯里(不允许自由文本 —— 见 AGENT §24)。
|
||||
|
||||
空串 = 保持基准单位。给出可用单位清单,用户不用猜为什么被拒。
|
||||
"""
|
||||
d = get_field(name)
|
||||
base = d.unit if d else ""
|
||||
if not unit.strip():
|
||||
return base
|
||||
if not unit_allowed(name, unit.strip()):
|
||||
allowed = "、".join(u for u, _ in unit_options(name)) or "(无可用单位)"
|
||||
raise ValueError(
|
||||
f"字段「{name}」不支持单位「{unit.strip()}」:可选 {allowed}。"
|
||||
"单位只能从这些里选,因为换算是按固定系数做的(引擎按基准单位比较)。"
|
||||
)
|
||||
return unit.strip()
|
||||
|
||||
|
||||
def _sort_order(d: FieldDef) -> int:
|
||||
"""同分组内保持注册表顺序(分组顺序 × 1000 + 组内序号)。"""
|
||||
try:
|
||||
g = GROUP_ORDER.index(d.group_name)
|
||||
except ValueError:
|
||||
g = len(GROUP_ORDER)
|
||||
return g * 1000 + _SORT_BASE
|
||||
|
||||
|
||||
def _from_def(d: FieldDef) -> ConditionField:
|
||||
return ConditionField(
|
||||
name=d.name,
|
||||
label=d.label,
|
||||
description=d.description,
|
||||
kind=d.kind,
|
||||
group_name=d.group_name,
|
||||
unit=d.unit,
|
||||
source="builtin",
|
||||
enabled=True,
|
||||
sort_order=_sort_order(d),
|
||||
)
|
||||
|
||||
|
||||
def sync_builtin_fields(repo: ConditionFieldRepository, session) -> int:
|
||||
"""补齐缺失的**默认内置字段**(幂等;已存在的一律不动,用户改过的文案得以保留)。
|
||||
|
||||
只 seed `curated=True` 的那批:`curated=False` 的字段是「引擎支持但默认不进库」的
|
||||
选项,留给用户在字段库里按需新增(见 `list_available`)—— 否则「新增字段」永远
|
||||
无字段可选。
|
||||
"""
|
||||
existing = {f.name for f in repo.list()}
|
||||
missing = [_from_def(d) for d in curated_fields() if d.name not in existing]
|
||||
added = repo.insert_missing(missing)
|
||||
if added:
|
||||
session.commit()
|
||||
return added
|
||||
|
||||
|
||||
def list_fields(repo: ConditionFieldRepository, session, include_disabled: bool = True):
|
||||
"""字段库列表(首次读取自动 seed)。按分组/注册表顺序排序。"""
|
||||
items = repo.list()
|
||||
if not items:
|
||||
sync_builtin_fields(repo, session)
|
||||
items = repo.list()
|
||||
# 代码里新注册的默认字段也要补上:按差集触发,稳态零写入
|
||||
if {d.name for d in curated_fields()} - {f.name for f in items}:
|
||||
sync_builtin_fields(repo, session)
|
||||
items = repo.list()
|
||||
if not include_disabled:
|
||||
items = [i for i in items if i.enabled]
|
||||
return items
|
||||
|
||||
|
||||
def list_available(repo: ConditionFieldRepository, session):
|
||||
"""引擎支持但尚未进目录的字段(「新增字段」的可选项)。"""
|
||||
return available_fields({f.name for f in repo.list()})
|
||||
|
||||
|
||||
def create_field(
|
||||
repo: ConditionFieldRepository,
|
||||
session,
|
||||
*,
|
||||
name: str,
|
||||
label: str = "",
|
||||
description: str = "",
|
||||
group_name: str = "",
|
||||
unit: str = "",
|
||||
enabled: bool = True,
|
||||
) -> ConditionField:
|
||||
"""新增自定义字段。
|
||||
|
||||
校验顺序有意如此:先看引擎能不能算(不能算就 422,绝不放行)→ 再看是否已存在。
|
||||
这样用户拿到的是「这个字段引擎算不出来」而不是含糊的「已存在」。
|
||||
"""
|
||||
key = name.strip()
|
||||
reason = reason_unsupported(key)
|
||||
if reason:
|
||||
raise ValueError(reason)
|
||||
if repo.get(key) is not None:
|
||||
raise ValueError(f"字段「{key}」已在字段库中:直接编辑它,或给它改个中文名/含义即可")
|
||||
d = get_field(key)
|
||||
assert d is not None # reason_unsupported 为空 ⇒ 注册表必有此字段
|
||||
chosen = _check_unit(key, unit)
|
||||
item = ConditionField(
|
||||
name=key,
|
||||
label=(label.strip() or d.label),
|
||||
description=(description.strip() or d.description),
|
||||
kind=d.kind, # 类型来自引擎,不接受调用方声明
|
||||
group_name=(group_name.strip() or d.group_name),
|
||||
unit=chosen,
|
||||
source="custom",
|
||||
enabled=enabled,
|
||||
sort_order=_sort_order(d),
|
||||
)
|
||||
saved = repo.save(item)
|
||||
session.commit()
|
||||
return saved
|
||||
|
||||
|
||||
def update_field(
|
||||
repo: ConditionFieldRepository,
|
||||
session,
|
||||
name: str,
|
||||
*,
|
||||
label: str | None = None,
|
||||
description: str | None = None,
|
||||
group_name: str | None = None,
|
||||
unit: str | None = None,
|
||||
enabled: bool | None = None,
|
||||
) -> ConditionField:
|
||||
"""编辑字段文案/分组/**界面单位**/启用状态。
|
||||
|
||||
``name`` / ``kind`` / ``source`` 不可改:前者是引擎字段名,后两者是引擎事实与
|
||||
条目来历 —— 允许改会让「字段库」与引擎脱节(§24 不做假支持)。
|
||||
``unit`` 可改但**只能在注册表给的阶梯里选**(如 万元 ⇄ 亿元):它是界面单位,
|
||||
提交/回显按固定系数换算,引擎始终用基准单位 —— 所以改它不会让历史策略变义。
|
||||
"""
|
||||
item = repo.get(name)
|
||||
if item is None:
|
||||
raise KeyError(name)
|
||||
patch = item.model_copy(
|
||||
update={
|
||||
k: v
|
||||
for k, v in {
|
||||
"label": label,
|
||||
"description": description,
|
||||
"group_name": group_name,
|
||||
"unit": _check_unit(name, unit) if unit is not None else None,
|
||||
"enabled": enabled,
|
||||
}.items()
|
||||
if v is not None
|
||||
}
|
||||
)
|
||||
if not patch.label.strip():
|
||||
raise ValueError("中文名不能为空(下拉里要显示它)")
|
||||
saved = repo.save(patch)
|
||||
session.commit()
|
||||
return saved
|
||||
|
||||
|
||||
def delete_field(repo: ConditionFieldRepository, session, name: str) -> None:
|
||||
"""删除自定义字段。内置字段不允许删除 —— 删了下次 seed 又会补回来,只会让人困惑。"""
|
||||
item = repo.get(name)
|
||||
if item is None:
|
||||
raise KeyError(name)
|
||||
if item.source == "builtin":
|
||||
raise ValueError(
|
||||
f"「{name}」是内置字段,不能删除(删掉下次读取也会自动补回)。"
|
||||
"如果不想在条件里看到它,请改为「停用」。"
|
||||
)
|
||||
repo.delete(name)
|
||||
session.commit()
|
||||
@@ -0,0 +1,811 @@
|
||||
"""数据同步服务:增量 + 新浪「两边一致」校验兜底(financial / daily)。
|
||||
|
||||
背景(AGENT.md §5/§7/§8):
|
||||
- Tushare 是首选源;新浪财经只作备用。任何切源都必须可追溯(写 sync_log),
|
||||
且禁止静默把未经核验的备用源数据并入主库。
|
||||
- 本模块把「切到新浪」从 FailoverProvider 的『主源报错即兜底』收紧为
|
||||
『校验兜底』:只有当某只股票**两边重叠的历史数据一致**时,才允许把新浪
|
||||
的**新数据**(本地缺失键的行)导入;无本地历史或校验不一致 → 拒绝并告警,
|
||||
留待 Tushare 恢复后重跑补齐(数据真实性优先)。
|
||||
|
||||
校验口径(经验证,见仓库数据):
|
||||
- 财务可比字段只有 eps / gross_margin —— 两源同报告期数值逐位一致;
|
||||
ROE 两边口径不同(Tushare 摊薄 vs 新浪加权),不作为一致性依据。
|
||||
- 日线新浪为前复权,与本地不复权行仅「最近无除权区间」相等,因此只拿
|
||||
两源重叠的最近若干个交易日做一致性校验(通道可信 → 才允许补缺)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import date, timedelta
|
||||
from decimal import Decimal
|
||||
|
||||
from app.domain.entities.market import DailyBar, FinancialIndicator, SyncLog
|
||||
from app.domain.providers import MarketDataProvider
|
||||
from app.domain.repositories.market import (
|
||||
AdjustFactorRepository,
|
||||
DailyBarRepository,
|
||||
FinancialRepository,
|
||||
)
|
||||
from app.infrastructure.data_sources.errors import DataSourceAuthenticationError
|
||||
|
||||
# 财务两源可比字段(其余字段两端口径不一致 / 单侧缺失,不能作校验依据)
|
||||
FINANCIAL_COMPARE_FIELDS = ("eps", "gross_margin")
|
||||
DAILY_COMPARE_FIELDS = ("open", "high", "low", "close")
|
||||
|
||||
# 新浪日 K 可达窗口(getKLineData datalen=320 自然日)
|
||||
SINA_KLINE_DAYS = 320
|
||||
# 校验回看:请求新浪时额外回看 begin 之前的天数,确保与本地近期历史有重叠可比
|
||||
SINA_VERIFY_LOOKBACK_DAYS = 45
|
||||
_EPOCH = date(1990, 1, 1)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- 一致性校验(纯函数)
|
||||
|
||||
@dataclass
|
||||
class OverlapVerdict:
|
||||
"""两边重叠一致性结论。ok=True 才允许导入新浪新数据。"""
|
||||
|
||||
ok: bool
|
||||
shared: int = 0 # 重叠报告期 / 重叠交易日数量
|
||||
compared: int = 0 # 实际参与数值比较的行/日数量
|
||||
mismatches: list[str] = field(default_factory=list)
|
||||
|
||||
def summary(self) -> str:
|
||||
if self.ok:
|
||||
return f"重叠 {self.shared} 项,数值一致(比较 {self.compared} 项)"
|
||||
why = f"重叠 {self.shared} 项不足/为空"
|
||||
if self.mismatches:
|
||||
why = ";".join(self.mismatches[:3])
|
||||
return f"校验未通过:{why}"
|
||||
|
||||
|
||||
def _close_enough(a: Decimal, b: Decimal, *, rel_tol: float, abs_tol: float) -> bool:
|
||||
if a is None or b is None:
|
||||
return False
|
||||
diff = abs(a - b)
|
||||
if diff <= Decimal(str(abs_tol)):
|
||||
return True
|
||||
scale = max(abs(a), abs(b))
|
||||
return diff <= Decimal(str(rel_tol)) * scale
|
||||
|
||||
|
||||
def financial_overlap_consistent(
|
||||
local_rows: Sequence[FinancialIndicator],
|
||||
sina_rows: Sequence[FinancialIndicator],
|
||||
*,
|
||||
min_shared: int = 2,
|
||||
rel_tol: float = 1e-4,
|
||||
abs_tol: float = 1e-3,
|
||||
) -> OverlapVerdict:
|
||||
"""新浪财务行与本地(Tushare)行按报告期重叠校验。
|
||||
|
||||
新浪每个报告期只保留最新一版(getFinanceReport2022 的 report_list 按
|
||||
报告期一份);本地同报告期可能有多版公告,取公告日最新者比较。
|
||||
要求:重叠报告期数 >= min_shared,且全部可比字段(两源都非空)一致。
|
||||
"""
|
||||
local_latest: dict[date, FinancialIndicator] = {}
|
||||
for row in local_rows:
|
||||
cur = local_latest.get(row.report_date)
|
||||
if cur is None or row.announce_date > cur.announce_date:
|
||||
local_latest[row.report_date] = row
|
||||
sina_by_report = {row.report_date: row for row in sina_rows}
|
||||
|
||||
verdict = OverlapVerdict(ok=False)
|
||||
shared_dates = sorted(set(local_latest) & set(sina_by_report), reverse=True)
|
||||
verdict.shared = len(shared_dates)
|
||||
for report in shared_dates:
|
||||
a = local_latest[report]
|
||||
b = sina_by_report[report]
|
||||
day_mismatch: list[str] = []
|
||||
compared = 0
|
||||
for f in FINANCIAL_COMPARE_FIELDS:
|
||||
va, vb = getattr(a, f), getattr(b, f)
|
||||
if va is None or vb is None:
|
||||
continue
|
||||
compared += 1
|
||||
if not _close_enough(va, vb, rel_tol=rel_tol, abs_tol=abs_tol):
|
||||
day_mismatch.append(f"{report}: {f} {va}≠{vb}")
|
||||
verdict.compared += compared
|
||||
verdict.mismatches.extend(day_mismatch)
|
||||
verdict.ok = (
|
||||
verdict.shared >= min_shared and verdict.compared > 0 and not verdict.mismatches
|
||||
)
|
||||
return verdict
|
||||
|
||||
|
||||
def daily_overlap_consistent(
|
||||
local_bars: Sequence[DailyBar],
|
||||
sina_bars: Sequence[DailyBar],
|
||||
*,
|
||||
min_shared: int = 3,
|
||||
max_recent: int = 8,
|
||||
rel_tol: float = 1e-4,
|
||||
abs_tol: float = Decimal("0.02"),
|
||||
) -> OverlapVerdict:
|
||||
"""新浪日 K(前复权)与本地(不复权)重叠校验。
|
||||
|
||||
前复权锚定最新价:仅「最近一次除权之后」的交易日两源数值相等,因此只
|
||||
比较两源重叠的、最近的 max_recent 个交易日(此时若有除权发生在该段,
|
||||
校验会判不一致 → 拒绝兜底,安全方向)。vol/amount 两源单位/口径不同,
|
||||
不参与比较。
|
||||
"""
|
||||
local_by_day = {b.trade_date: b for b in local_bars}
|
||||
sina_by_day = {b.trade_date: b for b in sina_bars}
|
||||
shared = sorted(set(local_by_day) & set(sina_by_day), reverse=True)
|
||||
|
||||
verdict = OverlapVerdict(ok=False)
|
||||
verdict.shared = len(shared)
|
||||
for day in shared[:max_recent]:
|
||||
a, b = local_by_day[day], sina_by_day[day]
|
||||
day_mismatch: list[str] = []
|
||||
compared = 0
|
||||
for f in DAILY_COMPARE_FIELDS:
|
||||
va, vb = getattr(a, f), getattr(b, f)
|
||||
if va is None or vb is None:
|
||||
continue
|
||||
compared += 1
|
||||
if not _close_enough(va, vb, rel_tol=rel_tol, abs_tol=abs_tol):
|
||||
day_mismatch.append(f"{day}: {f} {va}≠{vb}")
|
||||
verdict.compared += compared
|
||||
verdict.mismatches.extend(day_mismatch)
|
||||
checked = len(shared[:max_recent])
|
||||
verdict.ok = (
|
||||
checked >= min_shared and verdict.compared > 0 and not verdict.mismatches
|
||||
)
|
||||
return verdict
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- 报告期披露节奏
|
||||
|
||||
def latest_expected_report_period(today: date) -> date:
|
||||
"""当前「应已披露」的最新报告期(报告期结束日)。
|
||||
|
||||
用作财务增量的已最新判断:本地已含该报告期 → 该股票已跟进到最新一季,
|
||||
跳过(避免每轮全量重拉;--full 强制)。窗口按 A 股披露节奏划分:
|
||||
- 1/1~2/14:年报季未开 → 上年三季报(09-30)
|
||||
- 2/15~6/30:年报+一季报季 → 本年一季报(03-31)
|
||||
- 7/1~10/15:半年报季 → 本年半年报(06-30)
|
||||
- 10/16~12/31:三季报季 → 本年三季报(09-30)
|
||||
"""
|
||||
y = today.year
|
||||
md = (today.month, today.day)
|
||||
if md <= (2, 14):
|
||||
return date(y - 1, 9, 30)
|
||||
if md <= (6, 30):
|
||||
return date(y, 3, 31)
|
||||
if md <= (10, 15):
|
||||
return date(y, 6, 30)
|
||||
return date(y, 9, 30)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- 审计
|
||||
|
||||
def _audit_sync(
|
||||
audit: Callable[[SyncLog], None],
|
||||
*,
|
||||
source: str,
|
||||
api: str,
|
||||
success: bool,
|
||||
row_count: int = 0,
|
||||
reason: str | None = None,
|
||||
start: date | None = None,
|
||||
end: date | None = None,
|
||||
) -> None:
|
||||
audit(
|
||||
SyncLog(
|
||||
source=source,
|
||||
api=api,
|
||||
success=success,
|
||||
failure_reason=reason,
|
||||
row_count=row_count,
|
||||
data_start=start,
|
||||
data_end=end,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- 结果结构
|
||||
|
||||
@dataclass
|
||||
class FinancialSymbolResult:
|
||||
"""单只股票财务同步结果(status: skip|ok|sina|failed)。"""
|
||||
|
||||
symbol: str
|
||||
status: str
|
||||
source: str | None = None # tushare | sina
|
||||
fetched: int = 0 # 数据源返回行数
|
||||
written: int = 0 # 实际落库行数(新增;--full 时含更新)
|
||||
updated: int = 0 # --full 下覆盖的既有行数
|
||||
report_first: date | None = None
|
||||
report_last: date | None = None
|
||||
announce_first: date | None = None
|
||||
announce_last: date | None = None
|
||||
notes: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DailySymbolResult:
|
||||
"""单只股票日线同步结果(status: skip|ok|sina|failed)。"""
|
||||
|
||||
symbol: str
|
||||
status: str
|
||||
source: str | None = None # tushare | sina
|
||||
bars_fetched: int = 0
|
||||
bars_written: int = 0
|
||||
day_first: date | None = None
|
||||
day_last: date | None = None
|
||||
factors_written: int | None = None # None=未尝试(新浪兜底无因子)
|
||||
notes: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DailyBasicDayResult:
|
||||
"""单个交易日的每日指标同步结果(status: skip|ok|failed)。"""
|
||||
|
||||
trade_date: date
|
||||
status: str
|
||||
source: str | None = None # tushare(新浪不支持本接口)
|
||||
rows_fetched: int = 0
|
||||
rows_written: int = 0
|
||||
notes: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class NameHistoryChunkResult:
|
||||
"""单个时间分片的名称变更同步结果(status: ok|failed)。"""
|
||||
|
||||
start: date
|
||||
end: date
|
||||
status: str
|
||||
source: str | None = None
|
||||
rows_fetched: int = 0
|
||||
rows_written: int = 0
|
||||
notes: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- 财务同步服务
|
||||
|
||||
class VerifiedFinancialSyncer:
|
||||
"""财务指标增量同步:Tushare 窗口化拉取 → 失败则新浪校验兜底。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
primary: MarketDataProvider,
|
||||
fallback: MarketDataProvider | None,
|
||||
repo: FinancialRepository,
|
||||
audit: Callable[[SyncLog], None],
|
||||
today: date | None = None,
|
||||
min_shared: int = 2,
|
||||
) -> None:
|
||||
self.primary = primary
|
||||
self.fallback = fallback
|
||||
self.repo = repo
|
||||
self.audit = audit
|
||||
self.today = today or date.today()
|
||||
self.min_shared = min_shared
|
||||
|
||||
def sync_symbol(self, symbol: str, *, force_full: bool = False) -> FinancialSymbolResult:
|
||||
local = self.repo.list_symbol(symbol)
|
||||
local_keys = {(r.symbol, r.report_date, r.announce_date) for r in local}
|
||||
due = latest_expected_report_period(self.today)
|
||||
if not force_full and local and any(r.report_date == due for r in local):
|
||||
return FinancialSymbolResult(
|
||||
symbol=symbol,
|
||||
status="skip",
|
||||
notes=[f"本地已含最新报告期 {due.isoformat()},跳过(--full 强制重拉)"],
|
||||
)
|
||||
# 拉取窗口:有本地行则从最早本地报告期起(含更正/补缺),无则全历史;
|
||||
# 上限到最新应披露报告期。
|
||||
hi = due
|
||||
lo = min((r.report_date for r in local), default=None) or _EPOCH
|
||||
try:
|
||||
rows = self.primary.get_financial(symbol, lo, hi)
|
||||
except DataSourceAuthenticationError:
|
||||
# 凭证无效/接口无权限:属全局性故障,快速失败让用户修 token,
|
||||
# 不要对全市场逐只做无意义的新浪试探
|
||||
raise
|
||||
except Exception as exc: # noqa: BLE001 —— 与 FailoverProvider 一致,统一走审计
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.primary.name,
|
||||
api="get_financial",
|
||||
success=False,
|
||||
reason=str(exc),
|
||||
start=lo,
|
||||
end=hi,
|
||||
)
|
||||
return self._sina_fallback(symbol, local, local_keys, primary_error=str(exc))
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.primary.name,
|
||||
api="get_financial",
|
||||
success=True,
|
||||
row_count=len(rows),
|
||||
start=lo,
|
||||
end=hi,
|
||||
)
|
||||
if force_full:
|
||||
to_write = rows
|
||||
updated = sum(1 for r in rows if _fin_key(r) in local_keys)
|
||||
else:
|
||||
to_write = [r for r in rows if _fin_key(r) not in local_keys]
|
||||
updated = 0
|
||||
written = self.repo.upsert_many(to_write)
|
||||
return _fin_result(symbol, status="ok", source="tushare", written_rows=to_write,
|
||||
written=written, updated=updated)
|
||||
|
||||
# ---- 新浪校验兜底 ----
|
||||
|
||||
def _sina_fallback(
|
||||
self,
|
||||
symbol: str,
|
||||
local: list[FinancialIndicator],
|
||||
local_keys: set[tuple],
|
||||
*,
|
||||
primary_error: str,
|
||||
) -> FinancialSymbolResult:
|
||||
if self.fallback is None:
|
||||
return FinancialSymbolResult(
|
||||
symbol=symbol,
|
||||
status="failed",
|
||||
notes=[f"Tushare 失败且未配置新浪兜底: {primary_error}"],
|
||||
)
|
||||
if not local:
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.fallback.name,
|
||||
api="get_financial",
|
||||
success=False,
|
||||
reason=f"无本地历史可做两边一致性校验,跳过待 Tushare 恢复重试({primary_error})",
|
||||
)
|
||||
return FinancialSymbolResult(
|
||||
symbol=symbol,
|
||||
status="failed",
|
||||
source="sina",
|
||||
notes=[
|
||||
f"Tushare 失败且本地无历史({symbol}),无法确认新浪数据真实性,"
|
||||
f"跳过待重试。primary: {primary_error}"
|
||||
],
|
||||
)
|
||||
try:
|
||||
sina_rows = self.fallback.get_financial(symbol)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.fallback.name,
|
||||
api="get_financial",
|
||||
success=False,
|
||||
reason=f"primary: {primary_error}; fallback: {exc}",
|
||||
)
|
||||
return FinancialSymbolResult(
|
||||
symbol=symbol,
|
||||
status="failed",
|
||||
source="sina",
|
||||
notes=[f"主备数据源均失败: primary={primary_error}; sina={exc}"],
|
||||
)
|
||||
verdict = financial_overlap_consistent(local, sina_rows, min_shared=self.min_shared)
|
||||
if not verdict.ok:
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.fallback.name,
|
||||
api="get_financial",
|
||||
success=False,
|
||||
reason=f"{verdict.summary()}(新浪返回 {len(sina_rows)} 行; primary={primary_error})",
|
||||
)
|
||||
return FinancialSymbolResult(
|
||||
symbol=symbol,
|
||||
status="failed",
|
||||
source="sina",
|
||||
notes=[
|
||||
f"新浪数据与本地历史不一致/无法校验({symbol}),拒绝导入。"
|
||||
f"primary: {primary_error};{verdict.summary()}"
|
||||
],
|
||||
)
|
||||
new_rows = [r for r in sina_rows if _fin_key(r) not in local_keys]
|
||||
written = self.repo.upsert_many(new_rows)
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.fallback.name,
|
||||
api="get_financial",
|
||||
success=True,
|
||||
row_count=written,
|
||||
)
|
||||
return _fin_result(symbol, status="sina", source="sina", written_rows=new_rows,
|
||||
written=written, updated=0,
|
||||
note=f"新浪校验通过后补入 {written} 行(仅本地缺失键,source=sina)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- 日线同步服务
|
||||
|
||||
class VerifiedDailySyncer:
|
||||
"""日线同步:Tushare 失败 → 新浪校验兜底(仅补缺失交易日、无复权因子)。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
primary: MarketDataProvider,
|
||||
fallback: MarketDataProvider | None,
|
||||
bars: DailyBarRepository,
|
||||
factors: AdjustFactorRepository,
|
||||
audit: Callable[[SyncLog], None],
|
||||
today: date | None = None,
|
||||
min_shared: int = 3,
|
||||
) -> None:
|
||||
self.primary = primary
|
||||
self.fallback = fallback
|
||||
self.bars = bars
|
||||
self.factors = factors
|
||||
self.audit = audit
|
||||
self.today = today or date.today()
|
||||
self.min_shared = min_shared
|
||||
|
||||
def sync_symbol(self, symbol: str, begin: date, end: date) -> DailySymbolResult:
|
||||
try:
|
||||
bars = self.primary.get_daily(symbol, begin, end)
|
||||
except DataSourceAuthenticationError:
|
||||
raise # 凭证/权限故障 → 快速失败(见财务同步注释)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.primary.name,
|
||||
api="get_daily",
|
||||
success=False,
|
||||
reason=str(exc),
|
||||
start=begin,
|
||||
end=end,
|
||||
)
|
||||
return self._sina_fallback(symbol, begin, end, primary_error=str(exc))
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.primary.name,
|
||||
api="get_daily",
|
||||
success=True,
|
||||
row_count=len(bars),
|
||||
start=begin,
|
||||
end=end,
|
||||
)
|
||||
try:
|
||||
factors = self.primary.get_adjust_factor(symbol, begin, end)
|
||||
except DataSourceAuthenticationError:
|
||||
raise # 凭证/权限故障 → 快速失败(见财务同步注释)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
# 复权因子是日线配套:缺因子不写本段,避免 resume 按日线已最新而跳过、因子永远补不上
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.primary.name,
|
||||
api="get_adjust_factor",
|
||||
success=False,
|
||||
reason=str(exc),
|
||||
start=begin,
|
||||
end=end,
|
||||
)
|
||||
return DailySymbolResult(
|
||||
symbol=symbol,
|
||||
status="failed",
|
||||
source="tushare",
|
||||
bars_fetched=len(bars),
|
||||
notes=[f"日线拉取成功但复权因子失败,本段未落库(防因子缺口): {exc}"],
|
||||
)
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.primary.name,
|
||||
api="get_adjust_factor",
|
||||
success=True,
|
||||
row_count=len(factors),
|
||||
start=begin,
|
||||
end=end,
|
||||
)
|
||||
self.bars.upsert_many(bars)
|
||||
self.factors.upsert_many(factors)
|
||||
return DailySymbolResult(
|
||||
symbol=symbol,
|
||||
status="ok",
|
||||
source="tushare",
|
||||
bars_fetched=len(bars),
|
||||
bars_written=len(bars),
|
||||
factors_written=len(factors),
|
||||
day_first=min((b.trade_date for b in bars), default=None),
|
||||
day_last=max((b.trade_date for b in bars), default=None),
|
||||
)
|
||||
|
||||
# ---- 新浪校验兜底 ----
|
||||
|
||||
def _sina_fallback(
|
||||
self, symbol: str, begin: date, end: date, *, primary_error: str
|
||||
) -> DailySymbolResult:
|
||||
if self.fallback is None:
|
||||
return DailySymbolResult(
|
||||
symbol=symbol,
|
||||
status="failed",
|
||||
notes=[f"Tushare 失败且未配置新浪兜底: {primary_error}"],
|
||||
)
|
||||
# 新浪只有最近 SINA_KLINE_DAYS 自然日数据;为拿到「本地近期历史」重叠做
|
||||
# 校验,请求窗口需回看 begin 之前 SINA_VERIFY_LOOKBACK_DAYS 天
|
||||
# (见 daily_overlap_consistent:只比较两源重叠的最近交易日)。
|
||||
q_start = max(
|
||||
self.today - timedelta(days=SINA_KLINE_DAYS - 1),
|
||||
begin - timedelta(days=SINA_VERIFY_LOOKBACK_DAYS),
|
||||
)
|
||||
try:
|
||||
sina_bars = self.fallback.get_daily(symbol, q_start, end)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.fallback.name,
|
||||
api="get_daily",
|
||||
success=False,
|
||||
reason=f"primary: {primary_error}; fallback: {exc}",
|
||||
start=q_start,
|
||||
end=end,
|
||||
)
|
||||
return DailySymbolResult(
|
||||
symbol=symbol,
|
||||
status="failed",
|
||||
source="sina",
|
||||
notes=[f"主备数据源均失败: primary={primary_error}; sina={exc}"],
|
||||
)
|
||||
local_recent = self.bars.get_range(symbol, q_start, end)
|
||||
verdict = daily_overlap_consistent(local_recent, sina_bars, min_shared=self.min_shared)
|
||||
if not verdict.ok:
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.fallback.name,
|
||||
api="get_daily",
|
||||
success=False,
|
||||
reason=f"{verdict.summary()}(新浪返回 {len(sina_bars)} 行; primary={primary_error})",
|
||||
start=q_start,
|
||||
end=end,
|
||||
)
|
||||
return DailySymbolResult(
|
||||
symbol=symbol,
|
||||
status="failed",
|
||||
source="sina",
|
||||
bars_fetched=len(sina_bars),
|
||||
notes=[
|
||||
f"新浪数据与本地历史不一致/无法校验({symbol}),拒绝兜底补缺。"
|
||||
f"primary: {primary_error};{verdict.summary()}"
|
||||
],
|
||||
)
|
||||
local_dates = {b.trade_date for b in local_recent}
|
||||
new_bars = [
|
||||
b
|
||||
for b in sina_bars
|
||||
if begin <= b.trade_date <= end and b.trade_date not in local_dates
|
||||
]
|
||||
written = self.bars.upsert_many(new_bars)
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.fallback.name,
|
||||
api="get_daily",
|
||||
success=True,
|
||||
row_count=written,
|
||||
start=begin,
|
||||
end=end,
|
||||
)
|
||||
return DailySymbolResult(
|
||||
symbol=symbol,
|
||||
status="sina",
|
||||
source="sina",
|
||||
bars_fetched=len(sina_bars),
|
||||
bars_written=written,
|
||||
day_first=min((b.trade_date for b in new_bars), default=None),
|
||||
day_last=max((b.trade_date for b in new_bars), default=None),
|
||||
factors_written=None,
|
||||
notes=[
|
||||
f"新浪校验通过,仅补本地缺失交易日 {written} 根(前复权 source=sina,"
|
||||
f"无复权因子;Tushare 恢复后 --resume 会按日覆盖回不复权口径)"
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- 每日指标同步服务
|
||||
|
||||
class DailyBasicSyncer:
|
||||
"""每日指标(估值 / 股息率 / 市值)同步:按交易日整表拉取。
|
||||
|
||||
设计(与财务/日线的「逐股校验兜底」不同):
|
||||
- daily_basic 是**横截面整表**接口,Tushare 一次返回当日全市场,无法逐股兜底;
|
||||
- 新浪不提供本接口 → 主源失败时**如实失败并写 sync_log**,禁止静默留缺口;
|
||||
- 幂等键 (symbol, trade_date):重跑同日无副作用(upsert)。
|
||||
- 交易日集合由日历仓储给出:只同步 is_open 且本地缺失的日期
|
||||
(`missing_dates`),因此增量与断点续跑天然安全。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
primary: MarketDataProvider,
|
||||
repo,
|
||||
audit: Callable[[SyncLog], None],
|
||||
batch_flush: int = 5,
|
||||
) -> None:
|
||||
self.primary = primary
|
||||
self.repo = repo
|
||||
self.audit = audit
|
||||
self.batch_flush = max(batch_flush, 1)
|
||||
|
||||
def sync_day(self, trade_date: date) -> DailyBasicDayResult:
|
||||
try:
|
||||
rows = self.primary.get_daily_basic(trade_date)
|
||||
except DataSourceAuthenticationError:
|
||||
raise # 凭证/权限故障 → 快速失败(不要逐日重试把配额烧光)
|
||||
except Exception as exc: # noqa: BLE001 —— 逐日失败不中断整段
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.primary.name,
|
||||
api="get_daily_basic",
|
||||
success=False,
|
||||
reason=str(exc),
|
||||
start=trade_date,
|
||||
end=trade_date,
|
||||
)
|
||||
return DailyBasicDayResult(
|
||||
trade_date=trade_date,
|
||||
status="failed",
|
||||
source=self.primary.name,
|
||||
notes=[f"拉取失败: {exc}"],
|
||||
)
|
||||
# 过滤掉非法行(ts_code 缺失)—— 避免脏键污染幂等
|
||||
rows = [r for r in rows if r.symbol and r.symbol != "None"]
|
||||
written = self.repo.upsert_many(rows)
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.primary.name,
|
||||
api="get_daily_basic",
|
||||
success=True,
|
||||
row_count=len(rows),
|
||||
start=trade_date,
|
||||
end=trade_date,
|
||||
)
|
||||
return DailyBasicDayResult(
|
||||
trade_date=trade_date,
|
||||
status="ok",
|
||||
source=self.primary.name,
|
||||
rows_fetched=len(rows),
|
||||
rows_written=written,
|
||||
)
|
||||
|
||||
def sync_range(
|
||||
self,
|
||||
start: date,
|
||||
end: date,
|
||||
*,
|
||||
on_progress: Callable[[int, int, date], None] | None = None,
|
||||
should_stop: Callable[[], bool] | None = None,
|
||||
) -> list[DailyBasicDayResult]:
|
||||
"""补齐 [start, end] 内开市但本地缺失的交易日。
|
||||
|
||||
on_progress(done, total, trade_date):进度回调(CLI 打印 / 日志)。
|
||||
should_stop():返回 True 时提前收尾(优雅中断,已落库的行保持有效)。
|
||||
"""
|
||||
days = self.repo.missing_dates(start, end)
|
||||
results: list[DailyBasicDayResult] = []
|
||||
total = len(days)
|
||||
for idx, day in enumerate(days, start=1):
|
||||
if should_stop is not None and should_stop():
|
||||
break
|
||||
results.append(self.sync_day(day))
|
||||
if on_progress is not None:
|
||||
on_progress(idx, total, day)
|
||||
return results
|
||||
|
||||
|
||||
class NameHistorySyncer:
|
||||
"""股票名称变更历史同步(Tushare namechange)—— 时点 ST 判定的数据基础。
|
||||
|
||||
设计要点:
|
||||
- **按年分片**:namechange 支持区间批量查询(2020+ 仅 4031 行),但全历史
|
||||
(1990 起)会触及单次 6000 行上限被**静默截断**,因此按自然年分片调用,
|
||||
每片独立审计与计数,超限时 Provider 会告警。
|
||||
- **幂等键 (symbol, start_date)**:重跑无副作用(upsert),可安全续跑。
|
||||
- 新浪不提供本接口 → 主源失败如实写 sync_log(禁止静默留缺口,AGENT.md §24)。
|
||||
- 为什么值得同步:`stock.name` 只是最新名称快照,用最新名称做 `exclude_st`
|
||||
会把「曾为高股息、后变 ST/退市」的股息陷阱样本整段排除 ——
|
||||
实测影响约 3.70pp 收益(见 docs/DEV_PLAN_DIVIDEND_BACKTEST.md §10.5)。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
primary: MarketDataProvider,
|
||||
repo,
|
||||
audit: Callable[[SyncLog], None],
|
||||
) -> None:
|
||||
self.primary = primary
|
||||
self.repo = repo
|
||||
self.audit = audit
|
||||
|
||||
def sync_chunk(self, start: date, end: date) -> NameHistoryChunkResult:
|
||||
try:
|
||||
rows = self.primary.get_name_changes(start, end)
|
||||
except DataSourceAuthenticationError:
|
||||
raise # 凭证/权限故障 → 快速失败,不要逐片重试烧配额
|
||||
except Exception as exc: # noqa: BLE001 —— 单片失败不中断整段
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.primary.name,
|
||||
api="get_namechange",
|
||||
success=False,
|
||||
reason=str(exc),
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
return NameHistoryChunkResult(
|
||||
start=start, end=end, status="failed",
|
||||
source=self.primary.name, notes=[f"拉取失败: {exc}"],
|
||||
)
|
||||
rows = [r for r in rows if r.symbol and r.symbol != "None"]
|
||||
written = self.repo.upsert_many(rows)
|
||||
_audit_sync(
|
||||
self.audit,
|
||||
source=self.primary.name,
|
||||
api="get_namechange",
|
||||
success=True,
|
||||
row_count=len(rows),
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
return NameHistoryChunkResult(
|
||||
start=start, end=end, status="ok", source=self.primary.name,
|
||||
rows_fetched=len(rows), rows_written=written,
|
||||
)
|
||||
|
||||
def sync_range(
|
||||
self,
|
||||
start: date,
|
||||
end: date,
|
||||
*,
|
||||
on_progress: Callable[[int, int, date], None] | None = None,
|
||||
) -> list[NameHistoryChunkResult]:
|
||||
"""按自然年分片同步 [start, end](每片 1 次 API 调用)。"""
|
||||
chunks: list[tuple[date, date]] = []
|
||||
cursor = date(start.year, 1, 1)
|
||||
while cursor <= end:
|
||||
chunk_end = min(date(cursor.year, 12, 31), end)
|
||||
chunks.append((max(cursor, start), chunk_end))
|
||||
cursor = date(cursor.year + 1, 1, 1)
|
||||
results: list[NameHistoryChunkResult] = []
|
||||
total = len(chunks)
|
||||
for idx, (cs, ce) in enumerate(chunks, start=1):
|
||||
results.append(self.sync_chunk(cs, ce))
|
||||
if on_progress is not None:
|
||||
on_progress(idx, total, cs)
|
||||
return results
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- 小工具
|
||||
|
||||
def _fin_key(row: FinancialIndicator) -> tuple:
|
||||
return (row.symbol, row.report_date, row.announce_date)
|
||||
|
||||
|
||||
def _fin_result(
|
||||
symbol: str,
|
||||
*,
|
||||
status: str,
|
||||
source: str,
|
||||
written_rows: Sequence[FinancialIndicator],
|
||||
written: int,
|
||||
updated: int,
|
||||
note: str | None = None,
|
||||
) -> FinancialSymbolResult:
|
||||
reports = [r.report_date for r in written_rows]
|
||||
announces = [r.announce_date for r in written_rows]
|
||||
notes = [note] if note else []
|
||||
return FinancialSymbolResult(
|
||||
symbol=symbol,
|
||||
status=status,
|
||||
source=source,
|
||||
fetched=len(written_rows),
|
||||
written=written,
|
||||
updated=updated,
|
||||
report_first=min(reports, default=None),
|
||||
report_last=max(reports, default=None),
|
||||
announce_first=min(announces, default=None),
|
||||
announce_last=max(announces, default=None),
|
||||
notes=notes,
|
||||
)
|
||||
@@ -0,0 +1,288 @@
|
||||
"""Experiment 归档服务(AGENT.md §21「研究可复现」):同步与异步两条路径共用。
|
||||
|
||||
谁在用:
|
||||
- 异步 Job:`job_executor._execute_inner` 执行成功后归档;
|
||||
- 同步研究端点:`api/research.py` 的 `POST /backtests`、`POST /factor-tests`
|
||||
(Phase 3 遗留的同步路径,此前只塞进程内存 `_LAST_*`,重启即丢,完全不落库)。
|
||||
|
||||
为什么单独成模块:`job_executor.py` 顶部依赖较重(研究服务、子进程调度),
|
||||
把归档抽到此处既避免 `api/research.py` → `job_executor` 的循环导入风险,也让
|
||||
「一次研究结果如何落库」只有一份实现(AGENT.md §40 简单可替换优先)。
|
||||
|
||||
归档内容:spec_json(复现依据)、完整结果 JSON、summary_text(人读摘要)、
|
||||
code_version(git 短 rev)、**data_version(真实数据快照指纹,见 data_version 模块)**、
|
||||
job_id(可追溯执行记录)、created_at。同一份结果**只在 experiment 存一份**:
|
||||
|
||||
- 异步 Job 路径下 `job.result_json` 不再重复写入(见 `job_executor`),
|
||||
`GET /api/jobs/{id}` 经 `job.experiment_id` 回读 experiment;
|
||||
- 同步端点路径下不存在 Job 记录,`job_id=None`。
|
||||
|
||||
大小护栏(诚实 + 不炸库):`experiment.result_json` 是 MEDIUMTEXT,**上限
|
||||
16,777,215 字节**(MySQL 的 MEDIUMTEXT 按字节计,utf8mb4 下中文占 3 字节),
|
||||
故预算按 UTF-8 字节执行,默认 12,000,000 字节(约上限的 71.5%,给行内其它
|
||||
列、SQL 协议与未预料字段留余量)。超预算时按 |期末收益| 降序裁剪
|
||||
`symbol_curves`,并把证据写进结果的 `archive_meta`(机器可读)与
|
||||
`unimplemented`(人可读)——绝不静默丢弃。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import subprocess
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from app.domain.entities.research import ExperimentRecord
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# backend/app/application/services/experiment_archive.py → parents[4] = 项目根
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[4]
|
||||
|
||||
# 归档结果 JSON 的字节预算(可被 config research.archive_max_chars 覆盖)
|
||||
DEFAULT_ARCHIVE_MAX_CHARS = 12_000_000
|
||||
|
||||
# MEDIUMTEXT 硬上限(字节),仅用于计算/说明默认预算的余量
|
||||
MEDIUMTEXT_MAX_BYTES = 16_777_215
|
||||
|
||||
|
||||
def new_id(prefix: str) -> str:
|
||||
"""生成 `PREFIX-XXXXXXXX` 形式的业务主键(Job/Experiment/Selection 等共用)。"""
|
||||
return f"{prefix}-{uuid.uuid4().hex[:8].upper()}"
|
||||
|
||||
|
||||
def _git_short_rev() -> str | None:
|
||||
"""当前代码版本(git 短 rev);无 git / 超时则如实返回 None。"""
|
||||
try:
|
||||
out = subprocess.run(
|
||||
["git", "rev-parse", "--short", "HEAD"],
|
||||
cwd=PROJECT_ROOT,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=5,
|
||||
check=False,
|
||||
)
|
||||
return out.stdout.strip() or None
|
||||
except Exception: # noqa: BLE001
|
||||
return None
|
||||
|
||||
|
||||
def _summary_text(kind: str, result) -> str | None:
|
||||
"""人读摘要(列表页展示);未知 kind 返回 None 而不是编造。"""
|
||||
from app.domain.entities.research import BacktestResult, FactorTestReport
|
||||
from app.domain.entities.selection import SelectionResult
|
||||
|
||||
if kind == "backtest" and isinstance(result, BacktestResult):
|
||||
s = result.summary
|
||||
return (
|
||||
f"总收益 {s.total_return_pct:.2f}% · 年化 {s.annual_return_pct:.2f}% · "
|
||||
f"回撤 {s.max_drawdown_pct:.2f}%"
|
||||
)
|
||||
if isinstance(result, FactorTestReport):
|
||||
return (
|
||||
f"IC {result.ic_mean:.4f} · RankIC {result.rank_ic_mean:.4f} · "
|
||||
f"样本 {result.sample_days} 日"
|
||||
)
|
||||
if kind == "selection" and isinstance(result, SelectionResult):
|
||||
return (
|
||||
f"as_of {result.as_of_date} · 选出 {result.statistics.selected} / "
|
||||
f"评估 {result.statistics.evaluated}"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _archive_budget_bytes() -> int:
|
||||
"""归档字节预算:config `research.archive_max_chars`(名义字符数)> 代码默认。"""
|
||||
from app.core.config import get_settings
|
||||
|
||||
try:
|
||||
configured = get_settings().research_archive_max_chars
|
||||
except Exception as exc: # noqa: BLE001 —— 配置读不到不能阻断归档
|
||||
logger.warning("归档预算配置读取失败,用默认值:%s: %s", type(exc).__name__, exc)
|
||||
configured = None
|
||||
value = int(configured) if configured else DEFAULT_ARCHIVE_MAX_CHARS
|
||||
return value if value > 0 else DEFAULT_ARCHIVE_MAX_CHARS
|
||||
|
||||
|
||||
def _curve_sort_key(curve: dict):
|
||||
"""个股曲线排序键:|期末收益| 降序(与引擎默认输出顺序一致,保证可预期)。"""
|
||||
try:
|
||||
return abs(float(curve.get("final_return_pct") or 0.0))
|
||||
except (TypeError, ValueError):
|
||||
return 0.0
|
||||
|
||||
|
||||
def _truncation_note(stored: int, total: int, budget_bytes: int) -> str:
|
||||
"""裁剪说明(人可读):必须写清「实际存了多少 / 共多少 / 为什么」。"""
|
||||
return (
|
||||
f"归档体积超过预算({budget_bytes:,} 字节,MEDIUMTEXT 上限 16MB):"
|
||||
f"个股收益曲线按 |期末收益| 降序仅存 {stored} / 共 {total} 只"
|
||||
"(完整明细见 trades / signal_history;archive_meta.truncated=true)"
|
||||
)
|
||||
|
||||
|
||||
def _over_budget_note(budget_bytes: int, curves_total: int) -> str:
|
||||
"""「曲线全裁掉仍超预算」的说明:主体(trades/positions/signal_history 等)
|
||||
无法裁剪,如实标注而不是假装达标。"""
|
||||
return (
|
||||
f"归档主体(trades / positions / signal_history 等,期内持有 {curves_total} 只)"
|
||||
f"已超过配置的体积预算({budget_bytes:,} 字节):个股收益曲线裁剪至 0 只仍无法达标,"
|
||||
"本次仍尝试完整落库;若 MEDIUMTEXT 写入失败会如实报错(archive_meta.over_budget=true)"
|
||||
)
|
||||
|
||||
|
||||
def _fit_payload(payload: dict, *, budget_bytes: int) -> tuple[dict, str]:
|
||||
"""给结果 payload 写入 `archive_meta`,必要时按预算裁剪 symbol_curves。
|
||||
|
||||
返回 (payload, 已序列化的 result_json)。裁剪判定全部作用于「最终字符串的
|
||||
UTF-8 字节数」,不是估算:每次候选都真实序列化一次(超预算只发生在
|
||||
异常巨大的归档上,此时多几次序列化的代价可接受)。
|
||||
"""
|
||||
has_curves = isinstance(payload.get("symbol_curves"), list)
|
||||
curves: list[dict] = list(payload.get("symbol_curves") or [])
|
||||
total = len(curves)
|
||||
ordered = sorted(curves, key=_curve_sort_key, reverse=True)
|
||||
notes = list(payload.get("unimplemented") or [])
|
||||
|
||||
def build(kept: list[dict], *, over_budget: bool = False) -> tuple[dict, str, dict]:
|
||||
"""按 kept 组装候选 payload 并序列化(含 archive_meta 自收敛)。"""
|
||||
truncated = has_curves and len(kept) < total
|
||||
candidate = dict(payload)
|
||||
extra_notes: list[str] = []
|
||||
if has_curves:
|
||||
candidate["symbol_curves"] = kept
|
||||
if truncated:
|
||||
extra_notes.append(_truncation_note(len(kept), total, budget_bytes))
|
||||
if over_budget:
|
||||
extra_notes.append(_over_budget_note(budget_bytes, total))
|
||||
if extra_notes:
|
||||
candidate["unimplemented"] = [*notes, *extra_notes]
|
||||
meta = {
|
||||
"budget_chars": budget_bytes,
|
||||
"budget_bytes": budget_bytes,
|
||||
"over_budget": over_budget,
|
||||
"result_chars": 0,
|
||||
"result_bytes": 0,
|
||||
}
|
||||
if has_curves:
|
||||
meta = {
|
||||
"curves_stored": len(kept),
|
||||
"curves_total": total,
|
||||
"truncated": truncated,
|
||||
**meta,
|
||||
}
|
||||
text = ""
|
||||
# result_chars/result_bytes 会改变自身长度,迭代至收敛(通常 2 轮内)
|
||||
for _ in range(8):
|
||||
meta["result_chars"] = len(text)
|
||||
meta["result_bytes"] = len(text.encode("utf-8"))
|
||||
candidate["archive_meta"] = meta
|
||||
text = json.dumps(candidate, ensure_ascii=False)
|
||||
if meta["result_chars"] == len(text) and meta["result_bytes"] == len(
|
||||
text.encode("utf-8")
|
||||
):
|
||||
break
|
||||
return candidate, text, meta
|
||||
|
||||
if not has_curves: # 因子测试 / 选股等无曲线结果:不裁剪,只记录体积
|
||||
fitted, text, meta = build([])
|
||||
if meta["result_bytes"] > budget_bytes:
|
||||
# 无曲线可裁:如实标注超预算,仍尝试落库(写入失败会如实报错)
|
||||
logger.warning(
|
||||
"归档超过预算且无曲线可裁剪:%d 字节 > %d 字节", meta["result_bytes"], budget_bytes
|
||||
)
|
||||
fitted, text, meta = build([], over_budget=True)
|
||||
return fitted, text
|
||||
|
||||
fitted, text, meta = build(ordered)
|
||||
if meta["result_bytes"] <= budget_bytes:
|
||||
return fitted, text
|
||||
|
||||
# 超预算:二分找「能放进预算的最大曲线数」(降序保留收益绝对值最大的那些)
|
||||
lo, hi, best = 0, total, 0
|
||||
while lo <= hi:
|
||||
mid = (lo + hi) // 2
|
||||
_, candidate_text, _ = build(ordered[:mid])
|
||||
if len(candidate_text.encode("utf-8")) <= budget_bytes:
|
||||
best = mid
|
||||
lo = mid + 1
|
||||
else:
|
||||
hi = mid - 1
|
||||
fitted, text, meta = build(ordered[:best])
|
||||
if meta["result_bytes"] > budget_bytes:
|
||||
# 连 0 条曲线都放不下:主体(trades/positions/…)超预算,如实标注
|
||||
fitted, text, meta = build(ordered[:best], over_budget=True)
|
||||
logger.warning(
|
||||
"归档超过预算:symbol_curves 裁剪为 %d/%d(预算 %d 字节,实际 %d 字节)",
|
||||
meta.get("curves_stored"),
|
||||
total,
|
||||
budget_bytes,
|
||||
meta.get("result_bytes"),
|
||||
)
|
||||
return fitted, text
|
||||
|
||||
|
||||
def archive_experiment(
|
||||
*,
|
||||
session,
|
||||
kind: str,
|
||||
spec_json: str,
|
||||
result,
|
||||
job_id: str | None = None,
|
||||
experiment_repo=None,
|
||||
) -> ExperimentRecord:
|
||||
"""把一次研究结果落为 Experiment 归档,返回归档记录(已 save + commit)。
|
||||
|
||||
参数:
|
||||
- `session`:SQLAlchemy Session(由调用方持有生命周期;本函数负责 commit);
|
||||
- `kind`:`backtest` / `factor_test` / `selection`;
|
||||
- `spec_json`:完整研究 spec(复现依据);
|
||||
- `result`:pydantic 结果对象(`model_dump(mode="json")` 序列化);
|
||||
- `job_id`:来源 Job(同步端点没有 Job,留 None);
|
||||
- `experiment_repo`:ExperimentRepository(不传则用 SqlAlchemy 实现;测试可注入假仓储)。
|
||||
|
||||
`data_version` 由 `compute_data_version(session)` 计算(真实数据指纹,取不到则
|
||||
如实降级为 `unavailable`,见该模块 docstring)。
|
||||
"""
|
||||
from app.infrastructure.persistence.sqlalchemy.data_version import compute_data_version
|
||||
|
||||
payload = result.model_dump(mode="json")
|
||||
budget_bytes = _archive_budget_bytes()
|
||||
payload, result_json = _fit_payload(payload, budget_bytes=budget_bytes)
|
||||
meta = payload.get("archive_meta") or {}
|
||||
|
||||
# 把归档元数据同步写回结果对象:同步端点的响应体因此也如实标注裁剪/超预算
|
||||
#(AGENT.md §24:不假装支持 / 不静默降级)。没有该字段的结果类型跳过。
|
||||
if hasattr(result, "archive_meta"):
|
||||
result.archive_meta = meta
|
||||
archived_notes = payload.get("unimplemented")
|
||||
if isinstance(archived_notes, list) and hasattr(result, "unimplemented"):
|
||||
# 归档 JSON 里的 unimplemented 才是权威(含归档侧追加的裁剪/超预算说明),
|
||||
# 让同步响应体与归档内容保持一字不差
|
||||
result.unimplemented = list(archived_notes)
|
||||
|
||||
experiment = ExperimentRecord(
|
||||
id=new_id("EXP"),
|
||||
kind=kind,
|
||||
spec_json=spec_json,
|
||||
result_json=result_json,
|
||||
summary_text=_summary_text(kind, result),
|
||||
code_version=_git_short_rev(),
|
||||
data_version=compute_data_version(session),
|
||||
job_id=job_id,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
repo = experiment_repo if experiment_repo is not None else _default_experiment_repo(session)
|
||||
repo.save(experiment)
|
||||
session.commit()
|
||||
return experiment
|
||||
|
||||
|
||||
def _default_experiment_repo(session):
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyExperimentRepository,
|
||||
)
|
||||
|
||||
return SqlAlchemyExperimentRepository(session)
|
||||
@@ -0,0 +1,147 @@
|
||||
"""因子目录用例:注册表投影 + 参数化实例的新建/停用(M7.1 起,2026-10 参数化)。
|
||||
|
||||
## 目录与代码的分工(这是本模块的核心规矩)
|
||||
|
||||
- **能不能算** = 代码注册表(`quant/factors.py`)唯一决定。库里的行只要能解析出
|
||||
`(模板, 参数)` 就算得出来;解析不出来的行(历史手工登记)保留但在目录里标
|
||||
`resolvable=False`,引用时抛 `FactorError`(不假装支持)。
|
||||
- **有哪些因子** = 目录(DB)。内置实例由代码投影进来;**参数化实例**(如
|
||||
`momentum(window=90,direction=higher_is_better)`)由人在目录里创建 —— 参数写在名字里,
|
||||
所以「同一个模板的多个参数版本」天然并存,且任何一个都冻结了自己的口径。
|
||||
- **口径文案**(description/formula/brief/frequency/lookback/direction/requires)永远按
|
||||
代码收敛:只要这行算得出来,它的文案就是引擎的文案。手改会被改回 —— 因为
|
||||
「文档写一套、代码跑另一套」是本仓库明令禁止的;要改口径就改代码。
|
||||
- **唯一人配的字段**是 `enabled`(是否出现在因子下拉里):只对参数化实例生效;
|
||||
内置实例的开关同样由代码收敛(恒 True)。停用不影响已引用它的策略/归档解析 ——
|
||||
历史口径不能被一个开关改义。
|
||||
|
||||
稳态(目录 == 代码投影)零写入,不会每次 GET 都刷库。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
from app.domain.entities.factor import FactorDefinition
|
||||
from app.domain.repositories.factor import FactorRepository
|
||||
from app.quant.factors import (
|
||||
FactorError,
|
||||
build_factor_def,
|
||||
canonical_key,
|
||||
get_template,
|
||||
is_resolvable,
|
||||
list_factors,
|
||||
resolve_factor,
|
||||
)
|
||||
|
||||
# 「引擎口径」字段:一旦与代码不一致,目录就在撒谎,必须按代码改回。
|
||||
# enabled 也在其中 —— 但它只对**注册表实例**收敛(见 sync_registry_factors)。
|
||||
REGISTRY_FIELDS = (
|
||||
"description",
|
||||
"formula",
|
||||
"brief",
|
||||
"frequency",
|
||||
"lookback",
|
||||
"direction",
|
||||
"requires",
|
||||
"enabled",
|
||||
)
|
||||
|
||||
# 参数化实例同样要收敛的字段(不含 enabled:开关是人配的)。
|
||||
_ENGINE_FIELDS = tuple(f for f in REGISTRY_FIELDS if f != "enabled")
|
||||
|
||||
|
||||
def _registry_names() -> set[str]:
|
||||
return {d.name for d in list_factors()}
|
||||
|
||||
|
||||
def _wanted_rows(current: list[FactorDefinition]) -> list[FactorDefinition]:
|
||||
"""应然状态:内置实例按代码;能解析出来的参数化实例口径也按代码(开关保留)。"""
|
||||
registry = _registry_names()
|
||||
wanted = [FactorDefinition.from_registry_def(d) for d in list_factors()]
|
||||
for row in current:
|
||||
if row.name in registry:
|
||||
continue # 内置实例已由上面的投影覆盖
|
||||
try:
|
||||
defn, _fn = resolve_factor(row.name)
|
||||
except FactorError:
|
||||
continue # 算不出来的手登记行:不碰(只有人写的备注,没有可收敛的引擎口径)
|
||||
want = FactorDefinition.from_factor_def(defn, enabled=row.enabled)
|
||||
want = want.model_copy(update={"created_at": row.created_at, "version": row.version})
|
||||
wanted.append(want)
|
||||
return wanted
|
||||
|
||||
|
||||
def _differs(current: FactorDefinition | None, want: FactorDefinition, fields) -> bool:
|
||||
if current is None:
|
||||
return True
|
||||
return any(getattr(current, name) != getattr(want, name) for name in fields)
|
||||
|
||||
|
||||
def sync_registry_factors(repo: FactorRepository, session) -> int:
|
||||
"""把代码投影进目录(幂等);返回本次写入的行数,稳态为 0。"""
|
||||
current = repo.list()
|
||||
by_name = {f.name: f for f in current}
|
||||
registry = _registry_names()
|
||||
stale: list[FactorDefinition] = []
|
||||
for want in _wanted_rows(current):
|
||||
row = by_name.get(want.name)
|
||||
fields = REGISTRY_FIELDS if want.name in registry else _ENGINE_FIELDS
|
||||
if row is None or _differs(row, want, fields):
|
||||
stale.append(want)
|
||||
if not stale:
|
||||
return 0
|
||||
n = repo.upsert_many(stale)
|
||||
session.commit()
|
||||
return n
|
||||
|
||||
|
||||
def create_parameterized_factor(
|
||||
repo: FactorRepository,
|
||||
session,
|
||||
*,
|
||||
template: str,
|
||||
params: Mapping[str, Any] | None = None,
|
||||
) -> FactorDefinition:
|
||||
"""从模板 + 参数创建一个**新的参数化因子实例**(参数写在名字里,冻结口径)。
|
||||
|
||||
参数缺省项取模板默认值;越界/未知参数/重复的参数组合一律报错(不静默纠正)。
|
||||
"""
|
||||
tpl = get_template(template) # 未知模板 → FactorError
|
||||
key = canonical_key(tpl, params or {}) # 越界/未知参数 → FactorError
|
||||
if repo.get(key) is not None:
|
||||
raise ValueError(f"该参数组合的因子已存在:{key}(参数相同不会重复创建)")
|
||||
defn = build_factor_def(tpl, params or {}, name=key, source="custom")
|
||||
entity = FactorDefinition.from_factor_def(defn, enabled=True)
|
||||
repo.upsert_many([entity])
|
||||
session.commit()
|
||||
return entity
|
||||
|
||||
|
||||
def set_factor_enabled(
|
||||
repo: FactorRepository,
|
||||
session,
|
||||
*,
|
||||
name: str,
|
||||
enabled: bool,
|
||||
) -> FactorDefinition:
|
||||
"""启用/停用目录里的因子(只影响「能不能被选中」,不影响历史解析)。"""
|
||||
if name in _registry_names():
|
||||
raise ValueError(
|
||||
f"「{name}」是代码注册表里的内置因子,开关由代码决定,不能在目录里停用;"
|
||||
"如果要一个不同参数的版本,请从模板新建参数化因子。"
|
||||
)
|
||||
row = repo.get(name)
|
||||
if row is None:
|
||||
raise LookupError(f"因子 {name} 不在目录里")
|
||||
if not is_resolvable(name):
|
||||
raise ValueError(f"因子 {name} 引擎算不出来(未注册模板/参数非法),不能启用或停用")
|
||||
updated = row.model_copy(update={"enabled": enabled})
|
||||
repo.upsert_many([updated])
|
||||
session.commit()
|
||||
return updated
|
||||
|
||||
|
||||
# 兼容旧名(语义相同:把注册表同步进目录)。
|
||||
seed_registry_factors = sync_registry_factors
|
||||
@@ -0,0 +1,538 @@
|
||||
"""Job 执行编排与 Experiment 归档(Phase 4)+ 内存隔离执行(内存优化专项)。
|
||||
|
||||
execute_job 运行状态机 queued→running→(success|failed),
|
||||
成功时自动把 spec/result 存为 Experiment(含代码版本),实现「研究可复现」(AGENT.md §21)。
|
||||
独立 Session 生命周期(不依赖请求 scope),可被 FastAPI BackgroundTasks 或测试直接调用。
|
||||
|
||||
执行模式(config job.mode):
|
||||
- local —— 本进程内执行(开发 / 单测;无内存隔离)
|
||||
- subprocess —— 独立子进程执行(python -m app.cli.run_job <job_id>),
|
||||
子进程内施加 RLIMIT_AS 上限(config job.max_memory_gb),研究任务 OOM 只会
|
||||
MemoryError 失败归档,不会拖垮 API worker;父进程在子进程异常退出时补记 failed。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from app.application.services.experiment_archive import archive_experiment, new_id
|
||||
from app.domain.entities.research import (
|
||||
JobRecord,
|
||||
JobStatus,
|
||||
ResearchSpec,
|
||||
)
|
||||
from app.quant.service import ResearchService
|
||||
|
||||
# backend/app/application/services/job_executor.py → parents[4] = 项目根
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[4]
|
||||
BACKEND_ROOT = PROJECT_ROOT / "backend"
|
||||
|
||||
# 本进程内正在运行的子进程 Job(并发上限保护宿主内存)
|
||||
_active_jobs: dict[str, subprocess.Popen | None] = {}
|
||||
_active_lock = threading.Lock()
|
||||
|
||||
|
||||
def _execute_inner(
|
||||
job_id: str,
|
||||
*,
|
||||
session_factory: Callable,
|
||||
job_repo_factory: Callable,
|
||||
experiment_repo_factory: Callable,
|
||||
stock_repo_factory: Callable,
|
||||
daily_repo_factory: Callable,
|
||||
engine,
|
||||
basic_repo_factory: Callable | None = None,
|
||||
financial_repo_factory: Callable | None = None,
|
||||
index_repo_factory: Callable | None = None,
|
||||
name_repo_factory: Callable | None = None,
|
||||
) -> None:
|
||||
"""执行 Job。
|
||||
|
||||
可选工厂的语义(AGENT.md §24:不做静默降级):
|
||||
- `basic_repo_factory`:daily_basic(dv_ratio 等每日指标)—— 因子/条件引用时必需
|
||||
- 复权(qfq/hfq)折算在行情仓储 SQL 内完成,无需 adjust_factor 工厂
|
||||
- `financial_repo_factory`:财务表 —— 条件引用 fundamental.* 时必需
|
||||
- `index_repo_factory`:指数成分 —— universe.index_code 时必需
|
||||
- `name_repo_factory`:名称变更历史 —— universe.exclude_st 时点口径;未注入则
|
||||
回退最新名称(旧行为),结果如实标注残余偏差
|
||||
未注入且 spec 需要时,由 Service 抛出明确错误(而非返回空结果)。
|
||||
"""
|
||||
with session_factory() as session:
|
||||
job_repo = job_repo_factory(session)
|
||||
experiment_repo = experiment_repo_factory(session)
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
return
|
||||
job.status = JobStatus.RUNNING
|
||||
job.started_at = datetime.now()
|
||||
job_repo.update(job)
|
||||
session.commit()
|
||||
try:
|
||||
is_selection = job.kind == "selection"
|
||||
is_combo = job.kind == "combo"
|
||||
basic_repo = basic_repo_factory(session) if basic_repo_factory else None
|
||||
index_repo = index_repo_factory(session) if index_repo_factory else None
|
||||
name_repo = name_repo_factory(session) if name_repo_factory else None
|
||||
if is_combo:
|
||||
# 回测组合:解析 combo + 取齐选股策略 + 读公共配置 → ComboService.run
|
||||
from app.application.services.combo_service import ComboService
|
||||
from app.domain.entities.combo import BacktestCombo
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.combo_impl import (
|
||||
SqlAlchemyGlobalConfigRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
|
||||
SqlAlchemyStrategyRepository,
|
||||
)
|
||||
|
||||
combo = BacktestCombo.model_validate_json(job.spec_json)
|
||||
strategy_repo = SqlAlchemyStrategyRepository(session)
|
||||
strategies = []
|
||||
for sid in combo.strategy_ids:
|
||||
st = strategy_repo.get(sid)
|
||||
if st is None:
|
||||
raise ValueError(f"组合引用的选股策略 {sid} 不存在(可能已被删除)")
|
||||
strategies.append(st)
|
||||
config = SqlAlchemyGlobalConfigRepository(session).get()
|
||||
service = ComboService(
|
||||
stock_repo_factory(session),
|
||||
daily_repo_factory(session),
|
||||
index_repo=index_repo,
|
||||
basic_repo=basic_repo,
|
||||
financial_repo=(
|
||||
financial_repo_factory(session) if financial_repo_factory else None
|
||||
),
|
||||
name_repo=name_repo,
|
||||
)
|
||||
spec = combo # 仅用于下方分支判断占位;实际执行用 combo
|
||||
elif is_selection:
|
||||
from app.application.services.selection_service import SelectionService
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
|
||||
spec = SelectionQuery.model_validate_json(job.spec_json)
|
||||
service = SelectionService(
|
||||
stock_repo_factory(session),
|
||||
daily_repo_factory(session),
|
||||
financial_repo_factory(session) if financial_repo_factory else None,
|
||||
index_repo=index_repo,
|
||||
basic_repo=basic_repo,
|
||||
name_repo=name_repo,
|
||||
)
|
||||
else:
|
||||
spec = ResearchSpec.model_validate_json(job.spec_json)
|
||||
service = ResearchService(
|
||||
stock_repo_factory(session),
|
||||
daily_repo_factory(session),
|
||||
engine,
|
||||
index_repo=index_repo,
|
||||
basic_repo=basic_repo,
|
||||
financial_repo=(
|
||||
financial_repo_factory(session) if financial_repo_factory else None
|
||||
),
|
||||
name_repo=name_repo,
|
||||
)
|
||||
|
||||
def _set_stage(name: str) -> None:
|
||||
"""阶段上报(v3 §23):独立短会话写 job.stage 并 commit(子进程同样走 DB)。"""
|
||||
try:
|
||||
with session_factory() as st_sess:
|
||||
st = job_repo_factory(st_sess).get(job_id)
|
||||
if st is not None and st.status == JobStatus.RUNNING:
|
||||
st.stage = name
|
||||
job_repo_factory(st_sess).update(st)
|
||||
st_sess.commit()
|
||||
except Exception: # noqa: BLE001 —— 阶段上报失败不阻断执行
|
||||
pass
|
||||
|
||||
if is_combo:
|
||||
result = service.run(spec, strategies, config, on_stage=_set_stage)
|
||||
elif is_selection:
|
||||
_set_stage("selection")
|
||||
result = service.select(spec)
|
||||
elif spec.type == "backtest":
|
||||
result = service.run_backtest(spec, on_stage=_set_stage)
|
||||
else:
|
||||
result = service.run_factor_test(spec, on_stage=_set_stage)
|
||||
|
||||
# combo 的结果是 BacktestResult,归档 kind 记为 "backtest" 以便前端按回测渲染;
|
||||
# spec_json 仍存原始 combo(含 strategy_ids),可复现快照在 result.config_snapshot。
|
||||
archive_kind = "backtest" if is_combo else job.kind
|
||||
experiment = archive_experiment(
|
||||
session=session,
|
||||
kind=archive_kind,
|
||||
spec_json=job.spec_json,
|
||||
result=result,
|
||||
job_id=job.id,
|
||||
experiment_repo=experiment_repo,
|
||||
)
|
||||
|
||||
# 完整结果**只在 experiment 存一份**(此前 job.result_json 与
|
||||
# experiment.result_json 各存一份相同内容,完整存档后 2× 浪费)。
|
||||
# GET /api/jobs/{id} 经 job.experiment_id 回读 experiment;
|
||||
# 老记录(result_json 有值、experiment_id 为空)仍走 job 回退解码。
|
||||
job.result_json = None
|
||||
job.experiment_id = experiment.id
|
||||
job.status = JobStatus.SUCCESS
|
||||
job.error = None
|
||||
except Exception as exc: # noqa: BLE001 —— 统一记为 failed 供前端展示
|
||||
job.status = JobStatus.FAILED
|
||||
job.error = f"{type(exc).__name__}: {exc}"
|
||||
finally:
|
||||
job.finished_at = datetime.now()
|
||||
# 回读最后一次阶段上报(_set_stage 经独立会话写库),避免被本会话覆盖
|
||||
try:
|
||||
with session_factory() as last_sess:
|
||||
last = job_repo_factory(last_sess).get(job_id)
|
||||
if last is not None:
|
||||
job.stage = last.stage
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
job_repo.update(job)
|
||||
session.commit()
|
||||
|
||||
|
||||
def execute_job(
|
||||
job_id: str,
|
||||
*,
|
||||
session_factory: Callable,
|
||||
job_repo_factory: Callable,
|
||||
experiment_repo_factory: Callable,
|
||||
stock_repo_factory: Callable,
|
||||
daily_repo_factory: Callable,
|
||||
engine,
|
||||
basic_repo_factory: Callable | None = None,
|
||||
financial_repo_factory: Callable | None = None,
|
||||
index_repo_factory: Callable | None = None,
|
||||
name_repo_factory: Callable | None = None,
|
||||
) -> None:
|
||||
"""入口包装:任何未预期异常都将 Job 标记 failed(防止卡在 queued/running)。"""
|
||||
try:
|
||||
_execute_inner(
|
||||
job_id,
|
||||
session_factory=session_factory,
|
||||
job_repo_factory=job_repo_factory,
|
||||
experiment_repo_factory=experiment_repo_factory,
|
||||
stock_repo_factory=stock_repo_factory,
|
||||
daily_repo_factory=daily_repo_factory,
|
||||
engine=engine,
|
||||
basic_repo_factory=basic_repo_factory,
|
||||
financial_repo_factory=financial_repo_factory,
|
||||
index_repo_factory=index_repo_factory,
|
||||
name_repo_factory=name_repo_factory,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
try:
|
||||
with session_factory() as session:
|
||||
repo = job_repo_factory(session)
|
||||
job = repo.get(job_id)
|
||||
if job is not None:
|
||||
job.status = JobStatus.FAILED
|
||||
job.error = f"内部错误: {type(exc).__name__}: {exc}"
|
||||
job.finished_at = datetime.now()
|
||||
repo.update(job)
|
||||
session.commit()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
|
||||
def default_factories() -> dict:
|
||||
"""后台执行所需的独立 Session / Repository / 引擎装配(跨请求生命周期)。"""
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.index_impl import (
|
||||
SqlAlchemyIndexConstituentRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyExperimentRepository,
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyDailyBasicRepository,
|
||||
SqlAlchemyFinancialRepository,
|
||||
SqlAlchemyStockNameHistoryRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.session import SessionLocal
|
||||
from app.quant.engine import LocalEngine
|
||||
|
||||
return {
|
||||
"session_factory": SessionLocal,
|
||||
"job_repo_factory": lambda s: SqlAlchemyJobRepository(s),
|
||||
"experiment_repo_factory": lambda s: SqlAlchemyExperimentRepository(s),
|
||||
"stock_repo_factory": lambda s: SqlAlchemyStockRepository(s),
|
||||
"daily_repo_factory": lambda s: SqlAlchemyDailyBarRepository(s),
|
||||
# 新增装配(daily_basic / adjust_factor / 财务 / 指数成分):
|
||||
# 缺失时对应能力(dv_ratio 因子、qfq/hfq 复权、fundamental 条件、指数成份池)
|
||||
# 会由 Service 明确报错,绝不静默降级
|
||||
"basic_repo_factory": lambda s: SqlAlchemyDailyBasicRepository(s),
|
||||
"financial_repo_factory": lambda s: SqlAlchemyFinancialRepository(s),
|
||||
"index_repo_factory": lambda s: SqlAlchemyIndexConstituentRepository(s),
|
||||
"name_repo_factory": lambda s: SqlAlchemyStockNameHistoryRepository(s),
|
||||
"engine": LocalEngine(),
|
||||
}
|
||||
|
||||
|
||||
def submit_and_run(spec: ResearchSpec, *, factories: dict | None = None) -> JobRecord:
|
||||
"""创建并执行一个 Job(复用 Job 状态机与 Experiment 归档),返回终态 Job。
|
||||
|
||||
按 job.mode 调度:subprocess 模式在独立进程内执行(内存隔离),否则本进程。
|
||||
"""
|
||||
facts = factories or default_factories()
|
||||
session_factory = facts["session_factory"]
|
||||
job_repo = facts["job_repo_factory"]
|
||||
job = JobRecord(
|
||||
id=new_id("JOB"),
|
||||
kind=spec.type,
|
||||
spec_json=json.dumps(spec.model_dump(mode="json"), ensure_ascii=False),
|
||||
status=JobStatus.QUEUED,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
with session_factory() as session:
|
||||
job_repo(session).create(job)
|
||||
session.commit()
|
||||
if _job_mode() == "subprocess":
|
||||
_run_in_subprocess(job.id, session_factory=session_factory, job_repo_factory=job_repo)
|
||||
else:
|
||||
execute_job(
|
||||
job.id,
|
||||
session_factory=session_factory,
|
||||
job_repo_factory=job_repo,
|
||||
experiment_repo_factory=facts["experiment_repo_factory"],
|
||||
stock_repo_factory=facts["stock_repo_factory"],
|
||||
daily_repo_factory=facts["daily_repo_factory"],
|
||||
engine=facts["engine"],
|
||||
basic_repo_factory=facts.get("basic_repo_factory"),
|
||||
financial_repo_factory=facts.get("financial_repo_factory"),
|
||||
index_repo_factory=facts.get("index_repo_factory"),
|
||||
name_repo_factory=facts.get("name_repo_factory"),
|
||||
)
|
||||
with session_factory() as session:
|
||||
done = job_repo(session).get(job.id)
|
||||
# **仅内存**读透(不落库):完整结果只存 experiment 一份,但 submit_and_run 的
|
||||
# 既有调用方(scripts/run_dividend_case.py、agent 工具)习惯从 job.result_json
|
||||
# 取结果,这里按 experiment_id 回读一次填进返回对象,避免调用方静默拿到空结果。
|
||||
# 数据库中的 job.result_json 仍然保持 NULL(P1:不重复存第二份)。
|
||||
if done is not None and done.experiment_id and done.result_json is None:
|
||||
exp = facts["experiment_repo_factory"](session).get(done.experiment_id)
|
||||
if exp is not None:
|
||||
done.result_json = exp.result_json
|
||||
assert done is not None
|
||||
return done
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 执行模式调度(config job.mode:local | subprocess)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _job_mode() -> str:
|
||||
from app.core.config import get_settings
|
||||
|
||||
return get_settings().job_mode
|
||||
|
||||
|
||||
def _job_memory_limit_gb() -> int:
|
||||
from app.core.config import get_settings
|
||||
|
||||
return get_settings().job_memory_limit_gb
|
||||
|
||||
|
||||
def _job_max_concurrent() -> int:
|
||||
from app.core.config import get_settings
|
||||
|
||||
return get_settings().job_max_concurrent
|
||||
|
||||
|
||||
def _job_worker_cmd(job_id: str) -> list[str]:
|
||||
"""研究子进程命令行(独立进程执行,cwd=backend 使 `-m app.cli.run_job` 可导入)。"""
|
||||
return [sys.executable, "-m", "app.cli.run_job", job_id]
|
||||
|
||||
|
||||
def run_job_background(
|
||||
job_id: str,
|
||||
*,
|
||||
session_factory: Callable | None = None,
|
||||
job_repo_factory: Callable | None = None,
|
||||
) -> None:
|
||||
"""API 后台任务入口:按 job.mode 调度 Job 执行。"""
|
||||
if _job_mode() == "subprocess":
|
||||
_run_in_subprocess(
|
||||
job_id, session_factory=session_factory, job_repo_factory=job_repo_factory
|
||||
)
|
||||
else:
|
||||
execute_job(job_id, **default_factories())
|
||||
|
||||
|
||||
def _acquire_slot(job_id: str) -> bool:
|
||||
with _active_lock:
|
||||
if len(_active_jobs) >= max(_job_max_concurrent(), 1):
|
||||
return False
|
||||
_active_jobs[job_id] = None
|
||||
return True
|
||||
|
||||
|
||||
def _release_slot(job_id: str) -> None:
|
||||
with _active_lock:
|
||||
_active_jobs.pop(job_id, None)
|
||||
|
||||
|
||||
def _mark_failed(
|
||||
job_id: str,
|
||||
error: str,
|
||||
*,
|
||||
session_factory: Callable | None = None,
|
||||
job_repo_factory: Callable | None = None,
|
||||
) -> None:
|
||||
"""把非终态 Job 标记 failed(子进程异常退出 / 并发超限时兜底,防永久 running)。"""
|
||||
if session_factory is None or job_repo_factory is None:
|
||||
facts = default_factories()
|
||||
session_factory = session_factory or facts["session_factory"]
|
||||
job_repo_factory = job_repo_factory or facts["job_repo_factory"]
|
||||
try:
|
||||
with session_factory() as session:
|
||||
repo = job_repo_factory(session)
|
||||
job = repo.get(job_id)
|
||||
if job is not None and job.status in (JobStatus.QUEUED, JobStatus.RUNNING):
|
||||
job.status = JobStatus.FAILED
|
||||
job.error = error
|
||||
job.finished_at = datetime.now()
|
||||
repo.update(job)
|
||||
session.commit()
|
||||
except Exception: # noqa: BLE001 —— 兜底标记失败自身异常不向上抛
|
||||
pass
|
||||
|
||||
|
||||
def _subprocess_log_target():
|
||||
"""研究子进程 stderr 的落盘目标。
|
||||
|
||||
原先父进程用 DEVNULL,会吞掉子进程全部诊断(含 run_job 打印的内存上限设置失败
|
||||
警告),使「如实提示未实现项」落空。改为追加写入 <项目根>/logs/job-subprocess.log。
|
||||
记日志失败(无权限/磁盘满)绝不能导致 Job 失败 —— 此时退回 DEVNULL。
|
||||
"""
|
||||
try:
|
||||
log_dir = BACKEND_ROOT.parent / "logs"
|
||||
log_dir.mkdir(parents=True, exist_ok=True)
|
||||
return (log_dir / "job-subprocess.log").open("a", encoding="utf-8")
|
||||
except OSError:
|
||||
return subprocess.DEVNULL
|
||||
|
||||
|
||||
def _run_in_subprocess(
|
||||
job_id: str,
|
||||
*,
|
||||
session_factory: Callable | None = None,
|
||||
job_repo_factory: Callable | None = None,
|
||||
) -> None:
|
||||
"""在独立 python 进程执行 Job:子进程 RLIMIT 上限(防 OOM 整机),
|
||||
父进程等待;子进程异常退出(如被杀)时把 Job 补记 failed。"""
|
||||
if not _acquire_slot(job_id):
|
||||
_mark_failed(
|
||||
job_id,
|
||||
"系统繁忙:并发研究任务已达上限,请稍后重试",
|
||||
session_factory=session_factory,
|
||||
job_repo_factory=job_repo_factory,
|
||||
)
|
||||
return
|
||||
env = dict(os.environ)
|
||||
env["QLIB_JOB_MEM_LIMIT_GB"] = str(_job_memory_limit_gb())
|
||||
rc = -1
|
||||
already_marked = False
|
||||
errlog = _subprocess_log_target()
|
||||
try:
|
||||
proc = subprocess.Popen(
|
||||
_job_worker_cmd(job_id),
|
||||
cwd=BACKEND_ROOT,
|
||||
env=env,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=errlog,
|
||||
)
|
||||
with _active_lock:
|
||||
_active_jobs[job_id] = proc
|
||||
rc = proc.wait()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
_mark_failed(
|
||||
job_id,
|
||||
f"研究子进程启动失败: {type(exc).__name__}: {exc}",
|
||||
session_factory=session_factory,
|
||||
job_repo_factory=job_repo_factory,
|
||||
)
|
||||
already_marked = True
|
||||
finally:
|
||||
if errlog is not subprocess.DEVNULL:
|
||||
errlog.close()
|
||||
_release_slot(job_id)
|
||||
if rc != 0 and not already_marked:
|
||||
_mark_failed(
|
||||
job_id,
|
||||
f"研究子进程异常退出(code={rc}):任务未完成",
|
||||
session_factory=session_factory,
|
||||
job_repo_factory=job_repo_factory,
|
||||
)
|
||||
|
||||
|
||||
def terminate_active(job_id: str) -> bool:
|
||||
"""终止该 Job 的活动子进程(如有)。返回是否找到并终止。"""
|
||||
with _active_lock:
|
||||
proc = _active_jobs.get(job_id)
|
||||
if proc is None:
|
||||
return False
|
||||
import contextlib
|
||||
|
||||
with contextlib.suppress(Exception):
|
||||
proc.terminate()
|
||||
return True
|
||||
|
||||
|
||||
def cancel_job(
|
||||
job_id: str,
|
||||
*,
|
||||
session_factory: Callable | None = None,
|
||||
job_repo_factory: Callable | None = None,
|
||||
) -> bool:
|
||||
"""取消 queued/running Job(置 CANCELLED 并终止子进程);不可取消返回 False。"""
|
||||
facts = default_factories()
|
||||
sf = session_factory or facts["session_factory"]
|
||||
jr = job_repo_factory or facts["job_repo_factory"]
|
||||
with sf() as session:
|
||||
repo = jr(session)
|
||||
job = repo.get(job_id)
|
||||
if job is None or job.status not in (JobStatus.QUEUED, JobStatus.RUNNING):
|
||||
return False
|
||||
job.status = JobStatus.CANCELLED
|
||||
job.stage = None
|
||||
repo.update(job)
|
||||
session.commit()
|
||||
terminate_active(job_id)
|
||||
return True
|
||||
|
||||
|
||||
def mark_stale_jobs_failed(
|
||||
*,
|
||||
session_factory: Callable | None = None,
|
||||
job_repo_factory: Callable | None = None,
|
||||
reason: str | None = None,
|
||||
) -> int:
|
||||
"""服务启动调用:把上次进程异常退出遗留的 queued/running Job 标记 failed。
|
||||
|
||||
返回处理数量。防「进程被杀后 Job 永久 running」(AGENT.md §20 状态机闭环)。
|
||||
"""
|
||||
facts = default_factories()
|
||||
sf = session_factory or facts["session_factory"]
|
||||
jr = job_repo_factory or facts["job_repo_factory"]
|
||||
counted = 0
|
||||
with sf() as session:
|
||||
repo = jr(session)
|
||||
for status in (JobStatus.QUEUED, JobStatus.RUNNING):
|
||||
for job in repo.list_by_status(status):
|
||||
job.status = JobStatus.FAILED
|
||||
job.error = reason or "服务重启:上次未完成任务被中断"
|
||||
job.finished_at = datetime.now()
|
||||
repo.update(job)
|
||||
counted += 1
|
||||
session.commit()
|
||||
return counted
|
||||
@@ -0,0 +1,101 @@
|
||||
"""Bar Replay 服务(M9-6):线性逐交易日重放选股+信号(as_of 语义)。
|
||||
|
||||
- 范围约束:universe.symbols 必填(≤ 40 只)、重放交易日 ≤ 90 —— 避免全市场长任务
|
||||
- 每日计算只使用 <= as_of 数据(与 select/signal/回测同一引擎与口径)
|
||||
- ReplayDay.events 按 rank 升序;top 取前 N(意图排名,与回测 selection_history 对齐)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, timedelta
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from app.domain.entities.replay import ReplayDay, ReplayResult, ReplayTop
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
from app.domain.entities.signal import SignalRules
|
||||
from app.domain.repositories.market import DailyBarRepository, StockRepository
|
||||
from app.quant.selection import factor_columns
|
||||
from app.quant.service import load_daily_df
|
||||
from app.quant.signal import generate_signals
|
||||
from app.quant.universe import filter_stocks, names_as_of, resolve_members
|
||||
|
||||
MAX_SYMBOLS = 40
|
||||
MAX_DAYS = 90
|
||||
|
||||
|
||||
class ReplayService:
|
||||
def __init__(
|
||||
self,
|
||||
stock_repo: StockRepository,
|
||||
daily_repo: DailyBarRepository,
|
||||
index_repo=None,
|
||||
name_repo=None,
|
||||
) -> None:
|
||||
self._stock_repo = stock_repo
|
||||
self._daily_repo = daily_repo
|
||||
self._index_repo = index_repo
|
||||
# 名称变更历史仓储:exclude_st 的时点口径(与选股/回测一致,v2 §25)
|
||||
self._name_repo = name_repo
|
||||
|
||||
def replay(
|
||||
self,
|
||||
query: SelectionQuery,
|
||||
rules: SignalRules,
|
||||
start: date,
|
||||
end: date,
|
||||
top_n: int = 5,
|
||||
) -> ReplayResult:
|
||||
symbols = list(query.universe.symbols or [])
|
||||
if not symbols:
|
||||
raise ValueError("Bar Replay 需要 universe.symbols 白名单(≤40 只),避免全市场长任务")
|
||||
if len(symbols) > MAX_SYMBOLS:
|
||||
raise ValueError(f"Bar Replay 白名单最多 {MAX_SYMBOLS} 只,当前 {len(symbols)}")
|
||||
all_stocks = self._stock_repo.list()
|
||||
name_at, _applied = names_as_of(all_stocks, start, self._name_repo)
|
||||
stocks = filter_stocks(
|
||||
all_stocks, query.universe, as_of=start,
|
||||
members=resolve_members(self._index_repo, query.universe, start),
|
||||
name_at=name_at,
|
||||
)
|
||||
if not stocks:
|
||||
return ReplayResult(start=start, end=end, top_n=top_n)
|
||||
|
||||
columns = sorted(factor_columns(query))
|
||||
daily = load_daily_df(
|
||||
self._daily_repo,
|
||||
symbols,
|
||||
start - timedelta(days=query.warmup_days),
|
||||
end,
|
||||
columns,
|
||||
adjust=query.price_adjustment,
|
||||
)
|
||||
if daily.empty:
|
||||
return ReplayResult(start=start, end=end, top_n=top_n)
|
||||
trading_days = sorted(
|
||||
pd.to_datetime(daily["trade_date"].unique())
|
||||
)
|
||||
days = [d for d in trading_days if start <= d.date() <= end]
|
||||
if len(days) > MAX_DAYS:
|
||||
raise ValueError(f"重放区间交易日 {len(days)} > 上限 {MAX_DAYS},请缩短区间")
|
||||
|
||||
out_days: list[ReplayDay] = []
|
||||
for d in days:
|
||||
res = generate_signals(daily, query, rules, as_of=d.date())
|
||||
top = [
|
||||
ReplayTop(symbol=e.symbol, score=e.score or 0.0)
|
||||
for e in res.events[:top_n]
|
||||
]
|
||||
out_days.append(
|
||||
ReplayDay(
|
||||
as_of=d.date(),
|
||||
top=top,
|
||||
events=res.events,
|
||||
counts={
|
||||
"buy": res.statistics.buy,
|
||||
"watch": res.statistics.watch,
|
||||
"sell": res.statistics.sell,
|
||||
},
|
||||
)
|
||||
)
|
||||
return ReplayResult(start=start, end=end, days=out_days, top_n=top_n)
|
||||
@@ -0,0 +1,186 @@
|
||||
"""选股用例入口(ARCHITECTURE_v2 §14 Selection Engine · 业务层)。
|
||||
|
||||
- 输入:SelectionQuery(universe + method + factors/conditions + top_n/pct + as_of)
|
||||
- 装配:股票池(universe 过滤)→ 行情长表(含预热窗口)→ Selection Engine
|
||||
(method=score 因子评分 / method=condition 结构化条件)
|
||||
- 输出:SelectionResult(可解释:factor_values / filter_status / selection_reason)
|
||||
- 未来函数红线:行情只取 <= as_of;财务条件只取 announce_date <= as_of 的已公告值(v2 §9)
|
||||
|
||||
MVP 为同步执行(单日全市场因子/条件计算量轻);如需异步可复用 Job 链路。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, timedelta
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from app.domain.entities.market import FinancialIndicator
|
||||
from app.domain.entities.selection import SelectionQuery, SelectionResult
|
||||
from app.domain.repositories.market import (
|
||||
DailyBarRepository,
|
||||
FinancialRepository,
|
||||
StockRepository,
|
||||
)
|
||||
from app.quant.selection import (
|
||||
condition_needed_columns,
|
||||
factor_columns,
|
||||
run_condition_selection,
|
||||
run_score_selection,
|
||||
)
|
||||
from app.quant.service import (
|
||||
load_basic_df,
|
||||
load_daily_df,
|
||||
merge_basic_into_daily,
|
||||
split_factor_columns,
|
||||
)
|
||||
from app.quant.universe import filter_stocks, names_as_of, resolve_members
|
||||
|
||||
_FUNDAMENTAL_PREFIX = "fundamental."
|
||||
|
||||
|
||||
def fill_candidate_names(result: SelectionResult, stocks: list) -> SelectionResult:
|
||||
"""把股票池的 `symbol → name` 回填进候选股(展示增强,未命中保持 None)。
|
||||
|
||||
为什么在业务层做:名称是展示数据而非选股语义,引擎(quant/selection.py)
|
||||
只做纯数值计算,不应感知名称;而本用例的 `stocks` 已是 universe 过滤后的
|
||||
股票实体列表(天然带 name),在这里一次性建立映射即可,不必在各出口各自查库。
|
||||
查不到名称的候选保持 None —— 前端按「名称未知」渲染,不伪造也不报错。
|
||||
"""
|
||||
if not result.candidates or not stocks:
|
||||
return result
|
||||
name_map = {s.symbol: s.name for s in stocks if getattr(s, "name", None)}
|
||||
if not name_map:
|
||||
return result
|
||||
for cand in result.candidates:
|
||||
if cand.name is None:
|
||||
cand.name = name_map.get(cand.symbol)
|
||||
return result
|
||||
|
||||
|
||||
class SelectionService:
|
||||
"""选股用例入口:select(query) → SelectionResult(当前或历史 as_of)。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stock_repo: StockRepository,
|
||||
daily_repo: DailyBarRepository,
|
||||
financial_repo: FinancialRepository | None = None,
|
||||
index_repo=None,
|
||||
basic_repo=None,
|
||||
name_repo=None,
|
||||
) -> None:
|
||||
self._stock_repo = stock_repo
|
||||
self._daily_repo = daily_repo
|
||||
self._financial_repo = financial_repo
|
||||
self._index_repo = index_repo
|
||||
# 每日指标仓储(daily_basic):score 因子/条件引用 dv_ratio 等列时使用
|
||||
self._basic_repo = basic_repo
|
||||
# 名称变更历史仓储:exclude_st 的时点口径(与回测口径一致,v2 §25)
|
||||
self._name_repo = name_repo
|
||||
|
||||
def select(self, query: SelectionQuery) -> SelectionResult:
|
||||
as_of = query.as_of or date.today()
|
||||
all_stocks = self._stock_repo.list()
|
||||
name_at, _applied = names_as_of(all_stocks, as_of, self._name_repo)
|
||||
stocks = filter_stocks(
|
||||
all_stocks, query.universe, as_of=as_of,
|
||||
members=resolve_members(self._index_repo, query.universe, as_of),
|
||||
name_at=name_at,
|
||||
)
|
||||
if not stocks:
|
||||
return self._run(query, pd.DataFrame(), stocks, as_of, financial={})
|
||||
symbols = [s.symbol for s in stocks]
|
||||
if query.method == "score":
|
||||
columns = sorted(factor_columns(query))
|
||||
else:
|
||||
columns = sorted(condition_needed_columns(query))
|
||||
bar_cols, basic_cols = split_factor_columns(columns)
|
||||
data_start = as_of - timedelta(days=query.warmup_days)
|
||||
daily = load_daily_df(
|
||||
self._daily_repo,
|
||||
symbols,
|
||||
data_start,
|
||||
as_of,
|
||||
sorted(bar_cols),
|
||||
adjust="none",
|
||||
price_adjust=query.price_adjustment,
|
||||
)
|
||||
if basic_cols:
|
||||
daily = self._attach_basic(daily, symbols, data_start, as_of, sorted(basic_cols))
|
||||
financial: dict[str, FinancialIndicator] = {}
|
||||
if query.method == "condition" and self._uses_fundamental(query):
|
||||
financial = self._load_financial(symbols, as_of)
|
||||
return self._run(query, daily, stocks, as_of, financial)
|
||||
|
||||
# ---- 内部 ----
|
||||
|
||||
def _attach_basic(
|
||||
self,
|
||||
daily: pd.DataFrame,
|
||||
symbols: list[str],
|
||||
start: date,
|
||||
end: date,
|
||||
columns: list[str],
|
||||
) -> pd.DataFrame:
|
||||
"""并入 daily_basic 列(与 ResearchService 同一装配逻辑,保证 v2 §25 一致性)。"""
|
||||
if self._basic_repo is None:
|
||||
raise ValueError(
|
||||
f"选股条件/因子需要每日指标列 {columns}(daily_basic),"
|
||||
"但未注入 DailyBasicRepository。请检查 API 的依赖装配。"
|
||||
)
|
||||
basic = load_basic_df(self._basic_repo, symbols, start, end, columns)
|
||||
if basic.empty:
|
||||
raise ValueError(
|
||||
f"daily_basic 表在 {start}~{end} 无数据,无法计算需要 {columns} 的因子/条件。"
|
||||
"请先运行:python -m app.cli.sync daily_basic --start 20200101"
|
||||
)
|
||||
return merge_basic_into_daily(daily, basic)
|
||||
|
||||
def _run(
|
||||
self,
|
||||
query: SelectionQuery,
|
||||
daily: pd.DataFrame,
|
||||
stocks: list,
|
||||
as_of: date,
|
||||
financial: dict[str, FinancialIndicator],
|
||||
) -> SelectionResult:
|
||||
if query.method == "score":
|
||||
result = run_score_selection(daily, query, as_of)
|
||||
else:
|
||||
result = run_condition_selection(daily, stocks, query, as_of, financial)
|
||||
# 名称只在业务层回填:引擎(quant/selection.py)保持纯符号计算,
|
||||
# 而 `stocks` 是本用例已经装配好的股票池,天然带 name,无需再查库/join。
|
||||
return fill_candidate_names(result, stocks)
|
||||
|
||||
@staticmethod
|
||||
def _uses_fundamental(query: SelectionQuery) -> bool:
|
||||
for c in query.conditions:
|
||||
if c.field.startswith(_FUNDAMENTAL_PREFIX) or (
|
||||
c.ref is not None and c.ref.startswith(_FUNDAMENTAL_PREFIX)
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
def _load_financial(
|
||||
self, symbols: list[str], as_of: date
|
||||
) -> dict[str, FinancialIndicator]:
|
||||
"""按 announce_date <= as_of 批量取财务,每 symbol 保留最新一版。"""
|
||||
if self._financial_repo is None:
|
||||
raise ValueError("condition 引用了 fundamental.* 字段,但未注入 FinancialRepository")
|
||||
getter = getattr(self._financial_repo, "list_announced_many", None)
|
||||
if getter is not None:
|
||||
rows = list(getter(symbols, as_of))
|
||||
else: # 回退逐只
|
||||
rows = []
|
||||
for sym in symbols:
|
||||
rows.extend(self._financial_repo.list_announced(sym, as_of))
|
||||
by_symbol: dict[str, FinancialIndicator] = {}
|
||||
for row in rows:
|
||||
cur = by_symbol.get(row.symbol)
|
||||
if cur is None or (row.announce_date, row.report_date) > (
|
||||
cur.announce_date,
|
||||
cur.report_date,
|
||||
):
|
||||
by_symbol[row.symbol] = row
|
||||
return by_symbol
|
||||
@@ -0,0 +1,55 @@
|
||||
"""信号用例(M8.1):基于选股评分排序 + 技术条件生成交易信号。
|
||||
|
||||
信号与回测买入逻辑同源(同一评分引擎、同一口径),保证「为什么 BUY/SELL」可解释。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, timedelta
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
from app.domain.entities.signal import SignalResult, SignalRules
|
||||
from app.domain.repositories.market import DailyBarRepository, StockRepository
|
||||
from app.quant.selection import factor_columns
|
||||
from app.quant.service import load_daily_df
|
||||
from app.quant.signal import generate_signals
|
||||
from app.quant.universe import filter_stocks, names_as_of, resolve_members
|
||||
|
||||
|
||||
class SignalService:
|
||||
def __init__(
|
||||
self,
|
||||
stock_repo: StockRepository,
|
||||
daily_repo: DailyBarRepository,
|
||||
index_repo=None,
|
||||
name_repo=None,
|
||||
) -> None:
|
||||
self._stock_repo = stock_repo
|
||||
self._daily_repo = daily_repo
|
||||
self._index_repo = index_repo
|
||||
# 名称变更历史仓储:exclude_st 的时点口径(与选股/回测一致,v2 §25)
|
||||
self._name_repo = name_repo
|
||||
|
||||
def signal(self, query: SelectionQuery, rules: SignalRules) -> SignalResult:
|
||||
as_of = query.as_of or date.today()
|
||||
all_stocks = self._stock_repo.list()
|
||||
name_at, _applied = names_as_of(all_stocks, as_of, self._name_repo)
|
||||
stocks = filter_stocks(
|
||||
all_stocks, query.universe, as_of=as_of,
|
||||
members=resolve_members(self._index_repo, query.universe, as_of),
|
||||
name_at=name_at,
|
||||
)
|
||||
if not stocks:
|
||||
return generate_signals(pd.DataFrame(), query, rules, as_of)
|
||||
columns = sorted(factor_columns(query))
|
||||
daily = load_daily_df(
|
||||
self._daily_repo,
|
||||
[s.symbol for s in stocks],
|
||||
as_of - timedelta(days=query.warmup_days),
|
||||
as_of,
|
||||
columns,
|
||||
adjust=query.price_adjustment,
|
||||
)
|
||||
return generate_signals(daily, query, rules, as_of)
|
||||
@@ -0,0 +1 @@
|
||||
"""命令行工具(数据同步等)。用法:uv run python -m app.cli.sync ..."""
|
||||
@@ -0,0 +1,125 @@
|
||||
"""实验归档保留策略 CLI(归档体积治理)。
|
||||
|
||||
用法(cd backend):
|
||||
# 只看会删什么(**默认就是 dry-run**,不写库)
|
||||
uv run python -m app.cli.prune_experiments --keep 30
|
||||
uv run python -m app.cli.prune_experiments --keep 20 --kind backtest
|
||||
uv run python -m app.cli.prune_experiments --older-than 180 # 180 天前的归档
|
||||
# 真正删除(必须显式加 --apply)
|
||||
uv run python -m app.cli.prune_experiments --keep 30 --apply
|
||||
|
||||
为什么默认 dry-run:归档是**研究结论的唯一凭据**(结果 + spec + 代码/数据版本),
|
||||
误删不可恢复。因此:
|
||||
- 默认只打印候选清单(含 id / 类型 / 时间 / 体积 / 摘要)与合计释放空间,不写库;
|
||||
- 只有显式 `--apply` 才执行删除,且删除前**再打印一次**清单;
|
||||
- 按 `created_at` **由新到旧保留**:先按过滤条件选出候选集合,再保留最新的 N 条,
|
||||
其余删除(`--keep` 与 `--older-than` 可同时给,取交集);
|
||||
- 只删 `experiment` 表(归档本身),**不动 `job_record`**(执行历史仍可追溯);
|
||||
但要注意:完整结果只存归档一份(`job.result_json` 对新记录为 NULL),
|
||||
因此**删除归档 = 该次回测结果不可再查看**。要留底先导出(归档页「导出完整 JSON」);
|
||||
- `--kind` 可只清理某一类(如只清 `factor_test`,保留回测结论)。
|
||||
|
||||
本模块是组装层(composition root):装配 Session 与 Repository,删除逻辑走仓储。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyExperimentRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.session import SessionLocal
|
||||
|
||||
|
||||
def select_prune_candidates(
|
||||
rows: list,
|
||||
*,
|
||||
keep: int | None,
|
||||
older_than_days: int | None,
|
||||
now: datetime | None = None,
|
||||
) -> list:
|
||||
"""从归档列表里挑出**该删除**的那些(纯函数,便于单测)。
|
||||
|
||||
`rows`:任意具备 `id` / `created_at` 的归档摘要(不依赖 ORM 与数据库)。
|
||||
规则(两个条件同时给出时取交集):
|
||||
- `keep=N`:按 `created_at` 由新到旧排序后,**保留最新 N 条**,其余为候选;
|
||||
- `older_than_days=D`:创建时间早于 `now - D 天` 的才是候选;
|
||||
- `created_at` 为空的归档视为**最旧**(排序末位)——缺失时间时宁可被列为候选,
|
||||
也不静默把它当作「最新」而永久留在库里。
|
||||
"""
|
||||
ordered = sorted(rows, key=lambda e: (e.created_at or datetime.min), reverse=True)
|
||||
candidates = list(ordered)
|
||||
if older_than_days is not None:
|
||||
cutoff = (now or datetime.now()) - timedelta(days=older_than_days)
|
||||
candidates = [e for e in candidates if (e.created_at or datetime.min) < cutoff]
|
||||
if keep is not None:
|
||||
keep_ids = {e.id for e in ordered[: max(keep, 0)]}
|
||||
candidates = [e for e in candidates if e.id not in keep_ids]
|
||||
# 输出仍按新 → 旧,便于人核对
|
||||
candidates.sort(key=lambda e: (e.created_at or datetime.min), reverse=True)
|
||||
return candidates
|
||||
|
||||
|
||||
def _parse_args(argv: list[str]) -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(
|
||||
prog="python -m app.cli.prune_experiments",
|
||||
description="实验归档保留策略(默认 dry-run,需 --apply 才真正删除)",
|
||||
)
|
||||
p.add_argument("--keep", type=int, default=None, help="按时间由新到旧保留的最新条数")
|
||||
p.add_argument("--kind", default=None, help="只处理某类归档(backtest/factor_test/selection)")
|
||||
p.add_argument(
|
||||
"--older-than",
|
||||
type=int,
|
||||
default=None,
|
||||
metavar="DAYS",
|
||||
help="只处理创建时间早于 N 天的归档",
|
||||
)
|
||||
p.add_argument("--apply", action="store_true", help="真正执行删除(否则仅预览)")
|
||||
args = p.parse_args(argv)
|
||||
if args.keep is None and args.older_than is None:
|
||||
p.error("至少给出 --keep 或 --older-than 之一,避免误删全部归档")
|
||||
if args.keep is not None and args.keep < 0:
|
||||
p.error("--keep 不能为负数")
|
||||
return args
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
args = _parse_args(sys.argv[1:] if argv is None else argv)
|
||||
with SessionLocal() as session:
|
||||
repo = SqlAlchemyExperimentRepository(session)
|
||||
# 列表查询刻意不取 result_json,体积用库侧算出的 result_bytes(单位:字符)
|
||||
rows = repo.list_filtered(kind=args.kind, limit=1_000_000, offset=0)
|
||||
candidates = select_prune_candidates(
|
||||
rows, keep=args.keep, older_than_days=args.older_than
|
||||
)
|
||||
|
||||
total_bytes = sum(e.result_bytes or 0 for e in candidates)
|
||||
print(f"归档总数 {len(rows)} 条,命中删除条件 {len(candidates)} 条,"
|
||||
f"预计释放约 {total_bytes / 1024 / 1024:.1f} MB")
|
||||
for e in candidates:
|
||||
created = e.created_at.strftime("%Y-%m-%d %H:%M") if e.created_at else "未知时间"
|
||||
print(
|
||||
f" - {e.id} {e.kind:<12} {created} "
|
||||
f"{((e.result_bytes or 0) / 1024):>8.0f} KB {(e.summary_text or '')[:60]}"
|
||||
)
|
||||
if not candidates:
|
||||
print("没有需要删除的归档。")
|
||||
return 0
|
||||
if not args.apply:
|
||||
print("\n[dry-run] 未删除任何记录;确认无误后加 --apply 执行。")
|
||||
return 0
|
||||
|
||||
deleted = 0
|
||||
for e in candidates:
|
||||
if repo.delete(e.id):
|
||||
deleted += 1
|
||||
session.commit()
|
||||
print(f"\n已删除 {deleted} 条归档(job_record 未改动)。")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,172 @@
|
||||
"""从 Job 结果副本重建 Experiment 归档(删除后的可选恢复路径)。
|
||||
|
||||
**为什么会有这个工具**:完整存档上线前,历史归档在 `job.result_json` 里另存了一份完整结果
|
||||
(双写遗留)。删除归档只删 `experiment` 行、不动 `job` 行,因此这些历史记录
|
||||
**可以从 Job 副本原样重建** —— 删除不等于数据永久消失。
|
||||
新归档(`job.result_json IS NULL`)没有副本,删除即不可恢复,本工具会明确拒绝而不是假装能救。
|
||||
|
||||
注意:本工具是**恢复手段**,不是"删除的撤销键"。删除本身是正常功能:
|
||||
已有历史归档可重建、新归档不可;是否恢复由人决定,工具默认 dry-run、不做任何自动动作。
|
||||
|
||||
用法(**默认 dry-run**,只打印将要写入的内容;`--apply` 才写库):
|
||||
|
||||
cd backend
|
||||
PYTHONPATH=. .venv/bin/python -m app.cli.restore_experiment_from_job --job-id JOB-XXXX
|
||||
PYTHONPATH=. .venv/bin/python -m app.cli.restore_experiment_from_job --job-id JOB-XXXX \
|
||||
--code-version 92627f5 --apply
|
||||
|
||||
诚实性要求(AGENT.md §7/§24):
|
||||
- **归档 id 用回原 id**(`job.experiment_id`),移动端/书签里的旧链接继续有效;
|
||||
- `spec_json` / `result_json` 逐字复制 Job 副本,不重新计算、不"顺手修正";
|
||||
- `summary_text` 由副本 JSON 反序列化后按归档同一函数重新生成(口径一致);
|
||||
- `code_version` **必须显式提供**,不做猜测:显式传 `--code-version ""` 表示"未知则留空";
|
||||
- `data_version` 留空 —— 历史归档当年没有数据指纹,补一个今天的指纹是伪造复现依据;
|
||||
- 重建前会检查该 job 的归档是否已存在,已存在则拒绝(除非 `--force`)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
from datetime import datetime
|
||||
|
||||
from app.application.services.experiment_archive import _summary_text
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
ExperimentRecord,
|
||||
FactorTestReport,
|
||||
)
|
||||
from app.domain.entities.selection import SelectionResult
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyExperimentRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.session import SessionLocal
|
||||
|
||||
|
||||
def _parse_result(kind: str, raw: str):
|
||||
"""按 kind 把 Job 里的结果 JSON 反序列化成领域对象(用于生成同一个摘要口径)。"""
|
||||
payload = json.loads(raw)
|
||||
if kind == "backtest":
|
||||
return BacktestResult.model_validate(payload)
|
||||
if kind == "factor_test":
|
||||
return FactorTestReport.model_validate(payload)
|
||||
if kind == "selection":
|
||||
return SelectionResult.model_validate(payload)
|
||||
return None
|
||||
|
||||
|
||||
def _load_job(session, job_id: str):
|
||||
"""按 ORM 读 Job(**不要用裸 SQL**)。
|
||||
|
||||
裸 `text()` 查询在 SQLite 下把 DateTime 列原样返回成字符串,写回 ORM 的 DateTime
|
||||
字段会抛 `TypeError: SQLite DateTime type only accepts Python datetime...`
|
||||
(MySQL+pymysql 返回 datetime 所以当时看不出问题)。ORM 读法跨方言类型一致,
|
||||
这个坑由 `tests/test_restore_experiment.py::test_restores_archive_from_job_copy` 钉住。
|
||||
"""
|
||||
from app.infrastructure.persistence.sqlalchemy.models.jobs import JobModel
|
||||
|
||||
return session.get(JobModel, job_id)
|
||||
|
||||
|
||||
def _experiment_exists(session, exp_id: str) -> bool:
|
||||
from app.infrastructure.persistence.sqlalchemy.models import ExperimentModel
|
||||
|
||||
return session.get(ExperimentModel, exp_id) is not None
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
ap = argparse.ArgumentParser(description="从 Job 结果副本重建 Experiment 归档(删除后的可选恢复路径)")
|
||||
ap.add_argument("--job-id", required=True, help="来源 Job id(如 JOB-D3C120DC)")
|
||||
ap.add_argument(
|
||||
"--code-version",
|
||||
default=None,
|
||||
help="重建记录的 code_version(必填;传空串表示未知留空)。不猜版本。",
|
||||
)
|
||||
ap.add_argument("--force", action="store_true", help="归档已存在时也覆盖(默认拒绝)")
|
||||
ap.add_argument("--apply", action="store_true", help="真正写库(默认 dry-run)")
|
||||
args = ap.parse_args(argv)
|
||||
|
||||
if args.code_version is None:
|
||||
print(
|
||||
"❌ 必须显式给出 --code-version(不猜版本);确定未知请传 --code-version \"\"。",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 2
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
row = _load_job(session, args.job_id)
|
||||
if row is None:
|
||||
print(f"❌ Job {args.job_id} 不存在", file=sys.stderr)
|
||||
return 1
|
||||
job_id = row.id
|
||||
kind = row.kind
|
||||
status = row.status
|
||||
spec_json = row.spec_json
|
||||
result_json = row.result_json
|
||||
exp_id = row.experiment_id
|
||||
finished_at = row.finished_at
|
||||
created_at = row.created_at
|
||||
if not exp_id:
|
||||
print(f"❌ Job {job_id} 没有 experiment_id,无法确定重建为哪个归档", file=sys.stderr)
|
||||
return 1
|
||||
if result_json is None:
|
||||
print(
|
||||
f"❌ Job {job_id} 的 result_json 为空(完整存档上线后结果只存归档一份),"
|
||||
"本工具无法重建 —— 这类归档删除后不可恢复。",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 1
|
||||
|
||||
exists = _experiment_exists(session, exp_id)
|
||||
result = _parse_result(kind, result_json)
|
||||
summary = _summary_text(kind, result) if result is not None else None
|
||||
restored_at = finished_at or created_at or datetime.now()
|
||||
|
||||
print(f"来源 Job : {job_id}(status={status},finished_at={finished_at})")
|
||||
print(f"归档 id : {exp_id}(已存在: {exists})")
|
||||
print(f"kind : {kind}")
|
||||
print(f"spec 长度 : {len(spec_json)} 字符")
|
||||
print(f"result 长度 : {len(result_json)} 字符")
|
||||
print(f"summary_text : {summary!r}")
|
||||
print(f"code_version : {args.code_version!r}")
|
||||
print("data_version : None(历史归档无指纹,不伪造)")
|
||||
print(f"created_at : {restored_at}")
|
||||
|
||||
if exists and not args.force:
|
||||
print("❌ 该归档已存在;如确要覆盖请加 --force", file=sys.stderr)
|
||||
return 1
|
||||
if not args.apply:
|
||||
print("\n(dry-run) 未写库。确认无误后加 --apply。")
|
||||
return 0
|
||||
|
||||
# 走仓储而不是直接操作 ORM 模型(AGENT §10:持久化只经 Repository)。
|
||||
# 这样 experiment 表 ↔ 领域实体的字段映射只有仓储一份:将来加列(尤其 NOT NULL)
|
||||
# 不会在这里静默漏写。覆盖语义由 `upsert`(merge)承担,归档 id 保持不变。
|
||||
record = ExperimentRecord(
|
||||
id=exp_id,
|
||||
kind=kind,
|
||||
spec_json=spec_json,
|
||||
result_json=result_json,
|
||||
summary_text=summary,
|
||||
code_version=args.code_version or None,
|
||||
data_version=None,
|
||||
job_id=job_id,
|
||||
created_at=restored_at,
|
||||
)
|
||||
repo = SqlAlchemyExperimentRepository(session)
|
||||
repo.upsert(record)
|
||||
session.commit()
|
||||
got = repo.get(exp_id)
|
||||
print(
|
||||
f"\n✅ 已重建:{got.id} / {got.kind} / 体积 {len(got.result_json)} 字符 / "
|
||||
f"code_version={got.code_version}"
|
||||
)
|
||||
return 0
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,64 @@
|
||||
"""独立进程执行研究 Job —— 内存隔离的执行体(内存优化专项)。
|
||||
|
||||
由 job_executor._run_in_subprocess 以 `python -m app.cli.run_job <JOB_ID>` 拉起
|
||||
(cwd=backend/),入口先施加 RLIMIT_AS 上限再导入重依赖,随后执行与 API 内
|
||||
execute_job 完全相同的状态机 / Experiment 归档逻辑。
|
||||
|
||||
环境变量:
|
||||
QLIB_JOB_MEM_LIMIT_GB —— 子进程虚拟地址空间上限(GB);父进程 spawn 时写入,
|
||||
手动直接运行本命令而未设置时回退到 Settings(config job.max_memory_gb)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
|
||||
def _apply_memory_limit() -> None:
|
||||
"""尝试施加 RLIMIT_AS:达到上限时分配抛 MemoryError → Job 归档 failed,
|
||||
而不是让内核 OOM 杀掉整机(8G Pi5 上保护同机其它服务)。
|
||||
|
||||
不按平台硬门禁:RLIMIT_AS 在 Linux / macOS 上都能设置,但**限额必须高于进程
|
||||
当前的虚拟地址空间基线**,否则内核返回 EINVAL(CPython 表现为
|
||||
`ValueError: current limit exceeds maximum limit`)。macOS 上 Python 进程的
|
||||
地址空间基线可达数十 GB(共享缓存 + malloc 区预留),故 `job.max_memory_gb: 6`
|
||||
在这类机器上会失败 —— 这是**配置阈值问题,不是平台不支持**,不应掩盖。
|
||||
|
||||
设置失败不阻断 Job,但如实打印原因与后果(AGENT.md §24「未实现项须如实标注」)。
|
||||
"""
|
||||
try:
|
||||
gb = int(os.environ.get("QLIB_JOB_MEM_LIMIT_GB") or "")
|
||||
except (TypeError, ValueError):
|
||||
from app.core.config import get_settings
|
||||
|
||||
gb = get_settings().job_memory_limit_gb
|
||||
limit = gb * 1024**3
|
||||
import resource
|
||||
|
||||
try:
|
||||
resource.setrlimit(resource.RLIMIT_AS, (limit, limit))
|
||||
except (ValueError, OSError) as exc:
|
||||
print(
|
||||
f"[run_job] 内存上限 {gb}GB 设置失败({type(exc).__name__}: {exc});"
|
||||
"本 Job 无内存隔离保护。若限额低于进程虚拟地址空间基线(macOS 常见),"
|
||||
"可调大 config.yaml job.max_memory_gb 后重试。",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
args = list(sys.argv[1:] if argv is None else argv)
|
||||
if len(args) < 1:
|
||||
print("用法: python -m app.cli.run_job <JOB_ID>", file=sys.stderr)
|
||||
return 2
|
||||
_apply_memory_limit()
|
||||
# 先设内存上限再导入重依赖(pandas / sqlalchemy 等)
|
||||
from app.application.services.job_executor import default_factories, execute_job
|
||||
|
||||
execute_job(args[0], **default_factories())
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,642 @@
|
||||
"""Phase 1 数据同步 CLI(Tushare 首选 → SQLite,新浪校验兜底)。
|
||||
|
||||
用法(cd backend):
|
||||
uv run python -m app.cli.sync basic
|
||||
uv run python -m app.cli.sync calendar --start 20240101 --end 20241231
|
||||
uv run python -m app.cli.sync daily --symbols 600519.SH,000001.SZ --start 20240101
|
||||
uv run python -m app.cli.sync daily --all --start 20240101 # 全市场
|
||||
uv run python -m app.cli.sync financial --all # 财务指标(增量)
|
||||
uv run python -m app.cli.sync financial --all --full # 财务指标(强制全量重拉)
|
||||
uv run python -m app.cli.sync verify --symbol 600519.SH # 新浪交叉验证
|
||||
uv run python -m app.cli.sync daily_basic --start 20200101 # 每日指标(股息率等)
|
||||
|
||||
增量与兜底:
|
||||
- daily --resume:从本地最新交易日续传(已有);Tushare 失败时走新浪校验兜底,
|
||||
只有「两源重叠历史一致」才用新浪补本地缺失交易日(source=sina/前复权)。
|
||||
- financial:默认增量——本地已含最新应披露报告期则跳过;Tushare 失败时新浪
|
||||
数据须通过「两边一致」校验(重叠报告期 eps/销售毛利率逐期一致)才允许补入
|
||||
本地缺失键(source=sina)。失败股票留待下轮重跑补齐,不会静默导入未核验数据。
|
||||
- 每次拉取写入 sync_log 审计(来源 / 成功与否 / 行数 / 区间),禁止静默切源。
|
||||
|
||||
本模块是组装层(composition root):在此装配 Provider / Repository / Session,
|
||||
业务逻辑在 application.services.data_sync,业务层仍只依赖抽象。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
import time
|
||||
from datetime import date, datetime, timedelta
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.application.services.data_sync import (
|
||||
DailyBasicSyncer,
|
||||
DailySymbolResult,
|
||||
FinancialSymbolResult,
|
||||
NameHistorySyncer,
|
||||
VerifiedDailySyncer,
|
||||
VerifiedFinancialSyncer,
|
||||
)
|
||||
from app.core.config import get_settings
|
||||
from app.infrastructure.data_sources.errors import DataSourceError
|
||||
from app.infrastructure.data_sources.sina import SinaProvider
|
||||
from app.infrastructure.data_sources.tushare import TushareProvider
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import StockModel
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.index_impl import (
|
||||
SqlAlchemyIndexConstituentRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyAdjustFactorRepository,
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyDailyBasicRepository,
|
||||
SqlAlchemyFinancialRepository,
|
||||
SqlAlchemyStockNameHistoryRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
SqlAlchemySyncLogRepository,
|
||||
SqlAlchemyTradingCalendarRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.session import SessionLocal
|
||||
|
||||
_DATE_FMT = "%Y%m%d"
|
||||
|
||||
|
||||
def _parse_day(text: str) -> date:
|
||||
return datetime.strptime(text, _DATE_FMT).date()
|
||||
|
||||
|
||||
def _failover_provider(session):
|
||||
"""Tushare 首选 + 新浪兜底(basic/calendar 用;daily/financial 走校验兜底服务)。
|
||||
|
||||
FailoverProvider 每次尝试写 sync_log(AGENT.md §7)。能力矩阵:新浪仅提供
|
||||
日线/财务,basic/calendar 新浪不支持 → 抛错保留单源语义,日志可见。
|
||||
"""
|
||||
from app.infrastructure.data_sources.failover import FailoverProvider
|
||||
from app.infrastructure.data_sources.sina import SinaProvider
|
||||
|
||||
audit_repo = SqlAlchemySyncLogRepository(session)
|
||||
primary = TushareProvider(token=get_settings().tushare_token)
|
||||
return FailoverProvider(primary, fallback=SinaProvider(), audit=audit_repo.add)
|
||||
|
||||
|
||||
def _session_ctx():
|
||||
return SessionLocal()
|
||||
|
||||
|
||||
def cmd_basic(args) -> int:
|
||||
"""股票基础信息。--include-delisted 同时拉取已退市/暂停上市(幸存者偏差修正)。"""
|
||||
statuses = ["L", "P", "D"] if getattr(args, "include_delisted", False) else ["L"]
|
||||
with _session_ctx() as session:
|
||||
provider = _failover_provider(session)
|
||||
repo = SqlAlchemyStockRepository(session)
|
||||
total = 0
|
||||
failed: list[str] = []
|
||||
for st in statuses:
|
||||
try:
|
||||
stocks = provider.get_stock_basic(st)
|
||||
except DataSourceError as exc:
|
||||
# failover 会把备用源(新浪)的 NotSupported 包装成 DataSourceError,
|
||||
# 因此这里必须捕获 DataSourceError 而不是 DataSourceNotSupported,
|
||||
# 否则 --include-delisted 在主源抖动时会整体抛栈退出(L 已写、P/D 静默缺失)。
|
||||
print(f"[error] list_status={st} 拉取失败:{exc}", file=sys.stderr)
|
||||
failed.append(st)
|
||||
continue
|
||||
touched = repo.upsert_many(stocks)
|
||||
session.commit()
|
||||
total += touched
|
||||
n_delisted = sum(1 for s in stocks if s.delist_date is not None)
|
||||
print(
|
||||
f"[basic] list_status={st} 拉取 {len(stocks)} 只(含 delist_date {n_delisted} 只),"
|
||||
f"落库 {touched} 条"
|
||||
)
|
||||
if len(statuses) > 1:
|
||||
print(
|
||||
"[basic] 已退市股票已入库;其历史行情需另行同步(否则回测仍无法使用):\n"
|
||||
" 请执行 sync daily --symbols <退市代码逗号分隔> --start 20200101"
|
||||
)
|
||||
print(f"[basic] 合计落库 {total} 条")
|
||||
if failed:
|
||||
print(
|
||||
f"[error] 以下 list_status 拉取失败:{','.join(failed)};"
|
||||
"退市股缺失会让回测重新出现幸存者偏差,请重跑",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 1
|
||||
return 0
|
||||
|
||||
|
||||
def cmd_calendar(args) -> int:
|
||||
start = _parse_day(args.start)
|
||||
end = _parse_day(args.end)
|
||||
with _session_ctx() as session:
|
||||
provider = _failover_provider(session)
|
||||
days = provider.get_trade_cal(start, end)
|
||||
repo = SqlAlchemyTradingCalendarRepository(session)
|
||||
touched = repo.upsert_many(days)
|
||||
session.commit()
|
||||
open_days = sum(1 for d in days if d.is_open)
|
||||
print(f"[calendar] {start}~{end} 共 {len(days)} 条(交易日 {open_days}),落库 {touched} 条")
|
||||
return 0
|
||||
|
||||
|
||||
def cmd_namechange(args) -> int:
|
||||
"""股票名称变更历史同步(时点 ST 判定依据;按自然年分片)。"""
|
||||
start = _parse_day(args.start) if args.start else date(1990, 1, 1)
|
||||
end = _parse_day(args.end) if args.end else date.today()
|
||||
if start > end:
|
||||
print("[error] --start 不能晚于 --end", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
with _session_ctx() as session:
|
||||
audit_repo = SqlAlchemySyncLogRepository(session)
|
||||
primary = TushareProvider(token=get_settings().tushare_token)
|
||||
repo = SqlAlchemyStockNameHistoryRepository(session)
|
||||
syncer = NameHistorySyncer(primary=primary, repo=repo, audit=audit_repo.add)
|
||||
|
||||
total_span = (end.year - start.year) + 1
|
||||
print(f"[namechange] {start}~{end} 共 {total_span} 个年度分片")
|
||||
ok = failed = written = fetched = 0
|
||||
t0 = time.time()
|
||||
|
||||
def _on(idx: int, total: int, chunk_start: date) -> None:
|
||||
_progress(idx, total, t0, f"{chunk_start.year} 年")
|
||||
|
||||
results = syncer.sync_range(start, end, on_progress=_on)
|
||||
for res in results:
|
||||
if res.status == "ok":
|
||||
ok += 1
|
||||
written += res.rows_written
|
||||
fetched += res.rows_fetched
|
||||
else:
|
||||
failed += 1
|
||||
for note in res.notes:
|
||||
print(f"\n[warn] {res.start}~{res.end}: {note}", file=sys.stderr)
|
||||
session.commit()
|
||||
lo, hi = repo.namechange_dates()
|
||||
print(
|
||||
f"\n[namechange] 分片成功 {ok} / 失败 {failed};"
|
||||
f"拉取 {fetched} 行、写入 {written} 行;"
|
||||
f"本地生效起点 {lo} ~ {hi}"
|
||||
)
|
||||
return 1 if failed else 0
|
||||
|
||||
|
||||
def _symbols_of(args) -> list[str]:
|
||||
if getattr(args, "all", False):
|
||||
with _session_ctx() as session:
|
||||
symbols = list(session.scalars(select(StockModel.symbol).order_by(StockModel.symbol)))
|
||||
if not symbols:
|
||||
print("[error] stock 表为空,请先运行:python -m app.cli.sync basic")
|
||||
sys.exit(2)
|
||||
return symbols
|
||||
return [s.strip() for s in args.symbols.split(",") if s.strip()]
|
||||
|
||||
|
||||
def _stock_names(session, symbols: list[str]) -> dict[str, str]:
|
||||
"""一次性取出股票名称(进度描述用);批量查询避开 SQLite 变量上限。"""
|
||||
names: dict[str, str] = {}
|
||||
for i in range(0, len(symbols), 500):
|
||||
chunk = symbols[i : i + 500]
|
||||
rows = session.execute(
|
||||
select(StockModel.symbol, StockModel.name).where(StockModel.symbol.in_(chunk))
|
||||
)
|
||||
names.update({sym: nm for sym, nm in rows})
|
||||
return names
|
||||
|
||||
|
||||
def _warn_notes(notes: list[str]) -> None:
|
||||
for note in notes:
|
||||
print(f" [warn] {note}", file=sys.stderr)
|
||||
|
||||
|
||||
def cmd_daily(args) -> int:
|
||||
from sqlalchemy import func
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import StockDailyModel
|
||||
|
||||
symbols = _symbols_of(args)
|
||||
start = _parse_day(args.start) if args.start else date(2005, 1, 1)
|
||||
end = _parse_day(args.end) if args.end else date.today()
|
||||
started = time.monotonic()
|
||||
n_ok = n_sina = n_failed = n_skip = 0
|
||||
rows_tushare = rows_sina = 0
|
||||
with _session_ctx() as session:
|
||||
names = _stock_names(session, symbols)
|
||||
audit = SqlAlchemySyncLogRepository(session).add
|
||||
syncer = VerifiedDailySyncer(
|
||||
primary=TushareProvider(token=get_settings().tushare_token),
|
||||
fallback=SinaProvider(),
|
||||
bars=SqlAlchemyDailyBarRepository(session),
|
||||
factors=SqlAlchemyAdjustFactorRepository(session),
|
||||
audit=audit,
|
||||
)
|
||||
bar_repo = SqlAlchemyDailyBarRepository(session)
|
||||
# 增量基准:本地数据已到该日期即视为「已最新」,resume 时不再调 API
|
||||
global_latest = (
|
||||
session.scalar(select(func.max(StockDailyModel.trade_date))) if args.resume else None
|
||||
)
|
||||
for i, symbol in enumerate(symbols, start=1):
|
||||
begin = start
|
||||
if args.resume:
|
||||
latest = bar_repo.latest_date(symbol)
|
||||
if latest is not None:
|
||||
if global_latest is not None and latest >= global_latest:
|
||||
n_skip += 1 # 已同步到本地最新交易日,无需续拉
|
||||
continue
|
||||
begin = max(begin, latest + timedelta(days=1))
|
||||
if begin > end:
|
||||
n_skip += 1 # 无待拉区间(如区间已含在本地)
|
||||
continue
|
||||
if getattr(args, "sleep", 0) > 0:
|
||||
time.sleep(args.sleep)
|
||||
res: DailySymbolResult = syncer.sync_symbol(symbol, begin, end)
|
||||
session.commit() # 逐只落库:中断/报错只丢当前一只,重跑增量续传
|
||||
if res.status == "ok":
|
||||
n_ok += 1
|
||||
rows_tushare += res.bars_written
|
||||
elif res.status == "sina":
|
||||
n_sina += 1
|
||||
rows_sina += res.bars_written
|
||||
elif res.status == "failed":
|
||||
n_failed += 1
|
||||
_warn_notes(res.notes)
|
||||
if i % 100 == 0:
|
||||
name = names.get(symbol, "")
|
||||
print(
|
||||
f" ... {i}/{len(symbols)} {symbol} {name}: "
|
||||
f"累计 tushare {rows_tushare} 根 + 新浪补缺 {rows_sina} 根;"
|
||||
f"成功 {n_ok} / 新浪 {n_sina} / 失败待重试 {n_failed}"
|
||||
)
|
||||
elapsed = time.monotonic() - started
|
||||
detail = (
|
||||
f"[daily] {len(symbols)} 只股票:成功 {n_ok} / 新浪校验补缺 {n_sina} / "
|
||||
f"已最新跳过 {n_skip} / 失败待重试 {n_failed}"
|
||||
)
|
||||
if args.resume:
|
||||
detail += f"(本地最新 {global_latest})"
|
||||
detail += f";写入 {rows_tushare} 根(tushare 不复权)+ {rows_sina} 根(sina 前复权),耗时 {elapsed:.0f}s"
|
||||
print(detail)
|
||||
return 0
|
||||
|
||||
|
||||
def _fin_progress_line(i: int, n: int, symbol: str, name: str, res: FinancialSymbolResult) -> str:
|
||||
"""financial 逐只进度行:结果 + 导入内容简单描述(报告期/公告区间、来源)。"""
|
||||
head = f"[financial {i}/{n}] {symbol} {name or ''}".rstrip()
|
||||
if res.status == "skip":
|
||||
return f"{head}:已最新,跳过(增量)"
|
||||
if res.status == "failed":
|
||||
return f"{head}:失败待重试(tushare 失败;新浪源 {'未通过校验' if res.source == 'sina' else '不可用'})"
|
||||
if res.status == "sina":
|
||||
return (
|
||||
f"{head}:tushare 失败 → 新浪校验通过,补入 {res.written} 行(source=sina)"
|
||||
+ _fin_span(res)
|
||||
)
|
||||
# status == ok(tushare 成功)
|
||||
if res.written:
|
||||
updated = f",覆盖更新 {res.updated} 行" if res.updated else ""
|
||||
return f"{head}:tushare 返回 {res.fetched} 行 → 新增 {res.written} 行{updated}" + _fin_span(res)
|
||||
return f"{head}:tushare 返回 {res.fetched} 行,均已在库,无新增"
|
||||
|
||||
|
||||
def _fin_span(res: FinancialSymbolResult) -> str:
|
||||
if not res.written or res.report_first is None:
|
||||
return ""
|
||||
if res.announce_first is None or res.announce_last is None:
|
||||
return ""
|
||||
return (
|
||||
f";报告期 {res.report_first.isoformat()}~{res.report_last.isoformat()}"
|
||||
f"(公告 {res.announce_first.isoformat()}~{res.announce_last.isoformat()})"
|
||||
)
|
||||
|
||||
|
||||
def cmd_financial(args) -> int:
|
||||
symbols = _symbols_of(args)
|
||||
started = time.monotonic()
|
||||
n_ok = n_sina = n_failed = n_skip = 0
|
||||
rows_tushare = rows_sina = 0
|
||||
with _session_ctx() as session:
|
||||
names = _stock_names(session, symbols)
|
||||
audit = SqlAlchemySyncLogRepository(session).add
|
||||
syncer = VerifiedFinancialSyncer(
|
||||
primary=TushareProvider(token=get_settings().tushare_token),
|
||||
fallback=SinaProvider(),
|
||||
repo=SqlAlchemyFinancialRepository(session),
|
||||
audit=audit,
|
||||
)
|
||||
for i, symbol in enumerate(symbols, start=1):
|
||||
if getattr(args, "sleep", 0) > 0:
|
||||
time.sleep(args.sleep)
|
||||
res: FinancialSymbolResult = syncer.sync_symbol(symbol, force_full=args.full)
|
||||
session.commit() # 逐只落库:中断只丢当前一只,重跑增量续传
|
||||
print(_fin_progress_line(i, len(symbols), symbol, names.get(symbol, ""), res))
|
||||
_warn_notes(res.notes)
|
||||
if res.status == "ok":
|
||||
n_ok += 1
|
||||
rows_tushare += res.written
|
||||
elif res.status == "sina":
|
||||
n_sina += 1
|
||||
rows_sina += res.written
|
||||
elif res.status == "failed":
|
||||
n_failed += 1
|
||||
elif res.status == "skip":
|
||||
n_skip += 1
|
||||
elapsed = time.monotonic() - started
|
||||
mode = "全量重拉(--full)" if args.full else "增量"
|
||||
print(
|
||||
f"[financial] 共 {len(symbols)} 只({mode}):成功 {n_ok} / 新浪校验兜底 {n_sina} / "
|
||||
f"已最新跳过 {n_skip} / 失败待重试 {n_failed};"
|
||||
f"合计写入 {rows_tushare + rows_sina} 行(tushare {rows_tushare} + sina {rows_sina}),"
|
||||
f"耗时 {elapsed:.0f}s"
|
||||
)
|
||||
if n_failed:
|
||||
print(
|
||||
" [tip] 失败股票未写入未核验数据,重跑本命令即可续传补齐;"
|
||||
"若因频率超限,可用 --sleep 加大间隔(如 --sleep 60)分多次跑。",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
def cmd_index_weight(args) -> int:
|
||||
"""同步指数历史成分(Tushare index_weight;新浪不支持 → failover 审计留痕)。"""
|
||||
|
||||
code = args.code
|
||||
with _session_ctx() as session:
|
||||
provider = _failover_provider(session)
|
||||
try:
|
||||
rows = provider.get_index_weight(code)
|
||||
except DataSourceError as exc:
|
||||
print(f"[index_weight] {code} 失败:{exc}")
|
||||
return 1
|
||||
repo = SqlAlchemyIndexConstituentRepository(session)
|
||||
touched = repo.upsert_many(rows)
|
||||
latest = repo.latest_date(code)
|
||||
session.commit()
|
||||
print(
|
||||
f"[index_weight] {code} 拉取 {len(rows)} 期成分行,落库 {touched} 条"
|
||||
f",最新快照 {latest}(as_of 查询见 Universe.index_code)"
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
def cmd_verify(args) -> int:
|
||||
"""新浪交叉验证:取新浪最新前复权收盘,与本地最新交易日对照。
|
||||
|
||||
注意:新浪为前复权口径,数值不直接等于本地不复权收盘,
|
||||
本命令仅用于确认新浪可用性 / 最新交易日,不把新浪数据并入主库。
|
||||
"""
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
)
|
||||
|
||||
sina = SinaProvider()
|
||||
end = date.today()
|
||||
start = end - timedelta(days=20)
|
||||
try:
|
||||
bars = sina.get_daily(args.symbol, start, end)
|
||||
except DataSourceError as exc:
|
||||
print(f"[verify] 新浪不可用: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
if not bars:
|
||||
print(f"[verify] 新浪最近无数据({args.symbol})")
|
||||
return 1
|
||||
latest = max(bars, key=lambda b: b.trade_date)
|
||||
with _session_ctx() as session:
|
||||
local = SqlAlchemyDailyBarRepository(session).latest_date(args.symbol)
|
||||
print(
|
||||
f"[verify] {args.symbol}: 新浪最新 {latest.trade_date} 收盘(前复权) {latest.close};"
|
||||
f"本地最新交易日 {local}"
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
def cmd_export(args) -> int:
|
||||
"""把 SQLite 日线按年导出为 Parquet(data/parquet/stock_daily/<year>.parquet)。"""
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import StockDailyModel
|
||||
|
||||
settings = get_settings()
|
||||
out_root = settings.storage.get("parquet_dir") or Path("data/parquet")
|
||||
out_root.mkdir(parents=True, exist_ok=True)
|
||||
total = 0
|
||||
years = [int(y) for y in (args.years or "").split(",") if y.strip()] or None
|
||||
with _session_ctx() as session:
|
||||
all_bars = session.execute(
|
||||
select(StockDailyModel).order_by(StockDailyModel.trade_date)
|
||||
).scalars()
|
||||
frame = pd.DataFrame(
|
||||
[
|
||||
{
|
||||
"symbol": b.symbol,
|
||||
"trade_date": b.trade_date,
|
||||
"open": float(b.open) if b.open is not None else None,
|
||||
"high": float(b.high) if b.high is not None else None,
|
||||
"low": float(b.low) if b.low is not None else None,
|
||||
"close": float(b.close) if b.close is not None else None,
|
||||
"volume": float(b.volume) if b.volume is not None else None,
|
||||
"amount": float(b.amount) if b.amount is not None else None,
|
||||
}
|
||||
for b in all_bars
|
||||
]
|
||||
)
|
||||
if frame.empty:
|
||||
print("[export] 无日线数据,请先运行 sync daily")
|
||||
return 0
|
||||
frame["trade_date"] = pd.to_datetime(frame["trade_date"])
|
||||
out_dir = out_root / "stock_daily"
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
for year, group in frame.groupby(frame["trade_date"].dt.year):
|
||||
if years and int(year) not in years:
|
||||
continue
|
||||
path = out_dir / f"{year}.parquet"
|
||||
group.sort_values(["symbol", "trade_date"]).to_parquet(path, index=False)
|
||||
total += len(group)
|
||||
print(f"[export] {year} → {path}({len(group)} 行)")
|
||||
print(f"[export] 合计 {total} 行 → {out_dir}")
|
||||
return 0
|
||||
|
||||
|
||||
_DAILY_BASIC_DEFAULT_START = date(2020, 1, 1)
|
||||
|
||||
|
||||
def cmd_daily_basic(args) -> int:
|
||||
"""每日指标同步(按交易日整表;新浪不支持 → 失败如实记录,不留静默缺口)。"""
|
||||
start = _parse_day(args.start) if args.start else _DAILY_BASIC_DEFAULT_START
|
||||
end = _parse_day(args.end) if args.end else date.today()
|
||||
if start > end:
|
||||
print("[error] --start 不能晚于 --end", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
with _session_ctx() as session:
|
||||
audit_repo = SqlAlchemySyncLogRepository(session)
|
||||
primary = TushareProvider(token=get_settings().tushare_token)
|
||||
repo = SqlAlchemyDailyBasicRepository(session)
|
||||
syncer = DailyBasicSyncer(primary=primary, repo=repo, audit=audit_repo.add)
|
||||
|
||||
if args.full:
|
||||
cal = SqlAlchemyTradingCalendarRepository(session)
|
||||
days = [d.calendar_date for d in cal.list_range(start, end) if d.is_open]
|
||||
else:
|
||||
days = repo.missing_dates(start, end)
|
||||
total = len(days)
|
||||
if total == 0:
|
||||
print(f"[daily_basic] {start}~{end} 无待补交易日(本地已完整)")
|
||||
return 0
|
||||
print(f"[daily_basic] {start}~{end} 待同步 {total} 个交易日")
|
||||
|
||||
ok = failed = written = 0
|
||||
failures: list[date] = []
|
||||
t0 = time.time()
|
||||
for idx, day in enumerate(days, start=1):
|
||||
try:
|
||||
res = syncer.sync_day(day)
|
||||
except DataSourceError as exc:
|
||||
print(f"\n[error] 第 {idx}/{total} 日 {day} 权限/凭证故障,中止:{exc}", file=sys.stderr)
|
||||
session.commit()
|
||||
return 1
|
||||
if res.status == "ok":
|
||||
ok += 1
|
||||
written += res.rows_written
|
||||
else:
|
||||
failed += 1
|
||||
failures.append(day)
|
||||
_progress(idx, total, t0, f"{day} 行数={res.rows_fetched}")
|
||||
if idx % 20 == 0 or idx == total:
|
||||
session.commit() # 分批提交:中断时已完成的部分保持有效
|
||||
if args.sleep:
|
||||
time.sleep(args.sleep)
|
||||
session.commit()
|
||||
|
||||
span = time.time() - t0
|
||||
print(
|
||||
f"\n[daily_basic] 完成:成功 {ok} 日 / 失败 {failed} 日,累计写入 {written} 行,"
|
||||
f"耗时 {span:.1f}s"
|
||||
)
|
||||
if failures:
|
||||
print(
|
||||
"[daily_basic] 失败交易日(可重跑本命令补齐):"
|
||||
+ ", ".join(d.isoformat() for d in failures[:20])
|
||||
+ (" …" if len(failures) > 20 else ""),
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 1
|
||||
return 0
|
||||
|
||||
|
||||
def _progress(idx: int, total: int, t0: float, extra: str = "") -> None:
|
||||
"""单行进度条(同步长任务可读性)。"""
|
||||
elapsed = time.time() - t0
|
||||
rate = idx / elapsed if elapsed > 0 else 0.0
|
||||
eta = (total - idx) / rate if rate > 0 else 0.0
|
||||
pct = idx / total * 100 if total else 100.0
|
||||
end = "\n" if idx >= total else "\r"
|
||||
print(
|
||||
f" 进度 {idx}/{total} ({pct:5.1f}%) 已用 {elapsed:6.1f}s 预计剩余 {eta:6.1f}s {extra}",
|
||||
end=end,
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(prog="app.cli.sync", description="Tushare 数据同步 CLI")
|
||||
sub = parser.add_subparsers(dest="command", required=True)
|
||||
|
||||
p_nc = sub.add_parser("namechange", help="同步股票名称变更历史(时点 ST 判定)")
|
||||
p_nc.add_argument("--start", help="起始日 YYYYMMDD(默认 19900101,覆盖全历史)")
|
||||
p_nc.add_argument("--end", help="结束日 YYYYMMDD(默认今天)")
|
||||
p_nc.set_defaults(func=cmd_namechange)
|
||||
|
||||
p_basic = sub.add_parser("basic", help="同步股票基础信息")
|
||||
p_basic.add_argument(
|
||||
"--include-delisted",
|
||||
action="store_true",
|
||||
help="同时拉取已退市(D)与暂停上市(P),填充 delist_date(幸存者偏差修正)",
|
||||
)
|
||||
p_basic.set_defaults(func=cmd_basic)
|
||||
|
||||
p_cal = sub.add_parser("calendar", help="同步交易日历")
|
||||
p_cal.add_argument("--start", required=True, help="YYYYMMDD")
|
||||
p_cal.add_argument("--end", required=True, help="YYYYMMDD")
|
||||
p_cal.set_defaults(func=cmd_calendar)
|
||||
|
||||
p_db = sub.add_parser(
|
||||
"daily_basic",
|
||||
help="同步每日指标(估值/股息率/市值,按交易日整表;新浪不支持本接口)",
|
||||
)
|
||||
p_db.add_argument("--start", default="20200101", help="YYYYMMDD(默认 20200101)")
|
||||
p_db.add_argument("--end", default="", help="YYYYMMDD(默认今天)")
|
||||
p_db.add_argument(
|
||||
"--full",
|
||||
action="store_true",
|
||||
help="忽略本地已有日期,重拉区间内全部开市日(默认只补缺失日)",
|
||||
)
|
||||
p_db.add_argument(
|
||||
"--sleep",
|
||||
type=float,
|
||||
default=0,
|
||||
help="每个交易日请求间隔秒数(限速时加大,如 0.2 或 1)",
|
||||
)
|
||||
p_db.set_defaults(func=cmd_daily_basic)
|
||||
|
||||
p_daily = sub.add_parser("daily", help="同步日线与复权因子(Tushare 失败 → 新浪校验兜底补缺)")
|
||||
p_daily.add_argument("--symbols", default="", help="600519.SH,000001.SZ")
|
||||
p_daily.add_argument("--all", action="store_true", help="遍历 stock 表全部股票")
|
||||
p_daily.add_argument("--start", default="", help="YYYYMMDD(默认 20050101)")
|
||||
p_daily.add_argument("--end", default="", help="YYYYMMDD(默认今天)")
|
||||
p_daily.add_argument("--resume", action="store_true", help="从本地最新交易日续传(增量)")
|
||||
p_daily.add_argument(
|
||||
"--sleep",
|
||||
type=float,
|
||||
default=0,
|
||||
help="每只股票请求间隔秒数(限速时加大,如 1 或 60)",
|
||||
)
|
||||
p_daily.set_defaults(func=cmd_daily)
|
||||
|
||||
p_fin = sub.add_parser("financial", help="同步财务指标快照(默认增量;Tushare 失败 → 新浪校验兜底)")
|
||||
p_fin.add_argument("--symbols", default="")
|
||||
p_fin.add_argument("--all", action="store_true", help="遍历 stock 表全部股票")
|
||||
p_fin.add_argument(
|
||||
"--full",
|
||||
action="store_true",
|
||||
help="强制全量重拉并覆盖既有行(默认只补本地缺失/更新的报告期,已最新跳过)",
|
||||
)
|
||||
p_fin.add_argument(
|
||||
"--sleep",
|
||||
type=float,
|
||||
default=0,
|
||||
help="每只股票请求间隔秒数(限速时加大,如 1 或 60)",
|
||||
)
|
||||
p_fin.set_defaults(func=cmd_financial)
|
||||
|
||||
p_idx = sub.add_parser("index_weight", help="同步指数历史成分(如沪深300 000300.SH)")
|
||||
p_idx.add_argument("--code", required=True, help="指数代码,如 000300.SH / 000905.SH")
|
||||
p_idx.set_defaults(func=cmd_index_weight)
|
||||
|
||||
p_verify = sub.add_parser("verify", help="新浪交叉验证最新行情")
|
||||
p_verify.add_argument("--symbol", required=True)
|
||||
p_verify.set_defaults(func=cmd_verify)
|
||||
|
||||
p_export = sub.add_parser("export", help="日线按年导出 Parquet(data/parquet)")
|
||||
p_export.add_argument("--years", default="", help="逗号分隔年份,留空导出全部")
|
||||
p_export.set_defaults(func=cmd_export)
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
args = build_parser().parse_args(argv)
|
||||
try:
|
||||
return args.func(args)
|
||||
except DataSourceError as exc:
|
||||
print(f"[error] {exc}", file=sys.stderr)
|
||||
return 1
|
||||
except KeyboardInterrupt:
|
||||
print("\n[interrupt] 已中止", file=sys.stderr)
|
||||
return 130
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
+164
-1
@@ -68,12 +68,43 @@ class Settings:
|
||||
data_source_primary: str
|
||||
data_source_fallback: str
|
||||
tushare_token: str
|
||||
llm_api_key: str
|
||||
llm_base_url: str | None
|
||||
llm_model: str
|
||||
storage: dict[str, Path]
|
||||
job_mode: str
|
||||
job_memory_limit_gb: int
|
||||
job_max_concurrent: int
|
||||
# research.archive_curve_limit:归档个股收益曲线数量上限;None = 完整存档(默认)
|
||||
research_archive_curve_limit: int | None
|
||||
# research.archive_max_chars:归档结果 JSON 的字节预算(MEDIUMTEXT 上限 16MB 的余量)
|
||||
research_archive_max_chars: int | None
|
||||
config_path: Path = CONFIG_PATH
|
||||
env_path: Path = ENV_PATH
|
||||
project_root: Path = PROJECT_ROOT
|
||||
|
||||
|
||||
def _env_or(value_env: str, configured: object) -> object:
|
||||
"""环境变量优先于 config.yaml(空串视为未设置)。"""
|
||||
env = os.environ.get(value_env)
|
||||
if env not in (None, ""):
|
||||
return env
|
||||
return configured
|
||||
|
||||
|
||||
def _optional_int(value: object) -> int | None:
|
||||
"""可选整数配置:None / 空串 / 非法值都视为「未配置」(返回 None)。
|
||||
|
||||
用于 `research.archive_curve_limit`:None 表示不截断(完整存档)。
|
||||
"""
|
||||
if value is None or value == "":
|
||||
return None
|
||||
try:
|
||||
return int(value) # type: ignore[arg-type]
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _resolve_storage_dirs(cfg: dict) -> dict[str, Path]:
|
||||
storage_cfg = _deep(cfg, "storage") or {}
|
||||
return {
|
||||
@@ -93,6 +124,94 @@ def _normalize_sqlite_url(url: str) -> str:
|
||||
return f"{_SQLITE_PREFIX}{(PROJECT_ROOT / rest).resolve()}"
|
||||
|
||||
|
||||
_URL_RESERVED = set("@:/?#%")
|
||||
|
||||
|
||||
def _quote_userinfo(value: str) -> str:
|
||||
"""仅编码会破坏 SQLAlchemy URL 解析的字符(@ : / ? # % 与空白)。
|
||||
|
||||
其余字符(含 ! - _ . 等非保留符)原样保留:避免 URL 出现 %XX 干扰
|
||||
Alembic configparser 的 interpolation。
|
||||
"""
|
||||
out: list[str] = []
|
||||
for ch in value:
|
||||
if ch in _URL_RESERVED or ch.isspace():
|
||||
out.append(f"%{ord(ch):02X}")
|
||||
elif ord(ch) > 127: # 非 ASCII:按 UTF-8 逐字节 percent 编码
|
||||
out.append("".join(f"%{b:02X}" for b in ch.encode("utf-8")))
|
||||
else:
|
||||
out.append(ch)
|
||||
return "".join(out)
|
||||
|
||||
|
||||
def _build_mysql_url(mysql: dict | None) -> str | None:
|
||||
"""由 config.yaml database.mysql 段组装 mysql+pymysql URL。
|
||||
|
||||
规则:仅当 enabled 且 host/user/db 齐全时返回 URL;密码经 password_env
|
||||
指定的环境变量读取(AGENT.md §33 密钥只放 .env),未设置则按空密码处理。
|
||||
"""
|
||||
if not mysql or not mysql.get("enabled"):
|
||||
return None
|
||||
host = mysql.get("host")
|
||||
user = mysql.get("user")
|
||||
db = mysql.get("db")
|
||||
if not (host and user and db):
|
||||
return None
|
||||
port = mysql.get("port") or 3306
|
||||
charset = mysql.get("charset") or "utf8mb4"
|
||||
password = os.environ.get(mysql.get("password_env") or "MYSQL_PASSWORD", "")
|
||||
auth = (
|
||||
f"{_quote_userinfo(user)}:{_quote_userinfo(password)}"
|
||||
if password
|
||||
else _quote_userinfo(user)
|
||||
)
|
||||
return f"mysql+pymysql://{auth}@{host}:{port}/{db}?charset={charset}"
|
||||
|
||||
|
||||
# 被禁止的数据库目标主机(用户明确要求:禁止使用 192.168.1.10 作为 DB 目标)。
|
||||
# 可用环境变量 QLIB_FORBIDDEN_DB_HOSTS 覆盖(逗号分隔;设为空串表示不做限制)。
|
||||
_FORBIDDEN_DB_HOSTS_ENV = "QLIB_FORBIDDEN_DB_HOSTS"
|
||||
_FORBIDDEN_DB_HOSTS_DEFAULT = "192.168.1.10"
|
||||
|
||||
|
||||
def _forbidden_db_hosts() -> set[str]:
|
||||
raw = os.environ.get(_FORBIDDEN_DB_HOSTS_ENV)
|
||||
if raw is None:
|
||||
raw = _FORBIDDEN_DB_HOSTS_DEFAULT
|
||||
return {h.strip().lower() for h in raw.split(",") if h.strip()}
|
||||
|
||||
|
||||
def assert_db_target_allowed(database_url: str) -> None:
|
||||
"""拒绝把数据库指向被禁主机;命中即**抛错**,不允许「起得来但连错库」。
|
||||
|
||||
为什么必须硬失败(AGENT.md §7 不静默):库目标错了不会有任何报错或界面异常——
|
||||
回测结果、实验归档、策略、信号会**安静地读写另一台机器上的数据**,
|
||||
而用户以为看的是本机库;这类错误事后极难发现(数据看似正常,只是「不对」)。
|
||||
因此在这里直接失败并说明原因与修改方法。
|
||||
|
||||
只拦主机名/IP 精确匹配(不做网段推断):`192.168.1.10` 与写进 userinfo 的
|
||||
同名字符串、SQLite 路径都不会误判;IPv6/带方括号的地址会去掉括号后比较。
|
||||
"""
|
||||
forbidden = _forbidden_db_hosts()
|
||||
if not forbidden or not database_url:
|
||||
return
|
||||
try:
|
||||
from sqlalchemy.engine import make_url
|
||||
|
||||
url = make_url(database_url)
|
||||
except Exception: # noqa: BLE001 —— 解析失败交给 SQLAlchemy 自己报错,这里不抢
|
||||
return
|
||||
host = (url.host or "").lower().strip("[]")
|
||||
if host and host in forbidden:
|
||||
raise RuntimeError(
|
||||
f"数据库目标 {host} 已被禁止({_FORBIDDEN_DB_HOSTS_ENV}="
|
||||
f"{','.join(sorted(forbidden))})。"
|
||||
"本项目只允许使用**本机 MariaDB**:把 config.yaml 的 database.mysql.host "
|
||||
"设为 127.0.0.1,或清空 .env 的 DATABASE_URL 让它走本机配置。"
|
||||
"确需临时放开请显式设置 QLIB_FORBIDDEN_DB_HOSTS(设为空串表示不限制)。"
|
||||
)
|
||||
|
||||
|
||||
@lru_cache
|
||||
def get_settings() -> Settings:
|
||||
"""加载并缓存 Settings。config / env 路径可通过参数覆盖以便测试。"""
|
||||
@@ -102,14 +221,50 @@ def get_settings() -> Settings:
|
||||
secret_env = _deep(cfg, "app.secret_key_env") or "APP_SECRET_KEY"
|
||||
url_env = _deep(cfg, "database.url_env") or "DATABASE_URL"
|
||||
token_env = _deep(cfg, "data_source.tushare_token_env") or "TUSHARE_TOKEN"
|
||||
# Agent LLM:URL/模型名在 config.yaml 明文,api_key 只从 .env 读(env 可覆盖 url/model)
|
||||
agent_llm = _deep(cfg, "agent.llm") or {}
|
||||
llm_key_env = agent_llm.get("api_key_env") or "LLM_API_KEY"
|
||||
llm_url_env = agent_llm.get("base_url_env") or "LLM_BASE_URL"
|
||||
llm_model_env = agent_llm.get("model_env") or "LLM_MODEL"
|
||||
|
||||
default_url = "sqlite:///./data/quant.db"
|
||||
database_url = os.environ.get(url_env) or default_url
|
||||
# URL 优先级:环境变量(DATABASE_URL) > config.yaml database.mysql 段 > SQLite 兜底
|
||||
database_url = os.environ.get(url_env) or _build_mysql_url(
|
||||
_deep(cfg, "database.mysql") or {}
|
||||
) or default_url
|
||||
|
||||
# 用户约束:禁止把数据写向 192.168.1.10(远端库)——命中直接抛错,不静默降级
|
||||
assert_db_target_allowed(database_url)
|
||||
|
||||
migrations_rel = _deep(cfg, "database.migrations_dir") or (
|
||||
"app/infrastructure/persistence/migrations"
|
||||
)
|
||||
|
||||
# Job 执行:环境变量 JOB_MODE / QLIB_JOB_MEM_LIMIT_GB 可覆盖 config.yaml(测试用)
|
||||
job_cfg = _deep(cfg, "job") or {}
|
||||
try:
|
||||
job_memory_limit_gb = int(
|
||||
os.environ.get("QLIB_JOB_MEM_LIMIT_GB") or job_cfg.get("max_memory_gb") or 6
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
job_memory_limit_gb = 6
|
||||
try:
|
||||
job_max_concurrent = int(
|
||||
os.environ.get("QLIB_JOB_MAX_CONCURRENT") or job_cfg.get("max_concurrent_jobs") or 2
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
job_max_concurrent = 2
|
||||
|
||||
# 归档(research 段):曲线数量上限(None = 完整存档)+ 结果 JSON 字节预算
|
||||
research_cfg = _deep(cfg, "research")
|
||||
research_cfg = research_cfg if isinstance(research_cfg, dict) else {}
|
||||
archive_curve_limit = _optional_int(
|
||||
_env_or("QLIB_ARCHIVE_CURVE_LIMIT", research_cfg.get("archive_curve_limit"))
|
||||
)
|
||||
archive_max_chars = _optional_int(
|
||||
_env_or("QLIB_ARCHIVE_MAX_CHARS", research_cfg.get("archive_max_chars"))
|
||||
)
|
||||
|
||||
return Settings(
|
||||
app_name=str(_deep(cfg, "app.name") or "qlib-platform"),
|
||||
app_version=str(_deep(cfg, "app.version") or "0.1.0"),
|
||||
@@ -122,5 +277,13 @@ def get_settings() -> Settings:
|
||||
data_source_primary=str(_deep(cfg, "data_source.primary") or "tushare"),
|
||||
data_source_fallback=str(_deep(cfg, "data_source.fallback") or "sina"),
|
||||
tushare_token=os.environ.get(token_env, ""),
|
||||
llm_api_key=os.environ.get(llm_key_env, ""),
|
||||
llm_base_url=os.environ.get(llm_url_env) or (agent_llm.get("base_url") or None),
|
||||
llm_model=os.environ.get(llm_model_env) or agent_llm.get("model") or "qwen-plus",
|
||||
storage=_resolve_storage_dirs(cfg),
|
||||
job_mode=os.environ.get("JOB_MODE") or job_cfg.get("mode") or "subprocess",
|
||||
job_memory_limit_gb=job_memory_limit_gb,
|
||||
job_max_concurrent=job_max_concurrent,
|
||||
research_archive_curve_limit=archive_curve_limit,
|
||||
research_archive_max_chars=archive_max_chars,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
"""统一量化可视化 DTO(v3 §20 Chart Service)。
|
||||
|
||||
原则:前端只展示本结构,不得自行重算选股/信号/成交(v3 §20.1)。
|
||||
价格口径:bars 已按请求 adjust 折算(显示层);成交/信号标记与 K 线同坐标系;
|
||||
metadata 记录 adjust_mode 与回测执行价 basis,杜绝图表与回测口径混用(v3 §20.5)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class OHLC(BaseModel):
|
||||
time: date
|
||||
open: float | None = None
|
||||
high: float | None = None
|
||||
low: float | None = None
|
||||
close: float | None = None
|
||||
|
||||
|
||||
class VolumePoint(BaseModel):
|
||||
time: date
|
||||
value: float | None = None
|
||||
|
||||
|
||||
class SeriesPoint(BaseModel):
|
||||
"""指标/分数等 (时间, 值) 序列点。"""
|
||||
|
||||
time: date
|
||||
value: float | None = None
|
||||
|
||||
|
||||
class EventMarker(BaseModel):
|
||||
"""K 线上可点击的事件标记(选股/信号/实际成交)。"""
|
||||
|
||||
time: date
|
||||
kind: str = Field(
|
||||
description="selection | signal_buy | signal_sell | signal_watch | fill_buy | fill_sell"
|
||||
)
|
||||
symbol: str = ""
|
||||
price: float | None = None
|
||||
score: float | None = None
|
||||
text: list[str] = Field(default_factory=list, description="原因/说明(tooltip)")
|
||||
ref_id: str | None = Field(default=None, description="关联 selection/signal 记录 id")
|
||||
|
||||
|
||||
class ChartMetadata(BaseModel):
|
||||
symbol: str
|
||||
name: str = ""
|
||||
adjust_mode: str = Field(default="none", description="显示口径:none | qfq | hfq")
|
||||
price_basis: str = Field(default="chart_display", description="显示价基准(显示层折算)")
|
||||
execution_price_basis: str | None = Field(
|
||||
default=None, description="回测执行价口径(如 none),与显示口径不同时用于解释"
|
||||
)
|
||||
start: date | None = None
|
||||
end: date | None = None
|
||||
bar_count: int = 0
|
||||
indicator_windows: list[int] = Field(default_factory=list)
|
||||
|
||||
|
||||
class ChartResult(BaseModel):
|
||||
"""单只股票 / 回测个股的统一图表数据(v3 §20.2)。"""
|
||||
|
||||
metadata: ChartMetadata
|
||||
bars: list[OHLC] = Field(default_factory=list)
|
||||
volume: list[VolumePoint] = Field(default_factory=list)
|
||||
indicators: dict[str, list[SeriesPoint]] = Field(default_factory=dict)
|
||||
selections: list[EventMarker] = Field(default_factory=list)
|
||||
signals: list[EventMarker] = Field(default_factory=list)
|
||||
fills: list[EventMarker] = Field(default_factory=list)
|
||||
holding_periods: list[dict] = Field(default_factory=list)
|
||||
factor_values: dict[str, list[SeriesPoint]] = Field(default_factory=dict)
|
||||
strategy_scores: dict[str, list[SeriesPoint]] = Field(default_factory=dict)
|
||||
unimplemented: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class SelectionHit(BaseModel):
|
||||
"""个股在选股历史中的命中(by-symbol 查询)。"""
|
||||
|
||||
selection_id: str
|
||||
as_of: date
|
||||
method: str
|
||||
symbol: str
|
||||
rank: int
|
||||
score: float
|
||||
selection_reason: list[str] = Field(default_factory=list)
|
||||
@@ -0,0 +1,162 @@
|
||||
"""回测组合与公共配置领域实体(2026-09 重构)。
|
||||
|
||||
把原来「一个策略 = 全套参数」拆成三件独立的事(用户目标):
|
||||
|
||||
1. **GlobalConfig(公共配置,全局唯一)** —— 费率 / 印花税 / 滑点 / 最低佣金 /
|
||||
复权口径 / 基准。所有回测共用,不再塞进每个策略。
|
||||
2. **SelectionStrategy(选股策略,见 strategy.py)** —— 只剩「选股条件组合」:
|
||||
股票池 + 因子 + 过滤条件。**不含**资金 / 持仓数 / 持仓时间 / 调仓 / 费率 / 区间。
|
||||
3. **BacktestCombo(回测组合)** —— 引用若干选股策略 + 回测时才定的参数:
|
||||
起始资金、持仓数量 N、持仓天数区间 [Tmin, Tmax]、调仓时机(日/周/月)、回测区间。
|
||||
|
||||
执行时由服务层把「组合 + 被引用的选股策略 + 公共配置快照」解析成一个
|
||||
`ComboRunSpec`,喂给组合引擎;该 spec 会原样写进归档的 config_snapshot,
|
||||
保证事后可复现(AGENT.md §21),即使之后公共配置被改也不影响历史结果。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
||||
from app.domain.entities.research import CostSpec
|
||||
|
||||
# ---------- 公共配置(全局唯一) ----------
|
||||
|
||||
|
||||
class GlobalConfig(BaseModel):
|
||||
"""全局交易成本与行情口径(单例,id 恒为 "default")。
|
||||
|
||||
为什么把复权口径也放这里:一次回测只能有一个复权口径(同一份行情不能既前复权
|
||||
又后复权),而多个选股策略可能想混用 —— 与其让它们在组合里打架,不如统一为
|
||||
全局口径,高股息默认 hfq。若将来确需按组合区分,再加字段即可(向前兼容)。
|
||||
|
||||
`extra="forbid"`:PUT /api/config 若带未知字段(拼错键名、旧版遗留键)直接报错,
|
||||
避免「以为改了某项、其实被静默忽略」。
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
id: str = Field(default="default", description="单例主键,恒为 default")
|
||||
commission_rate: float = Field(default=0.0003, ge=0, le=0.01, description="佣金率(如 0.0003 = 万三)")
|
||||
stamp_tax_rate: float = Field(default=0.0005, ge=0, le=0.01, description="印花税率(仅卖出)")
|
||||
slippage_rate: float = Field(default=0.001, ge=0, le=0.05, description="滑点率")
|
||||
min_commission: float = Field(default=5.0, ge=0, le=100.0, description="单笔最低佣金(元)")
|
||||
price_adjustment: str = Field(
|
||||
default="hfq", pattern="^(none|qfq|hfq)$",
|
||||
description="行情复权口径:none / qfq / hfq(高股息类建议 hfq)",
|
||||
)
|
||||
benchmark: str = Field(default="000300.SH", description="对照基准指数代码")
|
||||
updated_at: datetime | None = None
|
||||
|
||||
def to_cost_spec(self) -> CostSpec:
|
||||
"""转成引擎用的 CostSpec(benchmark 一并带入)。"""
|
||||
return CostSpec(
|
||||
commission_rate=self.commission_rate,
|
||||
stamp_tax_rate=self.stamp_tax_rate,
|
||||
slippage_rate=self.slippage_rate,
|
||||
min_commission=self.min_commission,
|
||||
benchmark=self.benchmark,
|
||||
)
|
||||
|
||||
|
||||
# ---------- 回测组合 ----------
|
||||
|
||||
# 调仓时机:日 / 周 / 月(在原有 weekly/monthly 之上新增 daily)
|
||||
REBALANCE_FREQS = ("daily", "weekly", "monthly")
|
||||
|
||||
|
||||
class BacktestCombo(BaseModel):
|
||||
"""一个可保存、可复跑的回测组合。
|
||||
|
||||
持仓模型(用户确认的语义):
|
||||
- `hold_count` = N:目标持仓只数(等权)。
|
||||
- `hold_min_days` = Tmin:个股**最少**持有天数 —— 掉出 TopN 时若未满 Tmin 不卖
|
||||
(防止频繁换手);但超过 Tmax 仍强制卖(安全阀优先)。
|
||||
- `hold_max_days` = Tmax:个股**最多**持有天数 —— 超过即强制了结(None = 不限)。
|
||||
- `rebalance_freq`:多久重新打分排序并调仓一次(日/周/月)。
|
||||
⚠️ Tmax 强制卖出**每个交易日**都检查(不只调仓日),否则月频下会远超 Tmax。
|
||||
|
||||
`extra="forbid"`:回测参数写错键名(如 hold_days、capital)时报错而非静默用默认值 ——
|
||||
静默用默认值会让「我明明设了 30 天」变成「其实没生效」,是本项目明确禁止的降级方式。
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
id: str = ""
|
||||
name: str = Field(min_length=1, max_length=64)
|
||||
description: str = ""
|
||||
strategy_ids: list[str] = Field(
|
||||
min_length=1, description="引用的选股策略 id(≥1 个;多策略取并集后 Borda 秩和打分)"
|
||||
)
|
||||
initial_capital: float = Field(default=1_000_000.0, gt=0, description="起始资金(元)")
|
||||
hold_count: int = Field(ge=1, le=1000, description="目标持仓只数 N")
|
||||
hold_min_days: int = Field(default=0, ge=0, description="个股最少持有天数 Tmin")
|
||||
hold_max_days: int | None = Field(
|
||||
default=None, ge=1, description="个股最多持有天数 Tmax;None = 不强制了结"
|
||||
)
|
||||
rebalance_freq: str = Field(
|
||||
default="monthly", description="调仓时机:daily / weekly / monthly"
|
||||
)
|
||||
period: tuple[date, date]
|
||||
version: str = "1"
|
||||
created_at: datetime | None = None
|
||||
|
||||
@field_validator("rebalance_freq")
|
||||
@classmethod
|
||||
def _freq(cls, v: str) -> str:
|
||||
if v not in REBALANCE_FREQS:
|
||||
raise ValueError(f"rebalance_freq 必须是 {REBALANCE_FREQS} 之一,收到 {v!r}")
|
||||
return v
|
||||
|
||||
@field_validator("period")
|
||||
@classmethod
|
||||
def _period_ordered(cls, period: tuple[date, date]) -> tuple[date, date]:
|
||||
if period[0] >= period[1]:
|
||||
raise ValueError("period 必须满足 start < end")
|
||||
return period
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _hold_band_and_strategies(self) -> BacktestCombo:
|
||||
if (
|
||||
self.hold_max_days is not None
|
||||
and self.hold_min_days > 0
|
||||
and self.hold_max_days < self.hold_min_days
|
||||
):
|
||||
raise ValueError(
|
||||
f"持仓上限 Tmax={self.hold_max_days} 不能小于下限 Tmin={self.hold_min_days}"
|
||||
)
|
||||
if len(set(self.strategy_ids)) != len(self.strategy_ids):
|
||||
raise ValueError("strategy_ids 存在重复的策略 id")
|
||||
return self
|
||||
|
||||
|
||||
class ComboRunSpec(BaseModel):
|
||||
"""解析后的、可复现的组合运行规格(写入归档 config_snapshot)。
|
||||
|
||||
为什么不直接存 BacktestCombo:组合只引用 strategy_ids,且费率/复权来自公共配置;
|
||||
若事后策略被删改、公共配置被调整,光凭 combo 无法复现。这里把「当时用到的策略定义
|
||||
+ 当时的成本/复权快照」一起固化,归档即可独立复现(AGENT.md §21)。
|
||||
"""
|
||||
|
||||
combo: BacktestCombo
|
||||
strategies: list[SelectionStrategyRef] = Field(
|
||||
description="运行时刻各选股策略的快照(name/universe/factors/conditions)"
|
||||
)
|
||||
costs: CostSpec
|
||||
price_adjustment: str = Field(pattern="^(none|qfq|hfq)$")
|
||||
config_version: str = "1"
|
||||
|
||||
|
||||
class SelectionStrategyRef(BaseModel):
|
||||
"""ComboRunSpec 内嵌的策略快照(只取选股相关字段,避免把已废弃字段带进归档)。"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
universe: dict # UniverseSpec.model_dump()
|
||||
factors: list[dict] # [{name, weight}]
|
||||
conditions: list[dict] = Field(default_factory=list)
|
||||
|
||||
|
||||
ComboRunSpec.model_rebuild()
|
||||
@@ -0,0 +1,33 @@
|
||||
"""因子组合(Composite Factor)领域实体(M7.2b,v2 §13)。
|
||||
|
||||
组合 = 一组 {因子, 权重} + method(MVP fixed:截面 zscore×方向×权重求和;
|
||||
方向在计算时取因子注册表元数据,落库时冗余快照以便列表展示)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
|
||||
class CompositeComponent(BaseModel):
|
||||
name: str
|
||||
weight: float = Field(default=1.0, gt=0)
|
||||
direction: str = Field(default="higher_is_better")
|
||||
|
||||
|
||||
class CompositeDefinition(BaseModel):
|
||||
id: str = ""
|
||||
name: str = Field(min_length=1, max_length=64)
|
||||
method: str = Field(default="fixed", pattern="^(fixed)$")
|
||||
description: str = ""
|
||||
components: list[CompositeComponent] = Field(min_length=1)
|
||||
created_at: datetime | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _no_duplicate(self) -> CompositeDefinition:
|
||||
names = [c.name for c in self.components]
|
||||
if len(set(names)) != len(names):
|
||||
raise ValueError("components 存在重复因子名")
|
||||
return self
|
||||
@@ -0,0 +1,54 @@
|
||||
"""字段库领域实体(2026-10:过滤条件字段目录 DB 化)。
|
||||
|
||||
与因子目录(``FactorDefinition``)同一套思路:
|
||||
|
||||
- **DB 是字段库的契约源**:中文名、含义、单位、是否启用、自定义条目都入库;
|
||||
- **引擎是字段可用性的唯一事实来源**:字段能不能算由 ``quant.condition_fields``
|
||||
对着引擎域校验,登记不出来的字段一律拒绝(防「建出来永远选不出股票」的伪字段)。
|
||||
|
||||
可编辑边界(有意为之):
|
||||
- ``name`` 是引擎字段名,**不可改**(改了就指向另一个字段,等于换字段);
|
||||
- ``kind`` 由引擎类型决定,**不可改**(字符串字段不能比大小);
|
||||
- ``label`` / ``description`` / ``group_name`` / ``unit`` / ``enabled`` 可编辑 ——
|
||||
内置字段也允许改文案(seed「只补不删」,不会覆盖用户的措辞)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, computed_field
|
||||
|
||||
FIELD_KINDS: tuple[str, ...] = ("num", "str")
|
||||
FIELD_SOURCES: tuple[str, ...] = ("builtin", "custom")
|
||||
|
||||
# 类型 → 可用比较符(单一事实来源)。
|
||||
# 字符串字段只能等值/集合:引擎 _compare 对字符串的 >/≥/</≤ 一律返回 False,
|
||||
# 若前端把「行业 > 5」这类选项摆出来,用户点出来的就是永远为假的条件。
|
||||
OPS_NUM: tuple[str, ...] = ("gt", "gte", "lt", "lte", "eq", "ne")
|
||||
OPS_STR: tuple[str, ...] = ("eq", "ne", "in", "not_in")
|
||||
OPS_BY_KIND: dict[str, tuple[str, ...]] = {"num": OPS_NUM, "str": OPS_STR}
|
||||
|
||||
|
||||
class ConditionField(BaseModel):
|
||||
"""字段库中的一个条件字段。"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
name: str = Field(min_length=1, max_length=64, description="引擎字段名,如 dv_ratio / static.industry")
|
||||
label: str = Field(default="", max_length=64, description="中文名(下拉里展示)")
|
||||
description: str = Field(default="", max_length=500, description="含义 / 口径(含单位)")
|
||||
kind: str = Field(default="num", pattern="^(num|str)$")
|
||||
group_name: str = Field(default="行情", max_length=32)
|
||||
unit: str = Field(default="", max_length=16)
|
||||
source: str = Field(default="builtin", pattern="^(builtin|custom)$")
|
||||
enabled: bool = Field(default=True, description="False=从选择器隐藏(内置不可删除,只能停用)")
|
||||
sort_order: int = 100
|
||||
created_at: datetime | None = None
|
||||
updated_at: datetime | None = None
|
||||
|
||||
@computed_field # type: ignore[prop-decorator]
|
||||
@property
|
||||
def ops(self) -> list[str]:
|
||||
"""该字段可用的比较符(前端据此收窄下拉,不自己猜)。"""
|
||||
return list(OPS_BY_KIND.get(self.kind, OPS_NUM))
|
||||
@@ -0,0 +1,91 @@
|
||||
"""因子目录领域实体(M7.1:因子元数据 DB 化,v2 §11;2026-10 参数化)。
|
||||
|
||||
DB 是因子目录的契约源:**哪些因子存在**(含用户从模板派生的参数化实例)入库;
|
||||
**能不能算**仍由代码注册表(quant/factors.py)唯一决定 —— 登记但解析不出来的因子
|
||||
在 score/condition 里引用时抛 FactorError(不假装支持)。
|
||||
|
||||
参数化的读法:参数化实例的名字本身就是身份(`momentum(window=90,direction=...…)`),
|
||||
所以 template / params / param_specs / label / source / resolvable 都是**由名字解析出来的
|
||||
投影**,不落库。落库的只有 `enabled`(是否出现在下拉里)—— 这是人做的配置,
|
||||
不是引擎事实。好处:参数不可能出现「表里一套、键里一套」的分裂。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class FactorParam(BaseModel):
|
||||
"""一个可编辑参数的约束(与 quant/factors.ParamSpec 对齐,供界面渲染表单)。"""
|
||||
|
||||
name: str
|
||||
label: str = ""
|
||||
kind: str = "int" # "int" | "enum"
|
||||
default: Any = None
|
||||
minimum: int | None = None
|
||||
maximum: int | None = None
|
||||
choices: list[str] = Field(default_factory=list)
|
||||
note: str = ""
|
||||
|
||||
@classmethod
|
||||
def from_spec(cls, spec) -> FactorParam:
|
||||
return cls(
|
||||
name=spec.name,
|
||||
label=spec.label,
|
||||
kind=spec.kind,
|
||||
default=spec.default,
|
||||
minimum=spec.minimum,
|
||||
maximum=spec.maximum,
|
||||
choices=list(spec.choices),
|
||||
note=spec.note,
|
||||
)
|
||||
|
||||
|
||||
class FactorDefinition(BaseModel):
|
||||
name: str = Field(min_length=1, max_length=128)
|
||||
description: str = ""
|
||||
formula: str = ""
|
||||
brief: str = ""
|
||||
frequency: str = "daily"
|
||||
lookback: int = 20
|
||||
direction: str = Field(default="higher_is_better", pattern="^(higher_is_better|lower_is_better)$")
|
||||
requires: list[str] = Field(default_factory=list)
|
||||
version: str = "1"
|
||||
created_at: datetime | None = None
|
||||
# ---- 参数化(落库的只有 enabled;其余由 name 解析投影而来)----
|
||||
enabled: bool = True
|
||||
template: str = "" # 模板名,如 "momentum"
|
||||
params: dict[str, Any] = Field(default_factory=dict) # 冻结的参数取值
|
||||
param_specs: list[FactorParam] = Field(default_factory=list) # 可编辑参数与约束
|
||||
label: str = "" # 中文显示名(含参数)
|
||||
source: str = "builtin" # builtin(代码注册表实例)| custom(目录里的参数化实例)
|
||||
resolvable: bool = True # False = 登记了但引擎算不出来(历史手工登记行)
|
||||
|
||||
@classmethod
|
||||
def from_registry_def(cls, d) -> FactorDefinition:
|
||||
"""由 quant/factors.FactorDef(dataclass)构造目录实体(seed 用)。"""
|
||||
return cls.from_factor_def(d, enabled=True)
|
||||
|
||||
@classmethod
|
||||
def from_factor_def(cls, d, *, enabled: bool = True) -> FactorDefinition:
|
||||
"""由因子实例(内置或参数化)构造目录实体(seed / 新建用)。"""
|
||||
return cls(
|
||||
name=d.name,
|
||||
description=d.description,
|
||||
formula=d.formula,
|
||||
brief=d.brief,
|
||||
frequency=d.frequency,
|
||||
lookback=d.lookback,
|
||||
direction=d.direction,
|
||||
requires=list(d.requires),
|
||||
enabled=enabled,
|
||||
template=d.template,
|
||||
params=dict(d.params),
|
||||
param_specs=[FactorParam.from_spec(s) for s in d.param_specs],
|
||||
label=d.label,
|
||||
source=d.source,
|
||||
resolvable=True,
|
||||
)
|
||||
@@ -0,0 +1,21 @@
|
||||
"""指数及其历史成分(v3 §9/§30:Survivorship-free Universe 的数据基础)。
|
||||
|
||||
index_weight:指数在某交易日的成分快照(来自指数权重表,每期含当时成分与权重)。
|
||||
历史成分语义:as_of 某日的成分 = 该日(<=as_of 最近一期)快照中的股票 ——
|
||||
禁止用今天的成分回测过去(未来函数/幸存者偏差红线)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class IndexWeight(BaseModel):
|
||||
index_code: str = Field(pattern=r"^\d{6}\.(SH|SZ|CSI|CI)$", description="如 000300.SH / 000905.SH")
|
||||
index_name: str | None = None
|
||||
trade_date: date # 该快照对应交易日(成分时点)
|
||||
symbol: str = Field(pattern=r"^\d{6}\.(SH|SZ|BJ)$")
|
||||
weight: Decimal | None = None
|
||||
@@ -0,0 +1,209 @@
|
||||
"""市场数据领域实体(Phase 1)。
|
||||
|
||||
约定(AGENT.md §8/§9):
|
||||
- 行情时间用 trade_date;财务数据同时区分 report_date(报告期)与 announce_date(公告日)
|
||||
- 禁止以 report_date 作可见性依据 —— 只允许 announce_date 已过的数据进入研究
|
||||
- 复权一律通过独立 AdjustFactor 表达,不在此层偷偷改前/后复权口径
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
# 常见精度:价格 4 位小数;成交量(股) 2 位;金额(元) 2 位
|
||||
PRICE_PLACES = Decimal("0.0001")
|
||||
AMOUNT_PLACES = Decimal("0.01")
|
||||
|
||||
# ---- 研究面板数值列白名单(单一事实来源) ----
|
||||
# 说明:研究装配(quant 层)与持久化列裁剪(infrastructure 层)必须用同一套列名,
|
||||
# 否则会出现「因子需要某列、仓储却拒绝」的隐蔽不一致。故在此统一定义,两处引用。
|
||||
# 只影响数值列;symbol / trade_date 恒返回。
|
||||
DAILY_BAR_NUMERIC_FIELDS: tuple[str, ...] = ("open", "high", "low", "close", "volume", "amount")
|
||||
DAILY_BASIC_NUMERIC_FIELDS: tuple[str, ...] = (
|
||||
"close",
|
||||
"turnover_rate",
|
||||
"volume_ratio",
|
||||
"pe",
|
||||
"pe_ttm",
|
||||
"pb",
|
||||
"ps",
|
||||
"ps_ttm",
|
||||
"dv_ratio",
|
||||
"dv_ttm",
|
||||
"total_share",
|
||||
"float_share",
|
||||
"free_share",
|
||||
"total_mv",
|
||||
"circ_mv",
|
||||
)
|
||||
|
||||
|
||||
class Stock(BaseModel):
|
||||
"""A 股基础信息。symbol 统一为 Tushare 风格,如 600519.SH。"""
|
||||
|
||||
model_config = ConfigDict(str_strip_whitespace=True)
|
||||
|
||||
symbol: str = Field(pattern=r"^\d{6}\.(SH|SZ|BJ)$", description="如 600519.SH")
|
||||
name: str
|
||||
industry: str | None = None
|
||||
area: str | None = None
|
||||
market: str | None = Field(default=None, description="主板/创业板/科创板/北交所")
|
||||
exchange: str | None = None
|
||||
list_date: date
|
||||
delist_date: date | None = None
|
||||
status: str = Field(default="L", description="L 上市 / D 退市 / P 暂停")
|
||||
|
||||
|
||||
class TradingCalendar(BaseModel):
|
||||
"""交易日历。"""
|
||||
|
||||
calendar_date: date
|
||||
is_open: bool = True
|
||||
|
||||
|
||||
class DailyBar(BaseModel):
|
||||
"""日线。默认不复权(source=tushare, adjust=none)。
|
||||
|
||||
备用源兜底行会标记 source=sina、adjust=qfq(新浪返回前复权价)。
|
||||
字段统一、可区分、可追溯(AGENT §5.2/§8):研究侧应优先消费
|
||||
source=tushare 且 adjust=none 的行;新浪行仅在 Tushare 不可用期间作为兜底,
|
||||
Tushare 恢复后重跑 --resume 会按日覆盖回不复权口径。
|
||||
"""
|
||||
|
||||
symbol: str
|
||||
trade_date: date
|
||||
source: str = Field(default="tushare", description="tushare | sina")
|
||||
adjust: str = Field(default="none", description="none 不复权 | qfq 前复权")
|
||||
open: Decimal | None = None
|
||||
high: Decimal | None = None
|
||||
low: Decimal | None = None
|
||||
close: Decimal | None = None
|
||||
volume: Decimal | None = Field(default=None, description="成交量(股)")
|
||||
amount: Decimal | None = Field(default=None, description="成交额(元)")
|
||||
|
||||
@property
|
||||
def is_complete(self) -> bool:
|
||||
"""基础行情字段是否齐全(供校验器使用)。"""
|
||||
return all(
|
||||
v is not None
|
||||
for v in (self.open, self.high, self.low, self.close, self.volume, self.amount)
|
||||
)
|
||||
|
||||
|
||||
class AdjustFactor(BaseModel):
|
||||
"""复权因子。因子原始口径由数据源决定,必须与数据源文档一致地存取。"""
|
||||
|
||||
symbol: str
|
||||
trade_date: date
|
||||
factor: Decimal
|
||||
|
||||
|
||||
class DailyBasic(BaseModel):
|
||||
"""每日指标快照(Tushare daily_basic)—— 估值 / 股息率 / 市值。
|
||||
|
||||
时点性说明(防未来函数,AGENT.md §9):
|
||||
- 本表每一行都是**该交易日收盘后**即可得的横截面指标(dv_ratio 由
|
||||
「过去 12 个月现金分红 / 当日总市值」逐日重算),属时点值;
|
||||
- 研究侧一律按 trade_date <= as_of_date 取值,不存在未来信息。
|
||||
|
||||
列语义:
|
||||
- dv_ratio 股息率(%):近 12 个月现金分红 / 总市值 × 100
|
||||
- dv_ttm 股息率(TTM,%):滚动 12 个月口径
|
||||
- 两者均可能因**特别分红**出现畸高值(实测 600738 在 2020-01-02 为 37.2%),
|
||||
使用时建议配合上限过滤。
|
||||
"""
|
||||
|
||||
symbol: str
|
||||
trade_date: date
|
||||
close: Decimal | None = Field(default=None, description="当日收盘价(不复权,与 stock_daily 一致)")
|
||||
turnover_rate: Decimal | None = Field(default=None, description="换手率(%)")
|
||||
volume_ratio: Decimal | None = Field(default=None, description="量比")
|
||||
pe: Decimal | None = None
|
||||
pe_ttm: Decimal | None = None
|
||||
pb: Decimal | None = None
|
||||
ps: Decimal | None = None
|
||||
ps_ttm: Decimal | None = None
|
||||
dv_ratio: Decimal | None = Field(default=None, description="股息率(%),近 12 个月现金分红/总市值")
|
||||
dv_ttm: Decimal | None = Field(default=None, description="股息率 TTM(%)")
|
||||
total_share: Decimal | None = Field(default=None, description="总股本(万股)")
|
||||
float_share: Decimal | None = Field(default=None, description="流通股本(万股)")
|
||||
free_share: Decimal | None = Field(default=None, description="自由流通股本(万股)")
|
||||
total_mv: Decimal | None = Field(default=None, description="总市值(万元)")
|
||||
circ_mv: Decimal | None = Field(default=None, description="流通市值(万元)")
|
||||
source: str = Field(default="tushare", description="tushare | sina(新浪不提供本接口)")
|
||||
|
||||
|
||||
class StockNameHistory(BaseModel):
|
||||
"""股票名称变更历史(Tushare namechange)—— 时点 ST / 风险警示判定的依据。
|
||||
|
||||
为什么需要它(实测背景):`stock.name` 只是**最新名称快照**,用它做
|
||||
`universe.exclude_st` 会把「曾为高股息、后来才变 ST/退市」的标的在**整段历史**里
|
||||
都排除掉 —— 而那正是「股息陷阱」样本。实测 `600565.SH` 2020 年叫「迪马股份」
|
||||
(dv_ratio 7.9%,当年高股息候选),2024-05-06 才变「ST迪马」,用最新名称判定
|
||||
会在 2020 年就把它排除,导致高股息回测收益被高估(对照组实测约 3.70pp)。
|
||||
|
||||
一行 = 一个「名称生效区间」:
|
||||
- `name` 在 `[start_date, end_date]` 内有效(`end_date` 为空表示至今有效)
|
||||
- `change_reason` 为 tushare 口径:ST / *ST / 撤销ST / 撤销*ST / 从ST变为*ST / 其他
|
||||
- 时点取值按**生效区间**:`start_date <= as_of <= end_date`(实现口径,无前视:
|
||||
名称自 `start_date` 起即对市场可见)。`ann_date` 为公告日,仅作留痕/审计,
|
||||
**不参与**判定 —— 实测数据中 `ann_date` 恒早于或等于 `start_date`,
|
||||
若改用「公告即改名」会让 `ann_date` 为空的记录整体丢失。
|
||||
"""
|
||||
|
||||
symbol: str
|
||||
name: str
|
||||
start_date: date
|
||||
end_date: date | None = None
|
||||
ann_date: date | None = None
|
||||
change_reason: str | None = None
|
||||
source: str = "tushare"
|
||||
|
||||
@property
|
||||
def is_risk_warned(self) -> bool:
|
||||
"""该区间名称是否含风险警示(ST / *ST)。"""
|
||||
return "ST" in self.name.upper()
|
||||
|
||||
|
||||
class FinancialIndicator(BaseModel):
|
||||
"""核心财务指标(快照)。
|
||||
|
||||
可见性红线:研究侧查询一律按 announce_date <= as_of_date 过滤,
|
||||
report_date 只表示报告所属期间,不代表公开时间。
|
||||
|
||||
source 标记数据来源:tushare(首选,字段全)| sina(兜底,字段
|
||||
可能不全——新浪关键指标只含 eps/roe/gross_margin 等少数项)。
|
||||
新浪兜底行只在「该股票本地历史与新浪重叠部分两边一致」通过校验后
|
||||
才导入(见 application/services/data_sync.py),且只补本地缺失键。
|
||||
研究侧对同一报告期应优先消费 source=tushare 的行。
|
||||
"""
|
||||
|
||||
symbol: str
|
||||
report_date: date
|
||||
announce_date: date
|
||||
source: str = Field(default="tushare", description="tushare | sina")
|
||||
eps: Decimal | None = None
|
||||
roe: Decimal | None = None
|
||||
total_revenue: Decimal | None = None
|
||||
net_profit: Decimal | None = None
|
||||
gross_margin: Decimal | None = None
|
||||
|
||||
def announced_by(self, as_of_date: date) -> bool:
|
||||
"""as_of_date(含当日)是否已可见。防未来函数的核心判断。"""
|
||||
return self.announce_date <= as_of_date
|
||||
|
||||
|
||||
class SyncLog(BaseModel):
|
||||
"""数据拉取审计记录(AGENT.md §7:来源必须可追踪,禁止静默切换)。"""
|
||||
|
||||
source: str
|
||||
api: str
|
||||
request_time: datetime = Field(default_factory=datetime.utcnow)
|
||||
success: bool
|
||||
failure_reason: str | None = None
|
||||
row_count: int = 0
|
||||
data_start: date | None = None
|
||||
data_end: date | None = None
|
||||
@@ -0,0 +1,34 @@
|
||||
"""Bar Replay(v3 §20.6,第二阶段 MVP)领域实体。
|
||||
|
||||
线性重放:对给定选股查询与信号规则,在交易日序列上逐日以「当日为止的数据」执行
|
||||
(as_of 语义),输出每日时间线(意图 Top + 信号 + 计数),用于核对:
|
||||
- 未来函数:每日计算只用 <= as_of 数据(与静态研究同一引擎)
|
||||
- 一致性:重放某日结果 == 该日独立 select/signal 结果(亦 == 回测该调仓日意图)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.domain.entities.signal import SignalEvent
|
||||
|
||||
|
||||
class ReplayTop(BaseModel):
|
||||
symbol: str
|
||||
score: float
|
||||
|
||||
|
||||
class ReplayDay(BaseModel):
|
||||
as_of: date
|
||||
top: list[ReplayTop] = Field(default_factory=list, description="意图排名前 N(score 降序)")
|
||||
events: list[SignalEvent] = Field(default_factory=list, description="当日信号(<=max_output_rank)")
|
||||
counts: dict[str, int] = Field(default_factory=dict, description="BUY/WATCH/SELL 计数")
|
||||
|
||||
|
||||
class ReplayResult(BaseModel):
|
||||
start: date
|
||||
end: date
|
||||
days: list[ReplayDay] = Field(default_factory=list)
|
||||
top_n: int = 5
|
||||
@@ -0,0 +1,555 @@
|
||||
"""研究领域对象:Research Specification、标准化研究结果。
|
||||
|
||||
原则(AGENT.md §16/§21/§24、ARCHITECTURE §14):
|
||||
- 前端 / Agent / 后端统一经 Research Specification 描述任务,禁止直接拼引擎配置
|
||||
- 回测结果一律标准化为 BacktestResult;未建模的成本/市场约束显式列在
|
||||
unimplemented,禁止默认「无成本 / 永远可成交」假设
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||
|
||||
# ---------- Research Specification ----------
|
||||
|
||||
|
||||
class UniverseSpec(BaseModel):
|
||||
"""股票池口径。MVP:市场 + 过滤条件;指数成分等 Phase 3 扩展。
|
||||
|
||||
symbols 白名单:非空时仅这些股票参与(再叠加其余过滤);供自选池/测试使用。
|
||||
market 目前为预留字段(stock.market 存储主板/创业板/科创板等中文枚举,过滤未启用)。
|
||||
"""
|
||||
|
||||
market: str = Field(default="CN_A", description="CN_A / CN_B / ...(预留)")
|
||||
exclude_st: bool = True
|
||||
exclude_suspended: bool = True
|
||||
min_listing_days: int = Field(default=250, ge=0, description="上市至少 N 个自然日")
|
||||
index_code: str | None = Field(
|
||||
default=None, description="指数成分过滤(如 000300.SH):按 as_of 当日历史成分(v3 §9)"
|
||||
)
|
||||
symbols: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="白名单(可选):非空时仅这些 symbol 参与选股/回测",
|
||||
)
|
||||
|
||||
|
||||
class FactorSpec(BaseModel):
|
||||
"""引用一个已注册因子并给定权重。"""
|
||||
|
||||
name: str
|
||||
weight: float = Field(default=1.0, gt=0)
|
||||
|
||||
|
||||
class ConditionSpec(BaseModel):
|
||||
"""结构化选股条件(回测与选股共用)。
|
||||
|
||||
字段域:
|
||||
- static.*:股票基础字段(industry / market / area / exchange / status…)
|
||||
- 行情/技术字段:close / ma20 / ma60 / volume 及全部已注册因子名(momentum_60 等),
|
||||
以及每日指标列(dv_ratio / dv_ttm / pe / pb / total_mv …)
|
||||
- fundamental.*:财务字段(eps / roe / total_revenue / net_profit / gross_margin),
|
||||
仅取 announce_date <= as_of 的最新已公告值(防未来函数)
|
||||
|
||||
右操作数取 value(字面量)或 ref(另一字段名),二者二选一。
|
||||
|
||||
字段域的事实来源:`quant/condition_fields.py`(字段库注册表,含中文名与口径)。
|
||||
`/api/condition-fields`(前端下拉)、该注册表与引擎求值共用同一份定义,
|
||||
避免「前端列一个、引擎算另一个」的漂移。
|
||||
|
||||
定义位置说明:本模型被 ResearchSpec(回测)与 SelectionQuery(选股)共用,
|
||||
故落在 research.py(被 selection.py 依赖的低层模块),selection.py 再 re-export,
|
||||
避免循环导入。
|
||||
"""
|
||||
|
||||
field: str
|
||||
op: str = Field(pattern="^(gt|gte|lt|lte|eq|ne|in|not_in)$")
|
||||
value: float | int | str | list | None = None
|
||||
ref: str | None = None # 与另一字段比较(如 close vs ma60)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _require_operand(self) -> ConditionSpec:
|
||||
if self.value is None and self.ref is None:
|
||||
raise ValueError("value 与 ref 必须提供一个")
|
||||
if self.value is not None and self.ref is not None:
|
||||
raise ValueError("value 与 ref 只能提供一个")
|
||||
if self.op in ("in", "not_in") and not isinstance(self.value, list):
|
||||
raise ValueError("in/not_in 的 value 必须是列表")
|
||||
return self
|
||||
|
||||
|
||||
class SelectionSpec(BaseModel):
|
||||
"""选股方式(两级截断)。
|
||||
|
||||
口径(用户案例「选 n 只 → 持仓前 x 只」):
|
||||
- `top_n` = n:**候选池**大小。universe ∩ conditions 过滤后,按复合因子分降序取前 n
|
||||
只 → 这就是「择股条件选出来的股数」(写入 selection_history)。
|
||||
- `hold_top_x` = x:**实际持仓数**,取候选池前 x 只等权。x 必须 ≤ n;
|
||||
另受「池内实际可买股票数」约束(过滤/缺数据会让实际池子小于 n)。
|
||||
None → 等于 top_n(此时与旧行为一致:选出多少就持多少)。
|
||||
|
||||
`allow_substitute` 与 `defer_buy` 决定「买不进」时的处理(两者互斥,只能选一个):
|
||||
- `allow_substitute=True`(**默认**,保持历史语义不变):从 n 名**之外**继续往下找
|
||||
可买标的补足 x 只 —— 引擎既有行为,见 v3 §20.3 的 Signal↔Fill 测试。
|
||||
- `defer_buy=True`(本项目「只买选出来的前 x 只」口径,推荐显式开启):
|
||||
**不替补**,把这只股票的买单**顺延到之后第一个可成交的交易日**(涨停/停牌解除后
|
||||
按当日收盘价买入);到下一次调仓仍未成交则作废,未投入资金留作现金。
|
||||
- 两者都 False:意图被拒后直接放弃,资金留现金(不替补也不顺延)。
|
||||
|
||||
默认值刻意保持「向后兼容」:既有 Strategy / Experiment 的语义不因本次扩展而静默改变
|
||||
(AGENT.md §35)。高股息案例在前端与 spec 中显式设置 defer_buy=True。
|
||||
"""
|
||||
|
||||
top_n: int = Field(default=30, ge=1, le=1000, description="n:候选池大小")
|
||||
hold_top_x: int | None = Field(
|
||||
default=None, ge=1, le=1000, description="x:实际持仓数;None → = top_n"
|
||||
)
|
||||
allow_substitute: bool = Field(
|
||||
default=True,
|
||||
description="True(默认,历史语义):从 n 名之外替补补足;False:不引入计划外标的",
|
||||
)
|
||||
defer_buy: bool = Field(
|
||||
default=False,
|
||||
description="True:买不进(涨停/停牌)时顺延到之后首个可成交日的收盘买入",
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_x_le_n(self) -> SelectionSpec:
|
||||
if self.hold_top_x is not None and self.hold_top_x > self.top_n:
|
||||
raise ValueError(
|
||||
f"hold_top_x(持仓 x={self.hold_top_x})不能大于 top_n(候选池 n={self.top_n})"
|
||||
)
|
||||
if self.allow_substitute and self.defer_buy:
|
||||
raise ValueError(
|
||||
"allow_substitute=True(往下替补)与 defer_buy=True(顺延买入)语义互斥,只能选一个"
|
||||
)
|
||||
return self
|
||||
|
||||
@property
|
||||
def x(self) -> int:
|
||||
"""实际持仓目标数(未显式给 x 时等于 n)。"""
|
||||
return self.hold_top_x if self.hold_top_x is not None else self.top_n
|
||||
|
||||
|
||||
class PortfolioSpec(BaseModel):
|
||||
"""组合构建(v2 §16)。MVP:等权;单股/行业上限等约束字段预留,
|
||||
未建模约束在回测结果 unimplemented 中如实标注(禁止假装支持)。
|
||||
"""
|
||||
|
||||
weighting: str = Field(default="equal", pattern="^(equal)$")
|
||||
max_position_pct: float | None = Field(
|
||||
default=None, gt=0, le=1, description="单股最大权重(预留,未建模)"
|
||||
)
|
||||
max_industry_weight_pct: float | None = Field(
|
||||
default=None, gt=0, le=1, description="行业最大权重(预留,未建模)"
|
||||
)
|
||||
|
||||
|
||||
class CostSpec(BaseModel):
|
||||
"""交易成本模型。
|
||||
|
||||
买:commission(≥ min_commission)+ slippage
|
||||
卖:commission(≥ min_commission)+ stamp_tax + slippage
|
||||
|
||||
`min_commission` 为**单笔最低佣金**(A 股常见 5 元)。默认 0.0 = 不启用,
|
||||
以保持既有回测数值不变(AGENT.md §35);高股息等实盘贴近场景建议显式设 5.0。
|
||||
注意:最低佣金对**小额单**影响显著,x 越多、单笔越小,成本占比越高。
|
||||
"""
|
||||
|
||||
commission_rate: float = Field(default=0.0003, ge=0, le=0.01)
|
||||
stamp_tax_rate: float = Field(default=0.0005, ge=0, le=0.01)
|
||||
slippage_rate: float = Field(default=0.001, ge=0, le=0.05)
|
||||
min_commission: float = Field(
|
||||
default=0.0, ge=0, le=100.0, description="单笔最低佣金(元);0 = 不启用"
|
||||
)
|
||||
benchmark: str = Field(default="000300.SH", description="对照基准指数代码")
|
||||
|
||||
|
||||
class ResearchSpec(BaseModel):
|
||||
"""一次研究的完整描述。type 决定执行路径。"""
|
||||
|
||||
type: str = Field(default="backtest", pattern="^(factor_test|backtest)$")
|
||||
universe: UniverseSpec = UniverseSpec()
|
||||
price_adjustment: str = Field(
|
||||
default="none", pattern="^(none|qfq|hfq)$",
|
||||
description=(
|
||||
"研究行情口径:none 不复权(默认)/ qfq 前复权 / hfq 后复权。"
|
||||
"qfq/hfq 基于 adjust_factor 折算(v3 §20.5);结果与 config_snapshot 中显式记录。"
|
||||
),
|
||||
)
|
||||
factors: list[FactorSpec] = Field(min_length=1)
|
||||
conditions: list[ConditionSpec] = Field(
|
||||
default_factory=list,
|
||||
description=(
|
||||
"选股过滤条件(AND,可选):universe 之后、因子排序之前执行。"
|
||||
"字段域同 SelectionQuery.conditions(static.* / 行情列 / 已注册因子 / "
|
||||
"fundamental.*),回测与 /api/selections 共用同一求值器(v2 §25 一致性)。"
|
||||
),
|
||||
)
|
||||
selection: SelectionSpec = SelectionSpec()
|
||||
rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$")
|
||||
selection_interval_months: int | None = Field(
|
||||
default=None, ge=1, le=60,
|
||||
description=(
|
||||
"m:择股间隔(月)。None → 每次调仓都重新择股(等价于 rebalance 频率)。"
|
||||
"择股日 = 起始月锚定,月序号 % m == 0 的月份的首个交易日。"
|
||||
),
|
||||
)
|
||||
rebalance_interval_months: int | None = Field(
|
||||
default=None, ge=1, le=60,
|
||||
description=(
|
||||
"y:调仓间隔(月)。None → 等于 selection_interval_months(未给则按 rebalance 频率)。"
|
||||
"y < m 时池子在下一次择股前保持不变(结果中会标注池子陈旧)。"
|
||||
),
|
||||
)
|
||||
period: tuple[date, date]
|
||||
costs: CostSpec = CostSpec()
|
||||
portfolio: PortfolioSpec = PortfolioSpec()
|
||||
initial_capital: float = Field(default=1_000_000.0, gt=0)
|
||||
|
||||
@field_validator("period")
|
||||
@classmethod
|
||||
def _period_ordered(cls, period: tuple[date, date]) -> tuple[date, date]:
|
||||
if period[0] >= period[1]:
|
||||
raise ValueError("period 必须满足 start < end")
|
||||
return period
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _no_duplicate_factors(self) -> ResearchSpec:
|
||||
names = [f.name for f in self.factors]
|
||||
if len(set(names)) != len(names):
|
||||
raise ValueError("factors 存在重复因子名")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_intervals(self) -> ResearchSpec:
|
||||
m = self.selection_interval_months
|
||||
y = self.rebalance_interval_months
|
||||
# 未给 m 却给了 y:语义不完整(y 无锚点可依)→ 明确拒绝而非猜
|
||||
if m is None and y is not None and y != 1:
|
||||
raise ValueError(
|
||||
"只给了 rebalance_interval_months(y) 而没给 selection_interval_months(m):"
|
||||
"请同时给出 m,否则无法确定择股日集合"
|
||||
)
|
||||
if m is not None and y is not None and y < m and y != 1:
|
||||
# 允许但不静默:池子会在多个调仓日复用(陈旧),交由结果 unimplemented 标注
|
||||
return self
|
||||
return self
|
||||
|
||||
@property
|
||||
def effective_selection_months(self) -> int | None:
|
||||
"""实际择股间隔(月);None 表示「每次调仓都择股」。"""
|
||||
return self.selection_interval_months
|
||||
|
||||
@property
|
||||
def effective_rebalance_months(self) -> int | None:
|
||||
"""实际调仓间隔(月);None 表示按 rebalance 频率(周/月)。"""
|
||||
if self.rebalance_interval_months is not None:
|
||||
return self.rebalance_interval_months
|
||||
return self.selection_interval_months
|
||||
|
||||
|
||||
# ---------- 回测结果 ----------
|
||||
|
||||
|
||||
class CurvePoint(BaseModel):
|
||||
date: date
|
||||
value: float
|
||||
|
||||
|
||||
class MonthlyReturn(BaseModel):
|
||||
year: int
|
||||
month: int
|
||||
return_pct: float # 百分数,如 3.2 表示 +3.2%
|
||||
|
||||
|
||||
class YearlyReturn(BaseModel):
|
||||
year: int
|
||||
return_pct: float
|
||||
|
||||
|
||||
class BacktestSummary(BaseModel):
|
||||
start: date
|
||||
end: date
|
||||
initial_capital: float
|
||||
final_equity: float
|
||||
total_return_pct: float
|
||||
annual_return_pct: float
|
||||
sharpe: float
|
||||
max_drawdown_pct: float
|
||||
volatility_pct: float
|
||||
win_rate_pct: float
|
||||
total_trades: int
|
||||
avg_turnover_pct: float
|
||||
benchmark_return_pct: float | None = None
|
||||
|
||||
|
||||
class TradeReason(BaseModel):
|
||||
"""一次交易意图 / 成交的**结构化理由**:用当时的真实数字解释「为什么买 / 为什么卖」。
|
||||
|
||||
为什么不让前端自己推:界面上出现的每个数字(排名、综合分、因子值、持有天数)
|
||||
都必须来自引擎当时的计算,否则就是「看着像真的」。理由因此分三层:
|
||||
|
||||
- `code`:机器可判定的原因分类(封闭取值,见 `quant/trade_reasons.py`)。
|
||||
前端据此筛选 / 上色,不去解析文案;
|
||||
- `text`:给人读的一句话(自带关键数字,可单独展示);
|
||||
- `data`:当时真实数值(`rank` / `total` / `score` / `top_n` / `hold_days` /
|
||||
`factors`(各因子当时的原始值)/ `budget` …)。前端只展示,不推算。
|
||||
"""
|
||||
|
||||
code: str
|
||||
text: str
|
||||
data: dict = Field(default_factory=dict)
|
||||
|
||||
|
||||
class Trade(BaseModel):
|
||||
entry_date: date
|
||||
exit_date: date
|
||||
symbol: str
|
||||
name: str | None = Field(default=None, description="股票名称(展示用)")
|
||||
entry_price: float
|
||||
exit_price: float
|
||||
return_pct: float
|
||||
entry_reason: TradeReason | None = Field(
|
||||
default=None, description="买入理由(建仓当日引擎给出的结构化理由)"
|
||||
)
|
||||
exit_reason: TradeReason | None = Field(
|
||||
default=None, description="卖出理由(了结当日引擎给出的结构化理由)"
|
||||
)
|
||||
|
||||
|
||||
class Position(BaseModel):
|
||||
date: date
|
||||
symbol: str
|
||||
name: str | None = Field(default=None, description="股票名称(展示用)")
|
||||
weight: float
|
||||
|
||||
|
||||
class RankedPick(BaseModel):
|
||||
"""调仓日选股意图候选(与 select(as_of) 同源;v3 §22.3 selection_history)。"""
|
||||
|
||||
date: date
|
||||
symbol: str
|
||||
name: str | None = Field(default=None, description="股票名称(展示用)")
|
||||
rank: int
|
||||
score: float
|
||||
|
||||
|
||||
class ActionRecord(BaseModel):
|
||||
"""一次交易意图(Signal)及其成交结果(Fill)—— v3 §20.3 Signal↔Fill 区分。
|
||||
|
||||
signal=BUY/SELL(策略意图);filled=是否实际成交;reject_reason 给出未成交原因
|
||||
(涨停/跌停/无价/现金不足等)。fills = [a for a in signal_history if a.filled]。
|
||||
|
||||
`reason` 是**数据化**的为什么:`reject_reason` 只说「没成交」(执行层),
|
||||
`reason` 同时覆盖成交与未成交(策略层 + 执行层),并带上当时的排名 / 综合分 /
|
||||
各因子原始值,前端「买卖说明」直接用,不再二次推断。
|
||||
|
||||
`name` 为展示增强字段:由服务层按股票池统一回填(未命中则为 None),
|
||||
引擎自身不感知名称 —— 引擎只处理 symbol,保持纯行情计算职责。
|
||||
"""
|
||||
|
||||
date: date
|
||||
symbol: str
|
||||
name: str | None = Field(default=None, description="股票名称(展示用)")
|
||||
signal: str = Field(pattern="^(BUY|SELL)$")
|
||||
filled: bool
|
||||
reject_reason: str | None = None
|
||||
price: float | None = Field(default=None, description="成交价(fill)或意图参考价")
|
||||
reason: TradeReason | None = Field(
|
||||
default=None, description="结构化理由(成交与未成交都有;旧归档为 null)"
|
||||
)
|
||||
|
||||
|
||||
class SymbolCurve(BaseModel):
|
||||
"""个股收益率趋势曲线 + 该股买卖点标注(回测结果可视化用)。
|
||||
|
||||
`points[].value` 语义:该股**持仓期间**的累计收益率(%,以建仓日收盘为 0% 基准,
|
||||
按日复利)。只在该股被持有的交易日落点(未持有期间不落点,以压缩结果体积);
|
||||
建仓当日会补一个基准点,保证买卖点标注总能在曲线上取到数值。多段持仓以累计值
|
||||
连乘衔接,读图时以 marks 中的 BUY/SELL 区分各段持仓区间。
|
||||
|
||||
`marks` 为该股实际成交(BUY/SELL fill)的日期与价格,与 `signal_history`
|
||||
中 filled=True 的记录一致(v3 §20.3 的成交口径)。
|
||||
"""
|
||||
|
||||
symbol: str
|
||||
name: str | None = Field(default=None, description="股票名称(展示用)")
|
||||
points: list[CurvePoint] = Field(default_factory=list)
|
||||
marks: list[ActionRecord] = Field(default_factory=list)
|
||||
final_return_pct: float = Field(
|
||||
default=0.0, description="该股持仓期累计收益率(%,多段持仓连乘)"
|
||||
)
|
||||
|
||||
|
||||
class FactorCurve(BaseModel):
|
||||
"""单个因子在回测期内的时间序列(**持仓组合加权平均原始值**)。
|
||||
|
||||
口径必须写死,否则读图会读反:
|
||||
|
||||
- 值为该因子在**当日持仓股票**上的权重加权平均(权重 = 该股当日市值 / 组合权益),
|
||||
是**原始值**:不做 z-score、不按方向取负 —— 图上看到的就是因子本身;
|
||||
- 空仓日不落点(不插值、不用 0 假填充),曲线中间会出现空档;
|
||||
- `direction` 一并归档:低为好的因子,曲线升高不等于「更好」;
|
||||
- `unit` 是代码注册表里的事实(`%` / `倍数` / `小数`),用于坐标轴与提示文案。
|
||||
"""
|
||||
|
||||
name: str
|
||||
label: str
|
||||
direction: str = "higher_is_better"
|
||||
unit: str | None = None
|
||||
points: list[CurvePoint] = Field(default_factory=list)
|
||||
|
||||
|
||||
class BacktestResult(BaseModel):
|
||||
"""标准化回测结果(ARCHITECTURE §14)。前端只依赖该结构。"""
|
||||
|
||||
summary: BacktestSummary
|
||||
equity_curve: list[CurvePoint]
|
||||
drawdown: list[CurvePoint]
|
||||
monthly_returns: list[MonthlyReturn]
|
||||
yearly_returns: list[YearlyReturn]
|
||||
positions: list[Position]
|
||||
trades: list[Trade]
|
||||
selection_history: list[RankedPick] = Field(
|
||||
default_factory=list, description="各调仓日选股意图候选(同 select(as_of))"
|
||||
)
|
||||
signal_history: list[ActionRecord] = Field(
|
||||
default_factory=list, description="交易意图与是否成交(v3 §20.3)"
|
||||
)
|
||||
fills: list[ActionRecord] = Field(
|
||||
default_factory=list, description="实际成交(signal_history 中 filled=True 的子集)"
|
||||
)
|
||||
symbol_curves: list[SymbolCurve] = Field(
|
||||
default_factory=list,
|
||||
description="个股收益率曲线 + 买卖点标注(按期末收益绝对值降序,体积可控)",
|
||||
)
|
||||
factor_curves: list[FactorCurve] = Field(
|
||||
default_factory=list,
|
||||
description=(
|
||||
"策略用到的每个因子的时间序列(持仓加权平均原始值):"
|
||||
"用来解释「买卖依据的那个因子在各时点是什么水平」"
|
||||
),
|
||||
)
|
||||
turnover_pct: float
|
||||
unimplemented: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="本结果中未建模的约束(AGENT §24:必须显式标注,禁止假装支持)",
|
||||
)
|
||||
config_snapshot: dict = Field(default_factory=dict, description="复现用完整配置快照")
|
||||
archive_meta: dict = Field(
|
||||
default_factory=dict,
|
||||
description=(
|
||||
"归档元数据(由 experiment_archive 在落库时写入):curves_stored / "
|
||||
"curves_total / truncated / budget_chars / budget_bytes / result_chars / "
|
||||
"result_bytes。用于说明归档是否因体积预算被裁剪(AGENT §24 不静默)"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
SymbolCurve.model_rebuild()
|
||||
|
||||
|
||||
# ---------- 因子测试结果 ----------
|
||||
|
||||
|
||||
class QuantileReturn(BaseModel):
|
||||
"""分层收益:按因子值升序分 N 层后各层等权组合的区间收益。"""
|
||||
|
||||
quantile: int
|
||||
return_pct: float
|
||||
|
||||
|
||||
class FactorTestReport(BaseModel):
|
||||
factor_name: str
|
||||
ic_mean: float
|
||||
icir: float
|
||||
rank_ic_mean: float
|
||||
positive_ratio_pct: float
|
||||
quantile_returns: list[QuantileReturn]
|
||||
spread_quantile: int | None = Field(
|
||||
default=None, description="分层价差 = 最高层收益 - 最低层收益(若多头/空头语义适用)"
|
||||
)
|
||||
sample_days: int
|
||||
unimplemented: list[str] = Field(default_factory=list)
|
||||
config_snapshot: dict = Field(default_factory=dict)
|
||||
|
||||
|
||||
# ---------- 因子相关性 / 暴露分析(C1,v3 §12) ----------
|
||||
|
||||
|
||||
class FactorCorrelationReport(BaseModel):
|
||||
"""多因子两两相关(横截面相关逐日均值;v3 §12 冗余剔除前置)。"""
|
||||
|
||||
factors: list[str]
|
||||
corr_matrix: dict[str, dict[str, float]] = Field(
|
||||
default_factory=dict, description="{f1: {f2: spearman 相关系数}}(对角线=1)"
|
||||
)
|
||||
sample_days: int = 0
|
||||
sample_min_symbols: int = 0
|
||||
unimplemented: list[str] = Field(default_factory=list)
|
||||
config_snapshot: dict = Field(default_factory=dict)
|
||||
|
||||
|
||||
# ---------- 异步 Job 与 Experiment(Phase 4) ----------
|
||||
|
||||
|
||||
class JobStatus(str):
|
||||
"""统一状态机(AGENT.md §20):queued→running→(success|failed|cancelled)。"""
|
||||
|
||||
QUEUED = "queued"
|
||||
RUNNING = "running"
|
||||
SUCCESS = "success"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
class JobRecord(BaseModel):
|
||||
"""一次异步研究任务。spec/result 以 JSON 文本存储(保持 Schema 演进自由)。"""
|
||||
|
||||
id: str
|
||||
kind: str # factor_test | backtest
|
||||
status: str = JobStatus.QUEUED
|
||||
stage: str | None = None
|
||||
spec_json: str
|
||||
error: str | None = None
|
||||
result_json: str | None = None
|
||||
experiment_id: str | None = None
|
||||
created_at: datetime | None = None
|
||||
started_at: datetime | None = None
|
||||
finished_at: datetime | None = None
|
||||
|
||||
|
||||
class ExperimentRecord(BaseModel):
|
||||
"""一次研究的可复现存档(AGENT.md §21)。"""
|
||||
|
||||
id: str
|
||||
kind: str # factor_test | backtest
|
||||
spec_json: str
|
||||
result_json: str
|
||||
summary_text: str | None = None # 便于列表展示的摘要(如 total_return_pct)
|
||||
code_version: str | None = None # git commit / 代码指纹
|
||||
data_version: str | None = None
|
||||
job_id: str | None = None
|
||||
created_at: datetime | None = None
|
||||
|
||||
|
||||
class ExperimentSummary(BaseModel):
|
||||
"""归档列表项(**不含 result_json**)。
|
||||
|
||||
列表接口一次可能返回上百条归档,而 `result_json` 是 MEDIUMTEXT(完整存档后
|
||||
单条可达数 MB):为避免把上百 MB 拉进内存,仓储的列表查询只取元数据列,
|
||||
`result_bytes` 由 SQL 的字符长度函数(MySQL CHAR_LENGTH / SQLite length)
|
||||
在库侧算出,不取回大字段本身。
|
||||
"""
|
||||
|
||||
id: str
|
||||
kind: str
|
||||
spec_json: str
|
||||
summary_text: str | None = None
|
||||
code_version: str | None = None
|
||||
data_version: str | None = None
|
||||
job_id: str | None = None
|
||||
created_at: datetime | None = None
|
||||
result_bytes: int = 0 # 归档 JSON 的字符数(SQL 侧计算,不拉大字段)
|
||||
@@ -0,0 +1,127 @@
|
||||
"""选股系统领域对象(ARCHITECTURE_v2 §14 Selection Engine)。
|
||||
|
||||
回答两个核心问题(v2 §8):
|
||||
- 「某历史日(as_of)为什么选出这些股票?」→ SelectionResult 带 factor_values / selection_reason
|
||||
- 「当前(as_of)有哪些股票满足策略?」→ 同一条查询对当前日期执行
|
||||
|
||||
设计:
|
||||
- SelectionQuery = v2 §14.2 的 Selection 输入(universe 范围 + 评分因子 + TopN 截断 + as_of)。
|
||||
- method=score:按因子加权复合分取 TopN(复用现有 9 个内置因子);
|
||||
method=condition:结构化条件选股(M6.2 加入 ConditionSpec)。
|
||||
- 结果不落库由本实体负责(落库表在 M6.3);本实体是前后端/Agent 的统一 DTO(v2 §21.1)。
|
||||
- 所有查询天然带 as_of 语义:只允许使用 <= as_of 的数据(v2 §9 防未来函数)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||
|
||||
from app.domain.entities.research import ConditionSpec, FactorSpec, UniverseSpec # noqa: F401
|
||||
|
||||
|
||||
class SelectionQuery(BaseModel):
|
||||
"""一次选股查询(v2 §14.2 Selection 输入)。"""
|
||||
|
||||
universe: UniverseSpec = UniverseSpec()
|
||||
price_adjustment: str = Field(
|
||||
default="none", pattern="^(none|qfq|hfq)$",
|
||||
description=(
|
||||
"行情口径:none 不复权(默认)/ qfq 前复权 / hfq 后复权(按 adjust_factor 折算)。"
|
||||
"与回测(ResearchSpec.price_adjustment)同一口径域,结果 config_snapshot 中显式记录"
|
||||
),
|
||||
)
|
||||
# 研究时点:None → 引擎用 <= 今天最近可用交易日;显式给历史日期即做历史选股
|
||||
as_of: date | None = Field(
|
||||
default=None, description="选股时点;历史回测/解释用具体日期,当前选股可留空"
|
||||
)
|
||||
method: str = Field(default="score", pattern="^(score|condition)$")
|
||||
# method=score:因子 + 权重(至少 1 个;方向由因子元数据决定)
|
||||
factors: list[FactorSpec] = Field(default_factory=list)
|
||||
# method=condition:结构化条件(M6.2 引入 ConditionSpec 后启用)
|
||||
conditions: list[ConditionSpec] = Field(default_factory=list)
|
||||
# 截断:top_n(绝对数量)与 top_pct(占可评分股票比例)二选一;可选 min_score 下限
|
||||
top_n: int | None = Field(default=None, ge=1, le=2000)
|
||||
top_pct: float | None = Field(default=None, gt=0, le=1)
|
||||
min_score: float | None = None
|
||||
# 因子预热窗口(自然日):覆盖 lookback 前导数据,Lookback 放大时需同步加大
|
||||
warmup_days: int = Field(default=300, ge=0)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_method_args(self) -> SelectionQuery:
|
||||
if self.method == "score":
|
||||
if not self.factors:
|
||||
raise ValueError("method=score 需要至少一个 factors")
|
||||
if self.top_n is None and self.top_pct is None:
|
||||
raise ValueError("method=score 需要 top_n 与 top_pct 至少提供一个")
|
||||
if self.method == "condition" and not self.conditions:
|
||||
raise ValueError("method=condition 需要至少一个 conditions")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _no_duplicate_factors(self) -> SelectionQuery:
|
||||
names = [f.name for f in self.factors]
|
||||
if len(set(names)) != len(names):
|
||||
raise ValueError("factors 存在重复因子名")
|
||||
return self
|
||||
|
||||
|
||||
# ConditionSpec 的实体定义已上移到 research.py(ResearchSpec 回测条件与
|
||||
# SelectionQuery 选股条件共用同一模型)。此处 re-export:`selection.ConditionSpec`
|
||||
# 就是 research.ConditionSpec 这**同一个类对象**,因此两侧 isinstance 判断一致。
|
||||
# (见文件顶部 import 处的引用)
|
||||
|
||||
|
||||
class SelectionCandidate(BaseModel):
|
||||
"""单只候选股(v2 §14.3/§21.1)。
|
||||
|
||||
`name` 为展示增强字段(可选默认 None):由 SelectionService 用已装配的股票池
|
||||
统一回填,引擎/选股算法本身不感知名称 —— 避免在每个出口零散 join。
|
||||
"""
|
||||
|
||||
symbol: str
|
||||
name: str | None = Field(default=None, description="股票名称(展示用)")
|
||||
rank: int
|
||||
score: float
|
||||
factor_values: dict[str, float] = Field(default_factory=dict)
|
||||
filter_status: list[str] = Field(default_factory=list, description="各条件通过/未通过")
|
||||
selection_reason: list[str] = Field(default_factory=list, description="为什么选它(可解释)")
|
||||
|
||||
|
||||
class SelectionStatistics(BaseModel):
|
||||
universe_size: int = 0 # 股票池过滤后数量
|
||||
evaluated: int = 0 # 有有效分数的股票数量
|
||||
selected: int = 0 # 最终选出数量
|
||||
|
||||
|
||||
class SelectionResult(BaseModel):
|
||||
"""选股结果(v2 §21.1)。前端 / Agent 只依赖该结构。"""
|
||||
|
||||
as_of_date: date
|
||||
method: str
|
||||
statistics: SelectionStatistics
|
||||
candidates: list[SelectionCandidate] = Field(default_factory=list)
|
||||
unimplemented: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="本结果中未建模的约束(如 exclude_suspended 依赖停牌数据未实现)",
|
||||
)
|
||||
config_snapshot: dict = Field(default_factory=dict, description="复现用查询快照")
|
||||
|
||||
@field_validator("candidates")
|
||||
@classmethod
|
||||
def _rank_sorted(cls, candidates: list[SelectionCandidate]) -> list[SelectionCandidate]:
|
||||
return sorted(candidates, key=lambda c: c.rank)
|
||||
|
||||
|
||||
SelectionQuery.model_rebuild()
|
||||
|
||||
class SelectionMeta(BaseModel):
|
||||
"""选股运行元数据(列表/历史查询用,不含候选明细)。"""
|
||||
|
||||
id: str
|
||||
as_of: date
|
||||
method: str
|
||||
universe_size: int = 0
|
||||
selected: int = 0
|
||||
created_at: datetime | None = None
|
||||
@@ -0,0 +1,68 @@
|
||||
"""交易信号领域实体(M8.1,v2 §15 Signal Engine)。
|
||||
|
||||
Signal 输入 = Selection 排序(score)+ 价格/技术条件 + 规则;输出事件可解释:
|
||||
BUY / WATCH / SELL(破位警示),每条带 trigger_reason —— 回答
|
||||
「某日为什么对该股票给 BUY/SELL」(v2 §8)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class SignalRules(BaseModel):
|
||||
"""规则(结构化,MVP):买入区间 + 趋势/动量条件 + 卖出/警示区间。"""
|
||||
|
||||
buy_rank_threshold: int = Field(default=20, ge=1, le=500, description="rank<=此值进入买入候选")
|
||||
buy_require_trend: bool = Field(default=True, description="买入需 close > MA(trend_ma)")
|
||||
buy_require_momentum: bool = Field(default=False, description="买入需 close > 20 日前 close")
|
||||
trend_ma: int = Field(default=60, ge=10, le=250)
|
||||
sell_rank_threshold: int = Field(default=50, ge=1, le=1000, description="rank>此值或破位 → SELL 警示")
|
||||
sell_on_trend_break: bool = Field(default=True, description="买入区间内 close < MA(trend_ma) → SELL")
|
||||
max_output_rank: int = Field(default=80, ge=1, le=2000, description="仅输出排名前 N 的信号")
|
||||
|
||||
|
||||
class SignalEvent(BaseModel):
|
||||
symbol: str
|
||||
signal_date: date
|
||||
signal_type: str = Field(pattern="^(BUY|WATCH|SELL)$")
|
||||
score: float | None = None
|
||||
price: float | None = None
|
||||
trigger_reason: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class SignalStatistics(BaseModel):
|
||||
universe_size: int = 0
|
||||
buy: int = 0
|
||||
watch: int = 0
|
||||
sell: int = 0
|
||||
|
||||
|
||||
class SignalResult(BaseModel):
|
||||
as_of_date: date
|
||||
rules: SignalRules
|
||||
statistics: SignalStatistics
|
||||
events: list[SignalEvent] = Field(default_factory=list)
|
||||
config_snapshot: dict = Field(default_factory=dict)
|
||||
|
||||
|
||||
class SignalMeta(BaseModel):
|
||||
id: str
|
||||
as_of: date
|
||||
buy: int = 0
|
||||
watch: int = 0
|
||||
sell: int = 0
|
||||
created_at: datetime | None = None
|
||||
|
||||
|
||||
class SignalHit(BaseModel):
|
||||
"""个股在信号历史中的命中(by-symbol 查询,供 Chart 标记)。"""
|
||||
|
||||
signal_id: str
|
||||
signal_date: date
|
||||
signal_type: str
|
||||
score: float | None = None
|
||||
price: float | None = None
|
||||
trigger_reason: list[str] = Field(default_factory=list)
|
||||
@@ -0,0 +1,57 @@
|
||||
"""选股策略领域实体(2026-09 重构:原 StrategyDefinition → SelectionStrategy)。
|
||||
|
||||
策略库现在**只存选股条件组合**:股票池 + 因子 + 过滤条件。
|
||||
资金 / 持仓数 / 持仓时间 / 调仓时机 / 费率 / 回测区间一律移到「回测组合」
|
||||
(BacktestCombo)与「公共配置」(GlobalConfig),在回测时才确定。
|
||||
|
||||
仍映射到 `strategy` 表(id/name/description/spec_type/config_json/version/created_at),
|
||||
config_json 只存 universe/factors/conditions —— 迁移会把旧行的多余键剥掉。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
from app.domain.entities.research import ConditionSpec, FactorSpec, UniverseSpec
|
||||
|
||||
|
||||
class SelectionStrategy(BaseModel):
|
||||
"""一个选股策略 = 选股条件组合(不含任何回测执行参数)。
|
||||
|
||||
`extra="forbid"`(2026-09 收尾):请求里若混入旧版的回测执行参数(selection /
|
||||
rebalance / costs / portfolio / initial_capital / period …),一律**报错**而不是
|
||||
静默丢弃 —— 否则调用方会以为「在策略上设了费率/调仓」,实际服务端根本没存
|
||||
(AGENT.md 禁止静默降级与假装支持)。回测参数的正确位置是 BacktestCombo + GlobalConfig。
|
||||
兼容性:历史行残留的旧键由仓储 `_to_entity` 在**读出前**剔除,因此不受 forbid 影响。
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
id: str = ""
|
||||
name: str = Field(min_length=1, max_length=64)
|
||||
description: str = ""
|
||||
# 取值域收敛为 selection:本实体就是「选股策略」。历史 DB 列里的 "backtest"
|
||||
# 不会被读出(仓储 _to_entity 丢弃该键并回落到默认值),故收紧不会破坏旧数据。
|
||||
spec_type: str = Field(default="selection", pattern="^selection$")
|
||||
universe: UniverseSpec = UniverseSpec()
|
||||
factors: list[FactorSpec] = Field(min_length=1, description="打分因子(至少 1 个)")
|
||||
conditions: list[ConditionSpec] = Field(
|
||||
default_factory=list, description="过滤条件(AND,universe 之后、因子排序之前)"
|
||||
)
|
||||
version: str = "1"
|
||||
created_at: datetime | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _no_duplicate_factors(self) -> SelectionStrategy:
|
||||
names = [f.name for f in self.factors]
|
||||
if len(set(names)) != len(names):
|
||||
raise ValueError("factors 存在重复因子名")
|
||||
return self
|
||||
|
||||
|
||||
# 兼容别名:重构前到处引用的旧名字。新代码请用 SelectionStrategy;
|
||||
# 保留别名是为了让尚未迁移的导入点(agent 工具等)在过渡期不炸,
|
||||
# 最终会全部替换掉(见各调用点的 TODO)。
|
||||
StrategyDefinition = SelectionStrategy
|
||||
@@ -0,0 +1,86 @@
|
||||
"""MarketDataProvider(数据源抽象)。
|
||||
|
||||
业务层只依赖本 Protocol(AGENT.md §6),禁止在业务代码中 import
|
||||
tushare / 新浪实现。数据源一律返回 domain.entities 中的归一化实体。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.index import IndexWeight
|
||||
from app.domain.entities.market import (
|
||||
AdjustFactor,
|
||||
DailyBar,
|
||||
DailyBasic,
|
||||
FinancialIndicator,
|
||||
Stock,
|
||||
StockNameHistory,
|
||||
TradingCalendar,
|
||||
)
|
||||
|
||||
|
||||
class MarketDataProvider(Protocol):
|
||||
"""统一市场数据源接口。
|
||||
|
||||
实现约定:
|
||||
- get_daily 返回**不复权**行情(复权经 AdjustFactor 显式计算,禁止静默改口径)
|
||||
- get_financial 返回带 announce_date 的指标,供上层按 as_of 过滤
|
||||
- 实现不得抛出裸连接异常以外的噪音;业务错误应转为 DataSourceError
|
||||
"""
|
||||
|
||||
name: str
|
||||
|
||||
def get_stock_basic(self, list_status: str = "L") -> list[Stock]:
|
||||
"""股票基础信息。list_status: L=上市 / D=退市 / P=暂停上市(tushare 口径)。
|
||||
|
||||
默认 "L" 保持既有行为不变;研究侧的**幸存者偏差**修正依赖 "D"(已退市)
|
||||
——退市股的历史行情与 delist_date 缺失会让回测系统性高估收益。
|
||||
"""
|
||||
...
|
||||
|
||||
def get_trade_cal(self, start: date, end: date) -> list[TradingCalendar]: ...
|
||||
|
||||
def get_daily(self, symbol: str, start: date, end: date) -> list[DailyBar]: ...
|
||||
|
||||
def get_adjust_factor(self, symbol: str, start: date, end: date) -> list[AdjustFactor]: ...
|
||||
|
||||
def get_financial(
|
||||
self,
|
||||
symbol: str,
|
||||
start: date | None = None,
|
||||
end: date | None = None,
|
||||
) -> list[FinancialIndicator]:
|
||||
"""财务指标快照。
|
||||
|
||||
start/end 为**报告期**窗口(对应 Tushare fina_indicator 的
|
||||
start_date/end_date 参数,按报告期过滤);不传表示全量历史。
|
||||
新浪接口不支持按窗口拉取,提供方会忽略窗口后由调用方自行过滤。
|
||||
"""
|
||||
|
||||
def get_index_weight(self, index_code: str) -> list[IndexWeight]:
|
||||
"""指数历史成分(含权重):每期成分快照 → IndexWeight(index_code, trade_date, symbol)。
|
||||
|
||||
供 index_weight 同步与历史成分 Universe(v3 §9)。"""
|
||||
|
||||
def get_name_changes(self, start: date, end: date) -> list[StockNameHistory]:
|
||||
"""区间内全市场股票名称变更(时点 ST / 风险警示判定依据)。
|
||||
|
||||
仅 Tushare 提供;备用源应抛 `DataSourceNotSupported`(不得静默返回空列表,
|
||||
否则名称历史会「看起来同步成功但一条没有」,导致时点 ST 静默降级)。
|
||||
"""
|
||||
...
|
||||
|
||||
def get_daily_basic(self, trade_date: date) -> list[DailyBasic]:
|
||||
"""单交易日全市场每日指标(估值 / 股息率 / 市值)。
|
||||
|
||||
按**交易日**整表拉取(Tushare daily_basic 支持 trade_date 参数一次返回全市场),
|
||||
幂等键 (symbol, trade_date)。
|
||||
|
||||
实现约定:
|
||||
- 只能返回该 trade_date 当天已可得的值(dv_ratio 为时点值,天然无未来函数);
|
||||
- 不支持本接口的数据源(如新浪)必须抛 DataSourceNotSupported,
|
||||
禁止返回空列表冒充成功(AGENT.md §7:禁止静默切换)。
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
"""回测组合 + 公共配置 Repository Protocol(2026-09 重构)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.combo import BacktestCombo, GlobalConfig
|
||||
|
||||
|
||||
class GlobalConfigRepository(Protocol):
|
||||
def get(self) -> GlobalConfig:
|
||||
"""读取全局配置;不存在则返回带默认值的实例(不写库)。"""
|
||||
|
||||
def save(self, config: GlobalConfig) -> GlobalConfig:
|
||||
"""upsert 单例(id 恒为 default)。"""
|
||||
|
||||
|
||||
class ComboRepository(Protocol):
|
||||
def save(self, combo: BacktestCombo) -> BacktestCombo:
|
||||
"""新建或更新(name 冲突抛 ValueError)。"""
|
||||
|
||||
def get(self, combo_id: str) -> BacktestCombo | None: ...
|
||||
|
||||
def list(self) -> list[BacktestCombo]: ...
|
||||
|
||||
def delete(self, combo_id: str) -> bool: ...
|
||||
@@ -0,0 +1,19 @@
|
||||
"""Composite Factor Repository Protocol(M7.2b)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.composite import CompositeDefinition
|
||||
|
||||
|
||||
class CompositeRepository(Protocol):
|
||||
def save(self, definition: CompositeDefinition) -> CompositeDefinition:
|
||||
"""新建组合(name 冲突抛 ValueError)。"""
|
||||
|
||||
def get(self, composite_id: str) -> CompositeDefinition | None: ...
|
||||
|
||||
def list(self) -> list[CompositeDefinition]: ...
|
||||
|
||||
def delete(self, composite_id: str) -> bool:
|
||||
"""删除;不存在返回 False。"""
|
||||
@@ -0,0 +1,29 @@
|
||||
"""字段库 Repository 协议(依赖倒置)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.condition_field import ConditionField
|
||||
|
||||
|
||||
class ConditionFieldRepository(Protocol):
|
||||
def list(self) -> list[ConditionField]:
|
||||
"""按 sort_order, name 排序返回全部条目(含停用项)。"""
|
||||
...
|
||||
|
||||
def get(self, name: str) -> ConditionField | None: ...
|
||||
|
||||
def insert_missing(self, items: list[ConditionField]) -> int:
|
||||
"""只插入不存在的条目(seed 内置字段用)。
|
||||
|
||||
语义上是「只补不删、不覆盖」:已存在的行**原样保留** —— 用户在字段库里
|
||||
改过的中文名/含义不会被下次 seed 冲掉(与 factor_definition 同规矩)。
|
||||
"""
|
||||
...
|
||||
|
||||
def save(self, item: ConditionField) -> ConditionField:
|
||||
"""新增或整体更新一条(自定义字段增改、内置字段改文案/停用)。"""
|
||||
...
|
||||
|
||||
def delete(self, name: str) -> bool: ...
|
||||
@@ -0,0 +1,16 @@
|
||||
"""因子目录 Repository Protocol(M7.1)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.factor import FactorDefinition
|
||||
|
||||
|
||||
class FactorRepository(Protocol):
|
||||
def upsert_many(self, definitions: list[FactorDefinition]) -> int:
|
||||
"""以 name 为幂等键批量写入/更新,返回处理条数。"""
|
||||
|
||||
def list(self) -> list[FactorDefinition]: ...
|
||||
|
||||
def get(self, name: str) -> FactorDefinition | None: ...
|
||||
@@ -0,0 +1,18 @@
|
||||
"""指数成分 Repository Protocol(B1)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from datetime import date
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.index import IndexWeight
|
||||
|
||||
|
||||
class IndexConstituentRepository(Protocol):
|
||||
def upsert_many(self, rows: Sequence[IndexWeight]) -> int: ...
|
||||
|
||||
def members_at(self, index_code: str, as_of: date) -> set[str]:
|
||||
"""as_of 当日成分:取 <=as_of 最近一期快照的股票集合(历史成分语义)。"""
|
||||
|
||||
def latest_date(self, index_code: str) -> date | None: ...
|
||||
@@ -0,0 +1,57 @@
|
||||
"""Job / Experiment Repository Protocol(Phase 4)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.research import ExperimentRecord, ExperimentSummary, JobRecord
|
||||
|
||||
|
||||
class JobRepository(Protocol):
|
||||
def create(self, job: JobRecord) -> JobRecord: ...
|
||||
|
||||
def get(self, job_id: str) -> JobRecord | None: ...
|
||||
|
||||
def update(self, job: JobRecord) -> None: ...
|
||||
|
||||
def list_recent(self, kind: str | None = None, limit: int = 20) -> list[JobRecord]: ...
|
||||
|
||||
def list_by_status(self, status: str, limit: int = 100) -> list[JobRecord]:
|
||||
"""按状态查询(服务启动清理残留 queued/running 用)。"""
|
||||
|
||||
|
||||
class ExperimentRepository(Protocol):
|
||||
def save(self, experiment: ExperimentRecord) -> ExperimentRecord: ...
|
||||
|
||||
def upsert(self, experiment: ExperimentRecord) -> ExperimentRecord:
|
||||
"""按 id 插入或**覆盖**(重建归档用:同一 id 已存在时替换整行)。
|
||||
|
||||
与 `save` 的区别:`save` 只插入(同 id 会主键冲突),用于新归档;
|
||||
`upsert` 用于灾备/重建场景 —— 例如从 Job 副本重建一条被删除的历史归档,
|
||||
此时归档 id 必须保持不变(外部链接、对比记录仍指向它)。
|
||||
"""
|
||||
...
|
||||
|
||||
def get(self, experiment_id: str) -> ExperimentRecord | None: ...
|
||||
|
||||
def list_recent(self, limit: int = 50) -> list[ExperimentRecord]: ...
|
||||
|
||||
def list_filtered(
|
||||
self,
|
||||
*,
|
||||
kind: str | None = None,
|
||||
q: str | None = None,
|
||||
limit: int = 200,
|
||||
offset: int = 0,
|
||||
) -> list[ExperimentSummary]:
|
||||
"""按 kind 精确过滤 + q 模糊过滤(id / 因子名 / summary_text,大小写不敏感)。
|
||||
|
||||
只返回元数据(ExperimentSummary,**不含 result_json**),过滤与分页在
|
||||
SQL 层完成;排序为 created_at 倒序 + id 倒序(同秒创建时保证分页稳定)。
|
||||
"""
|
||||
|
||||
def count_filtered(self, *, kind: str | None = None, q: str | None = None) -> int:
|
||||
"""与 `list_filtered` 同口径的过滤总数(列表接口 X-Total-Count 用)。"""
|
||||
|
||||
def delete(self, experiment_id: str) -> bool:
|
||||
"""删除归档本身,返回是否存在。**不触碰**关联的 Job 记录。"""
|
||||
@@ -0,0 +1,176 @@
|
||||
"""Repository Protocol(Phase 1 数据层)。
|
||||
|
||||
业务层只依赖这些 Protocol;具体实现位于 infrastructure/persistence。
|
||||
实体一律以 domain.entities 类型进出,禁止把 ORM Model 泄漏到上层。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterator, Sequence
|
||||
from datetime import date
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.market import (
|
||||
AdjustFactor,
|
||||
DailyBar,
|
||||
DailyBasic,
|
||||
FinancialIndicator,
|
||||
Stock,
|
||||
StockNameHistory,
|
||||
SyncLog,
|
||||
TradingCalendar,
|
||||
)
|
||||
|
||||
|
||||
class StockRepository(Protocol):
|
||||
def get_by_symbol(self, symbol: str) -> Stock | None: ...
|
||||
|
||||
def list(self) -> list[Stock]: ...
|
||||
|
||||
def upsert_many(self, stocks: Sequence[Stock]) -> int:
|
||||
"""批量写入,以 symbol 为幂等键,返回写入/更新的行数。"""
|
||||
|
||||
|
||||
class TradingCalendarRepository(Protocol):
|
||||
def upsert_many(self, days: Sequence[TradingCalendar]) -> int: ...
|
||||
|
||||
def list_range(self, start: date, end: date) -> list[TradingCalendar]: ...
|
||||
|
||||
def is_open(self, day: date) -> bool: ...
|
||||
|
||||
|
||||
class DailyBarRepository(Protocol):
|
||||
def upsert_many(self, bars: Sequence[DailyBar]) -> int: ...
|
||||
|
||||
def get_range(self, symbol: str, start: date, end: date) -> list[DailyBar]: ...
|
||||
|
||||
def get_range_many(
|
||||
self,
|
||||
symbols: Sequence[str],
|
||||
start: date,
|
||||
end: date,
|
||||
adjust: str = "none",
|
||||
) -> list[DailyBar]:
|
||||
"""批量区间查询(研究装配面板用);adjust 指定行情口径(none 不复权 / qfq)。"""
|
||||
|
||||
def stream_range_many_columns(
|
||||
self,
|
||||
symbols: Sequence[str],
|
||||
start: date,
|
||||
end: date,
|
||||
columns: Sequence[str],
|
||||
adjust: str = "none",
|
||||
) -> Iterator[tuple]:
|
||||
"""流式(分批 yield)返回 symbol, trade_date(iso str), 数值列(float) 元组。
|
||||
|
||||
研究装配大数据面板专用:只 SELECT 所需列并在 SQL 侧转 REAL,
|
||||
避免 ORM 对象 / Decimal 全量物化(内存大头,见内存优化专项)。
|
||||
实现可选——ResearchService 会对缺失该方法的老实现回退到 get_range_many。
|
||||
"""
|
||||
|
||||
def latest_date(self, symbol: str) -> date | None:
|
||||
"""断点续传用:该股票本地已有数据的最新交易日。"""
|
||||
|
||||
|
||||
class AdjustFactorRepository(Protocol):
|
||||
def upsert_many(self, factors: Sequence[AdjustFactor]) -> int: ...
|
||||
|
||||
def get_range(self, symbol: str, start: date, end: date) -> list[AdjustFactor]: ...
|
||||
|
||||
|
||||
class DailyBasicRepository(Protocol):
|
||||
"""每日指标(估值 / 股息率 / 市值)仓储。
|
||||
|
||||
幂等键 (symbol, trade_date)。研究侧一律按 trade_date <= as_of 取用,
|
||||
禁止用未来时点的指标回填历史(AGENT.md §9)。
|
||||
"""
|
||||
|
||||
def upsert_many(self, rows: Sequence[DailyBasic]) -> int: ...
|
||||
|
||||
def get_range(self, symbol: str, start: date, end: date) -> list[DailyBasic]: ...
|
||||
|
||||
def get_range_many(
|
||||
self,
|
||||
symbols: Sequence[str],
|
||||
start: date,
|
||||
end: date,
|
||||
) -> list[DailyBasic]:
|
||||
"""批量区间查询(接口对齐 DailyBarRepository)。"""
|
||||
|
||||
def stream_range_many_columns(
|
||||
self,
|
||||
symbols: Sequence[str],
|
||||
start: date,
|
||||
end: date,
|
||||
columns: Sequence[str],
|
||||
) -> Iterator[tuple]:
|
||||
"""流式返回 symbol, trade_date(iso str), *数值列(float)。
|
||||
|
||||
研究装配面板专用(避免 Decimal / ORM 对象全量物化)。
|
||||
"""
|
||||
|
||||
def latest_date(self) -> date | None:
|
||||
"""本表全局最新交易日(增量同步断点)。"""
|
||||
|
||||
def missing_dates(self, start: date, end: date) -> list[date]:
|
||||
"""区间内「交易日历为开市、但本表无任何行」的交易日(补齐用)。"""
|
||||
|
||||
|
||||
class StockNameHistoryRepository(Protocol):
|
||||
"""股票名称变更历史仓储(时点 ST / 风险警示判定)。
|
||||
|
||||
用于把 `universe.exclude_st` 从「最新名称快照」升级为**时点名称**:
|
||||
研究侧必须按 as_of 取当时生效的名称,否则「曾为高股息、后来才 ST/退市」的
|
||||
股息陷阱样本会被整段排除(实测影响约 3.70pp 收益)。
|
||||
"""
|
||||
|
||||
def upsert_many(self, rows: Sequence[StockNameHistory]) -> int: ...
|
||||
|
||||
def names_as_of(
|
||||
self, symbols: Sequence[str], as_of: date
|
||||
) -> dict[str, str]:
|
||||
"""返回 as_of 时点生效的名称(缺该股记录则不返回该键,由调用方回退最新名称)。"""
|
||||
|
||||
def name_spans(self, symbols: Sequence[str]) -> dict[str, list[tuple[date, date | None, str]]]:
|
||||
"""返回各股票的名称生效区间列表 [(start, end, name)](批量时点查询复用)。"""
|
||||
|
||||
def count_rows(self) -> int:
|
||||
"""本表总行数(未同步时用于降级回退最新名称)。"""
|
||||
|
||||
def namechange_dates(self) -> tuple[date | None, date | None]:
|
||||
"""已同步的 (最早 start_date, 最晚 start_date)(增量断点)。"""
|
||||
|
||||
|
||||
class FinancialRepository(Protocol):
|
||||
def upsert_many(self, rows: Sequence[FinancialIndicator]) -> int: ...
|
||||
|
||||
def list_symbol(self, symbol: str) -> list[FinancialIndicator]:
|
||||
"""该股票本地全部财务行(增量判断 / 新浪校验重叠用,量级小)。"""
|
||||
|
||||
def has_report_period(self, symbol: str, report_date: date) -> bool:
|
||||
"""本地是否已含该报告期(最新应披露报告期是否已入库)。"""
|
||||
|
||||
def list_announced(
|
||||
self,
|
||||
symbol: str,
|
||||
as_of_date: date,
|
||||
report_start: date | None = None,
|
||||
) -> list[FinancialIndicator]:
|
||||
"""只返回 announce_date <= as_of_date 的记录 —— 未来函数红线。"""
|
||||
|
||||
def list_announced_many(
|
||||
self,
|
||||
symbols: Sequence[str],
|
||||
as_of_date: date,
|
||||
) -> list[FinancialIndicator]:
|
||||
"""批量版:返回这些股票 announce_date <= as_of_date 的全部记录。
|
||||
|
||||
供选股/截面研究一次性取财务字段(调用方按需取每 symbol 最新一版)。
|
||||
实现可选 —— 未提供时 SelectionService 回退逐只 list_announced。
|
||||
"""
|
||||
|
||||
|
||||
class SyncLogRepository(Protocol):
|
||||
def add(self, log: SyncLog) -> SyncLog: ...
|
||||
|
||||
def recent(self, source: str | None = None, limit: int = 20) -> list[SyncLog]: ...
|
||||
@@ -0,0 +1,33 @@
|
||||
"""选股 Repository Protocol(M6.3 落库)。
|
||||
|
||||
业务层只依赖本 Protocol;实现位于 infrastructure/persistence。
|
||||
结果按 v2 §8:selection_snapshot(一次运行的查询与统计)+ selection_result(逐候选行),
|
||||
用于回答「2023-08-15 为什么选这只股票」「某日选了什么」。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.chart import SelectionHit
|
||||
from app.domain.entities.selection import SelectionMeta, SelectionResult
|
||||
|
||||
|
||||
class SelectionRepository(Protocol):
|
||||
def save(self, selection_id: str, result: SelectionResult) -> None:
|
||||
"""落库一次选股:snapshot 一行 + 候选逐行(同一事务,由调用方 commit)。"""
|
||||
|
||||
def get(self, selection_id: str) -> SelectionResult | None:
|
||||
"""按 id 读回完整结果(重建 SelectionResult)。"""
|
||||
|
||||
def list_recent(
|
||||
self,
|
||||
as_of: date | None = None,
|
||||
method: str | None = None,
|
||||
limit: int = 20,
|
||||
) -> list[SelectionMeta]:
|
||||
"""历史选股元数据列表(按 created_at 倒序;可选 as_of/method 过滤)。"""
|
||||
|
||||
def list_by_symbol(self, symbol: str, limit: int = 50) -> list[SelectionHit]:
|
||||
"""该股在历史选股中的命中(Chart 标记用,含 as_of/rank/reason)。"""
|
||||
@@ -0,0 +1,22 @@
|
||||
"""Signal Repository Protocol(M8.1 落库)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.signal import SignalHit, SignalMeta, SignalResult
|
||||
|
||||
|
||||
class SignalRepository(Protocol):
|
||||
def save(self, signal_id: str, result: SignalResult) -> None:
|
||||
"""snapshot 一行 + 事件逐行(同事务,调用方 commit)。"""
|
||||
|
||||
def get(self, signal_id: str) -> SignalResult | None: ...
|
||||
|
||||
def list_recent(
|
||||
self, as_of: date | None = None, limit: int = 20
|
||||
) -> list[SignalMeta]: ...
|
||||
|
||||
def list_by_symbol(self, symbol: str, limit: int = 50) -> list[SignalHit]:
|
||||
"""该股在历史信号中的命中(Chart 标记用,含 signal_id 溯源)。"""
|
||||
@@ -0,0 +1,20 @@
|
||||
"""选股策略 Repository Protocol。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
|
||||
|
||||
class StrategyRepository(Protocol):
|
||||
def save(self, definition: SelectionStrategy) -> SelectionStrategy:
|
||||
"""新建(name 冲突抛 ValueError)。"""
|
||||
|
||||
def get(self, strategy_id: str) -> SelectionStrategy | None: ...
|
||||
|
||||
def get_by_name(self, name: str) -> SelectionStrategy | None: ...
|
||||
|
||||
def list(self) -> list[SelectionStrategy]: ...
|
||||
|
||||
def delete(self, strategy_id: str) -> bool: ...
|
||||
@@ -0,0 +1,7 @@
|
||||
"""数据源基础设施:MarketDataProvider 的具体实现与 Failover。
|
||||
|
||||
业务层不直接 import 本目录(AGENT.md §5/§6)——统一经
|
||||
MarketDataProvider(domain.providers)注入;唯一例外是组装处的依赖装配。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -0,0 +1,13 @@
|
||||
"""数据源异常类型。"""
|
||||
|
||||
|
||||
class DataSourceError(Exception):
|
||||
"""数据源通用失败(网络、限频、解析等)。"""
|
||||
|
||||
|
||||
class DataSourceAuthenticationError(DataSourceError):
|
||||
"""凭证无效 / 权限不足(如 Tushare token 无该接口权限)。"""
|
||||
|
||||
|
||||
class DataSourceNotSupported(DataSourceError):
|
||||
"""该数据源不提供此能力(如新浪无复权因子),用于 Failover 判定。"""
|
||||
@@ -0,0 +1,213 @@
|
||||
"""Tushare → Sina Failover 包装(AGENT.md §7)。
|
||||
|
||||
规则:
|
||||
- 优先 primary;primary 抛错时才尝试 fallback(避免对空结果做无谓兜底请求)
|
||||
- fallback 不支持该 API(DataSourceNotSupported)或自身失败 → 抛 DataSourceError
|
||||
- 每次尝试都写 SyncLog(source / 成功与否 / 行数 / 区间),禁止静默切换
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from datetime import date
|
||||
from typing import Any
|
||||
|
||||
from app.domain.entities.market import SyncLog
|
||||
from app.domain.providers import MarketDataProvider
|
||||
from app.infrastructure.data_sources.errors import DataSourceError, DataSourceNotSupported
|
||||
|
||||
|
||||
class FailoverProvider:
|
||||
"""以 primary 为主、fallback 为辅的 MarketDataProvider 实现。"""
|
||||
|
||||
name = "failover"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
primary: MarketDataProvider,
|
||||
fallback: MarketDataProvider | None = None,
|
||||
*,
|
||||
audit: Callable[[SyncLog], None] | None = None,
|
||||
) -> None:
|
||||
self.primary = primary
|
||||
self.fallback = fallback
|
||||
self._audit = audit or (lambda _log: None)
|
||||
|
||||
# ---- 各 API 代理 ----
|
||||
|
||||
def get_stock_basic(self, list_status: str = "L") -> list:
|
||||
return self._with_failover(
|
||||
"get_stock_basic",
|
||||
primary_call=lambda: self.primary.get_stock_basic(list_status),
|
||||
fallback_call=lambda: self.fallback.get_stock_basic(list_status),
|
||||
)
|
||||
|
||||
def get_name_changes(self, start: date, end: date) -> list:
|
||||
return self._with_failover(
|
||||
"get_name_changes",
|
||||
start=start,
|
||||
end=end,
|
||||
primary_call=lambda: self.primary.get_name_changes(start, end),
|
||||
fallback_call=lambda: self.fallback.get_name_changes(start, end),
|
||||
)
|
||||
|
||||
def get_trade_cal(self, start: date, end: date) -> list:
|
||||
return self._with_failover(
|
||||
"get_trade_cal",
|
||||
start=start,
|
||||
end=end,
|
||||
primary_call=lambda: self.primary.get_trade_cal(start, end),
|
||||
fallback_call=lambda: self.fallback.get_trade_cal(start, end),
|
||||
)
|
||||
|
||||
def get_daily(self, symbol: str, start: date, end: date) -> list:
|
||||
return self._with_failover(
|
||||
"get_daily",
|
||||
start=start,
|
||||
end=end,
|
||||
primary_call=lambda: self.primary.get_daily(symbol, start, end),
|
||||
fallback_call=lambda: self.fallback.get_daily(symbol, start, end),
|
||||
)
|
||||
|
||||
def get_index_weight(self, index_code: str) -> list:
|
||||
return self._with_failover(
|
||||
"get_index_weight",
|
||||
primary_call=lambda: self.primary.get_index_weight(index_code),
|
||||
fallback_call=lambda: self.fallback.get_index_weight(index_code),
|
||||
)
|
||||
|
||||
def get_adjust_factor(self, symbol: str, start: date, end: date) -> list:
|
||||
return self._with_failover(
|
||||
"get_adjust_factor",
|
||||
start=start,
|
||||
end=end,
|
||||
primary_call=lambda: self.primary.get_adjust_factor(symbol, start, end),
|
||||
fallback_call=lambda: self.fallback.get_adjust_factor(symbol, start, end),
|
||||
)
|
||||
|
||||
def get_financial(
|
||||
self, symbol: str, start: date | None = None, end: date | None = None
|
||||
) -> list:
|
||||
return self._with_failover(
|
||||
"get_financial",
|
||||
start=start,
|
||||
end=end,
|
||||
primary_call=lambda: self.primary.get_financial(symbol, start, end),
|
||||
fallback_call=lambda: self.fallback.get_financial(symbol, start, end),
|
||||
)
|
||||
|
||||
def get_daily_basic(self, trade_date: date) -> list:
|
||||
"""每日指标:新浪不支持本接口 → 主源失败时如实报错,不做假兜底。"""
|
||||
return self._with_failover(
|
||||
"get_daily_basic",
|
||||
start=trade_date,
|
||||
end=trade_date,
|
||||
primary_call=lambda: self.primary.get_daily_basic(trade_date),
|
||||
fallback_call=(
|
||||
(lambda: self.fallback.get_daily_basic(trade_date))
|
||||
if self.fallback is not None
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
# ---- 内部 ----
|
||||
|
||||
def _with_failover(
|
||||
self,
|
||||
api: str,
|
||||
*,
|
||||
primary_call: Callable[[], list],
|
||||
fallback_call: Callable[[], list] | None = None,
|
||||
start: date | None = None,
|
||||
end: date | None = None,
|
||||
) -> list:
|
||||
try:
|
||||
rows = primary_call()
|
||||
except Exception as exc: # noqa: BLE001 —— 统一走审计
|
||||
self._log(
|
||||
source=self.primary.name,
|
||||
api=api,
|
||||
success=False,
|
||||
reason=str(exc),
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
return self._try_fallback(api, fallback_call, start=start, end=end, primary_error=exc)
|
||||
self._log(
|
||||
source=self.primary.name,
|
||||
api=api,
|
||||
success=True,
|
||||
row_count=_len(rows),
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
return rows
|
||||
|
||||
def _try_fallback(self, api, fallback_call, *, start, end, primary_error):
|
||||
if fallback_call is None or self.fallback is None:
|
||||
raise DataSourceError(
|
||||
f"{self.primary.name}.{api} 失败且无备用源: {primary_error}"
|
||||
) from primary_error
|
||||
try:
|
||||
rows = fallback_call()
|
||||
except DataSourceNotSupported as exc:
|
||||
self._log(
|
||||
source=self.fallback.name,
|
||||
api=api,
|
||||
success=False,
|
||||
reason=f"不支持: {exc}",
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
raise DataSourceError(
|
||||
f"{self.primary.name}.{api} 失败,备用源不支持: {primary_error}"
|
||||
) from primary_error
|
||||
except Exception as exc: # noqa: BLE001
|
||||
self._log(
|
||||
source=self.fallback.name,
|
||||
api=api,
|
||||
success=False,
|
||||
reason=str(exc),
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
raise DataSourceError(
|
||||
f"主备数据源均失败: primary[{self.primary.name}]={primary_error} "
|
||||
f"fallback[{self.fallback.name}]={exc}"
|
||||
) from exc
|
||||
self._log(
|
||||
source=self.fallback.name,
|
||||
api=api,
|
||||
success=True,
|
||||
row_count=_len(rows),
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
return rows
|
||||
|
||||
def _log(
|
||||
self,
|
||||
*,
|
||||
source: str,
|
||||
api: str,
|
||||
success: bool,
|
||||
reason: str | None = None,
|
||||
row_count: int = 0,
|
||||
start: date | None = None,
|
||||
end: date | None = None,
|
||||
) -> None:
|
||||
self._audit(
|
||||
SyncLog(
|
||||
source=source,
|
||||
api=api,
|
||||
success=success,
|
||||
failure_reason=reason,
|
||||
row_count=row_count,
|
||||
data_start=start,
|
||||
data_end=end,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _len(rows: Any) -> int:
|
||||
return len(rows) if rows is not None else 0
|
||||
@@ -0,0 +1,246 @@
|
||||
"""新浪财经 Provider —— 备用数据源。
|
||||
|
||||
通道(公开接口方案参考 cc-cursor/finance/data/sources/sina_source.py):
|
||||
1. 财务:quotes.sina.cn CompanyFinanceService.getFinanceReport2022(source=gjzb,
|
||||
匿名免费、一次多期),含披露日 publish_date → FinancialIndicator
|
||||
(symbol / report_date=end_date / announce_date=publish_date),schema 与
|
||||
Tushare fina_indicator 一致 —— 用于财务兜底(保留防未来函数所需的公告日)。
|
||||
2. 日 K:quotes.sina.cn getKLineData(jsonp,**前复权**)。新浪无「不复权 + 独立复权
|
||||
因子」,因此日线兜底行标记 source=sina、adjust=qfq,与主口径区分;Tushare 恢复
|
||||
后 --resume 会按日覆盖回不复权行。
|
||||
|
||||
能力边界(其余接口新浪不支持 → DataSourceNotSupported):
|
||||
get_stock_basic / get_trade_cal / get_adjust_factor。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from app.domain.entities.market import DailyBar, FinancialIndicator
|
||||
from app.infrastructure.data_sources.errors import (
|
||||
DataSourceError,
|
||||
DataSourceNotSupported,
|
||||
)
|
||||
|
||||
_UA = (
|
||||
"Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 "
|
||||
"(KHTML, like Gecko) Chrome/138.0.0.0 Safari/537.36"
|
||||
)
|
||||
_KLINE_JSONP = (
|
||||
"https://quotes.sina.cn/cn/api/jsonp_v2.php/var%20data=/CN_MarketDataService"
|
||||
".getKLineData?symbol={sina_symbol}&scale=240&ma=no&datalen={datalen}"
|
||||
)
|
||||
_FIN_BASE = "https://quotes.sina.cn/cn/api/openapi.php/CompanyFinanceService.getFinanceReport2022"
|
||||
|
||||
# 新浪「关键指标」中文项名 → 本项目 FinancialIndicator 字段(None 表示已具备/忽略)
|
||||
_FIN_FIELD_MAP = {
|
||||
"基本每股收益": "eps",
|
||||
"净资产收益率(ROE)": "roe",
|
||||
"加权净资产收益率": "roe",
|
||||
"销售毛利率": "gross_margin",
|
||||
"毛利率": "gross_margin",
|
||||
"营业总收入": "total_revenue",
|
||||
"净利润": "net_profit",
|
||||
}
|
||||
|
||||
|
||||
def _to_sina_symbol(symbol: str) -> str:
|
||||
"""600519.SH -> sh600519;000001.SZ -> sz000001;无后缀时按规则猜测。"""
|
||||
code = symbol.strip().upper()
|
||||
if code.endswith(".SH"):
|
||||
return "sh" + code[:-3]
|
||||
if code.endswith(".SZ"):
|
||||
return "sz" + code[:-3]
|
||||
if code.endswith(".BJ"):
|
||||
return "bj" + code[:-3]
|
||||
if code.startswith(("6", "9")):
|
||||
return "sh" + code
|
||||
if code.startswith(("4", "8")):
|
||||
return "bj" + code
|
||||
return "sz" + code
|
||||
|
||||
|
||||
def _extract_jsonp(payload: str) -> list[dict[str, Any]]:
|
||||
"""从 JSONP 中提取数组:容忍前导注释 / var data=([...]) 包裹 / 尾部杂字符。
|
||||
|
||||
直接取首个 '[' 与末个 ']' 之间的内容(行情数组为扁平结构,无嵌套数组)。
|
||||
"""
|
||||
start = payload.find("[")
|
||||
end = payload.rfind("]")
|
||||
if start == -1 or end <= start:
|
||||
raise DataSourceError("新浪行情返回格式无法解析")
|
||||
return json.loads(payload[start : end + 1])
|
||||
|
||||
|
||||
def _d(value) -> Decimal | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return Decimal(str(value))
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
def _to_date(value: str) -> date:
|
||||
"""兼容 20240831 / 2024-08-31 等格式。"""
|
||||
digits = re.sub(r"\D", "", str(value))[:8]
|
||||
return datetime.strptime(digits, "%Y%m%d").date()
|
||||
|
||||
|
||||
class SinaProvider:
|
||||
"""新浪财经备用数据源:财务(与 Tushare schema 一致)+ 日线(前复权)。"""
|
||||
|
||||
name = "sina"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
timeout: float = 10.0,
|
||||
retries: int = 2,
|
||||
urlopen=urllib.request.urlopen,
|
||||
) -> None:
|
||||
self._timeout = timeout
|
||||
self._retries = retries
|
||||
self._urlopen = urlopen
|
||||
|
||||
# ---- HTTP(统一 UA / 重试) ----
|
||||
|
||||
def _open(self, url: str) -> bytes:
|
||||
req = urllib.request.Request(url, headers={"User-Agent": _UA})
|
||||
last_error: Exception | None = None
|
||||
for attempt in range(self._retries):
|
||||
try:
|
||||
with self._urlopen(req, timeout=self._timeout) as resp:
|
||||
return resp.read()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
last_error = exc
|
||||
if attempt < self._retries - 1:
|
||||
time.sleep(0.5 * (attempt + 1))
|
||||
raise DataSourceError(f"sina 请求失败: {last_error}") from last_error
|
||||
|
||||
# ---- 财务(兜底 Tushare fina_indicator) ----
|
||||
|
||||
def get_financial(
|
||||
self,
|
||||
symbol: str,
|
||||
start: date | None = None,
|
||||
end: date | None = None,
|
||||
) -> list[FinancialIndicator]:
|
||||
"""新浪关键指标(source=gjzb),含披露日 publish_date → announce_date。
|
||||
|
||||
新浪不支持按报告期窗口拉取:忽略 start/end 时返回其全部返回的
|
||||
报告期;传入窗口则按 report_date 客户端过滤(新浪行 source=sina)。
|
||||
"""
|
||||
params = {
|
||||
"paperCode": _to_sina_symbol(symbol),
|
||||
"source": "gjzb",
|
||||
"type": "0",
|
||||
"page": "1",
|
||||
"num": "100",
|
||||
}
|
||||
url = f"{_FIN_BASE}?{urllib.parse.urlencode(params)}"
|
||||
payload = json.loads(self._open(url).decode("utf-8", errors="replace"))
|
||||
try:
|
||||
data = payload["result"]["data"]
|
||||
report_dates = [item["date_value"] for item in data["report_date"]]
|
||||
except (KeyError, TypeError) as exc:
|
||||
raise DataSourceError(f"新浪财务返回结构异常({symbol}): {exc}") from exc
|
||||
|
||||
rows: list[FinancialIndicator] = []
|
||||
for rd in report_dates:
|
||||
entry = data["report_list"].get(rd)
|
||||
if not entry:
|
||||
continue
|
||||
announce = entry.get("publish_date")
|
||||
if not announce:
|
||||
continue # 无披露日不可用于研究(防未来函数)
|
||||
report_day = _to_date(str(rd))
|
||||
if start is not None and report_day < start:
|
||||
continue
|
||||
if end is not None and report_day > end:
|
||||
continue
|
||||
fields: dict[str, Decimal | None] = {
|
||||
"eps": None,
|
||||
"roe": None,
|
||||
"total_revenue": None,
|
||||
"net_profit": None,
|
||||
"gross_margin": None,
|
||||
}
|
||||
for item in entry.get("data", []):
|
||||
std = _FIN_FIELD_MAP.get(item.get("item_title", ""))
|
||||
if std and fields.get(std) is None:
|
||||
fields[std] = _d(item.get("item_value"))
|
||||
rows.append(
|
||||
FinancialIndicator(
|
||||
symbol=symbol,
|
||||
report_date=report_day,
|
||||
announce_date=_to_date(str(announce)),
|
||||
source="sina",
|
||||
eps=fields["eps"],
|
||||
roe=fields["roe"],
|
||||
total_revenue=fields["total_revenue"],
|
||||
net_profit=fields["net_profit"],
|
||||
gross_margin=fields["gross_margin"],
|
||||
)
|
||||
)
|
||||
return rows
|
||||
|
||||
# ---- 日 K(前复权兜底,标记 adjust=qfq) ----
|
||||
|
||||
def get_daily(self, symbol: str, start: date, end: date, datalen: int = 320) -> list[DailyBar]:
|
||||
"""拉取前复权日 K(新浪仅支持最近 datalen 个自然日窗口)。"""
|
||||
url = _KLINE_JSONP.format(sina_symbol=_to_sina_symbol(symbol), datalen=datalen)
|
||||
payload = self._open(url).decode("utf-8", errors="replace")
|
||||
|
||||
bars: list[DailyBar] = []
|
||||
for rec in _extract_jsonp(payload):
|
||||
day = datetime.strptime(rec["day"], "%Y-%m-%d").date()
|
||||
if day < start or day > end:
|
||||
continue
|
||||
bars.append(
|
||||
DailyBar(
|
||||
symbol=symbol,
|
||||
trade_date=day,
|
||||
source="sina",
|
||||
adjust="qfq",
|
||||
open=_d(rec.get("open")),
|
||||
high=_d(rec.get("high")),
|
||||
low=_d(rec.get("low")),
|
||||
close=_d(rec.get("close")),
|
||||
volume=_d(rec.get("volume")),
|
||||
)
|
||||
)
|
||||
return bars
|
||||
|
||||
# ---- 不支持 ----
|
||||
|
||||
def get_stock_basic(self, list_status: str = "L"):
|
||||
raise DataSourceNotSupported("新浪不提供股票基础信息列表")
|
||||
|
||||
def get_name_changes(self, start, end):
|
||||
raise DataSourceNotSupported("新浪不提供股票名称变更历史(namechange)")
|
||||
|
||||
def get_index_weight(self, index_code):
|
||||
raise DataSourceNotSupported("新浪不提供指数成分接口")
|
||||
|
||||
def get_trade_cal(self, start, end):
|
||||
raise DataSourceNotSupported("新浪不提供交易日历")
|
||||
|
||||
def get_adjust_factor(self, symbol, start, end):
|
||||
raise DataSourceNotSupported("新浪不提供复权因子(日线接口为前复权口径)")
|
||||
|
||||
def get_daily_basic(self, trade_date):
|
||||
"""新浪无估值/股息率接口。
|
||||
|
||||
禁止返回空列表冒充成功 —— 否则 daily_basic 同步会把「接口不支持」
|
||||
误记成「当日无数据」,导致缺口被静默固化(AGENT.md §7)。
|
||||
"""
|
||||
raise DataSourceNotSupported("新浪不提供每日指标(估值/股息率/市值)接口")
|
||||
@@ -0,0 +1,474 @@
|
||||
"""Tushare Provider —— 首选数据源实现。
|
||||
|
||||
依赖注入:pro 客户端(tushare.pro.client 或测试 Fake)。真实运行时惰性加载
|
||||
tushare 库(pyproject optional:uv sync --extra datasource-tushare)。
|
||||
归一化函数只依赖 list[dict],便于无 pandas 环境下单测。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from datetime import date, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from app.domain.entities.index import IndexWeight
|
||||
from app.domain.entities.market import (
|
||||
AdjustFactor,
|
||||
DailyBar,
|
||||
DailyBasic,
|
||||
FinancialIndicator,
|
||||
Stock,
|
||||
StockNameHistory,
|
||||
TradingCalendar,
|
||||
)
|
||||
from app.infrastructure.data_sources.errors import (
|
||||
DataSourceAuthenticationError,
|
||||
DataSourceError,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_TS_DATE = "%Y%m%d"
|
||||
|
||||
|
||||
def _to_date(value: str | None) -> date | None:
|
||||
"""日期归一:None / NaN / 空串 → None。
|
||||
|
||||
pandas 读到的缺失日期是 float NaN(namechange 的 end_date、部分财务字段),
|
||||
若不拦住会抛 `time data 'nan' does not match format '%Y%m%d'` 并**中断整批拉取**
|
||||
——实测 namechange 按年分片时 32/37 片因此失败。
|
||||
"""
|
||||
if value is None or value == "":
|
||||
return None
|
||||
if isinstance(value, float) and value != value: # NaN
|
||||
return None
|
||||
text = str(value).strip()
|
||||
if text == "" or text.lower() in {"nan", "none", "null", "nat"}:
|
||||
return None
|
||||
return datetime.strptime(text[:10], _TS_DATE).date()
|
||||
|
||||
|
||||
# 本地 symbol 规范:6 位数字 + 交易所后缀(与 Stock 实体的 pattern 校验一致)
|
||||
_SYMBOL_RE = re.compile(r"^\d{6}\.(SH|SZ|BJ)$")
|
||||
|
||||
|
||||
def _to_opt_str(value) -> str | None:
|
||||
"""可选字符串字段归一:None/NaN/空串 → None。
|
||||
|
||||
pandas 读到的缺失值是 float NaN(如退市股的 industry/area),直接塞进
|
||||
`str | None` 字段会被 pydantic 拒绝(string_type)——实测退市股拉取时命中。
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, float) and value != value: # NaN
|
||||
return None
|
||||
text = str(value).strip()
|
||||
if text == "" or text.lower() in {"nan", "none", "null"}:
|
||||
return None
|
||||
return text
|
||||
|
||||
|
||||
def _to_decimal(value) -> Decimal | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
num = float(value)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
if num != num: # NaN
|
||||
return None
|
||||
return Decimal(str(num))
|
||||
|
||||
|
||||
class TushareProvider:
|
||||
"""封装 Tushare Pro(ts.pro_api)。所有输出已归一化为领域实体。"""
|
||||
|
||||
name = "tushare"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
token: str = "",
|
||||
*,
|
||||
pro: object | None = None,
|
||||
max_retries: int = 3,
|
||||
rate_limit_wait: float = 30.0,
|
||||
) -> None:
|
||||
self._pro = pro if pro is not None else _build_pro(token)
|
||||
self._max_retries = max_retries
|
||||
self._rate_limit_wait = rate_limit_wait
|
||||
|
||||
# ---- 归一化(纯函数,输入 list[dict],可单测) ----
|
||||
|
||||
@staticmethod
|
||||
def normalize_stock(
|
||||
records: list[dict[str, Any]], default_status: str = "L"
|
||||
) -> list[Stock]:
|
||||
"""归一化为 Stock。
|
||||
|
||||
`default_status`:tushare `stock_basic(list_status='D')` 返回的 status 字段
|
||||
为空(实测 None),若一律兜底成 "L" 会把退市股标成在市 → 调用方按查询的
|
||||
list_status 传入,保证 status 与 delist_date 语义一致。
|
||||
"""
|
||||
stocks: list[Stock] = []
|
||||
for rec in records:
|
||||
stocks.append(
|
||||
Stock(
|
||||
symbol=str(rec.get("ts_code") or rec.get("symbol") or ""),
|
||||
name=str(rec.get("name") or "").strip(),
|
||||
industry=_to_opt_str(rec.get("industry")),
|
||||
area=_to_opt_str(rec.get("area")),
|
||||
market=_to_opt_str(rec.get("market")),
|
||||
exchange=_to_opt_str(rec.get("exchange")),
|
||||
list_date=_to_date(rec.get("list_date")) or date.min,
|
||||
delist_date=_to_date(rec.get("delist_date")),
|
||||
status=_to_opt_str(rec.get("status")) or default_status,
|
||||
)
|
||||
)
|
||||
return stocks
|
||||
|
||||
@staticmethod
|
||||
def normalize_calendar(records: list[dict[str, Any]]) -> list[TradingCalendar]:
|
||||
return [
|
||||
TradingCalendar(
|
||||
calendar_date=_to_date(rec.get("cal_date")) or date.min,
|
||||
is_open=bool(rec.get("is_open")),
|
||||
)
|
||||
for rec in records
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def normalize_index_weight(
|
||||
records: list[dict[str, Any]], index_code_fallback: str = ""
|
||||
) -> list[IndexWeight]:
|
||||
"""index_weight 接口行 → IndexWeight(index_code/con_code/trade_date/weight)。"""
|
||||
out: list[IndexWeight] = []
|
||||
for rec in records:
|
||||
code = str(rec.get("index_code") or index_code_fallback or "")
|
||||
symbol = str(rec.get("con_code") or "")
|
||||
if not code or not symbol:
|
||||
continue
|
||||
out.append(
|
||||
IndexWeight(
|
||||
index_code=code,
|
||||
index_name=rec.get("index_name"),
|
||||
trade_date=_to_date(rec.get("trade_date")) or date.min,
|
||||
symbol=symbol,
|
||||
weight=_to_decimal(rec.get("weight")),
|
||||
)
|
||||
)
|
||||
return out
|
||||
|
||||
@staticmethod
|
||||
def normalize_daily(records: list[dict[str, Any]]) -> list[DailyBar]:
|
||||
bars: list[DailyBar] = []
|
||||
for rec in records:
|
||||
vol = _to_decimal(rec.get("vol"))
|
||||
amount = _to_decimal(rec.get("amount"))
|
||||
bars.append(
|
||||
DailyBar(
|
||||
symbol=str(rec.get("ts_code") or ""),
|
||||
trade_date=_to_date(rec.get("trade_date")) or date.min,
|
||||
source="tushare",
|
||||
adjust="none",
|
||||
open=_to_decimal(rec.get("open")),
|
||||
high=_to_decimal(rec.get("high")),
|
||||
low=_to_decimal(rec.get("low")),
|
||||
close=_to_decimal(rec.get("close")),
|
||||
volume=vol * 100 if vol is not None else None,
|
||||
amount=amount * 1000 if amount is not None else None,
|
||||
)
|
||||
)
|
||||
return bars
|
||||
|
||||
@staticmethod
|
||||
def normalize_adj_factor(records: list[dict[str, Any]]) -> list[AdjustFactor]:
|
||||
return [
|
||||
AdjustFactor(
|
||||
symbol=str(rec.get("ts_code") or ""),
|
||||
trade_date=_to_date(rec.get("trade_date")) or date.min,
|
||||
factor=_to_decimal(rec.get("adj_factor")) or Decimal(1),
|
||||
)
|
||||
for rec in records
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def normalize_daily_basic(records: list[dict[str, Any]]) -> list[DailyBasic]:
|
||||
"""daily_basic → DailyBasic。
|
||||
|
||||
单位保持 Tushare 原样(不做隐式换算,避免口径漂移):
|
||||
- dv_ratio / dv_ttm / turnover_rate / volume_ratio / pe / pb / ps … 为百分数或倍数
|
||||
- total_share / float_share / free_share 单位万股;total_mv / circ_mv 单位万元
|
||||
- close 为**不复权**收盘价,与 stock_daily(adjust=none) 同口径
|
||||
"""
|
||||
rows: list[DailyBasic] = []
|
||||
for rec in records:
|
||||
rows.append(
|
||||
DailyBasic(
|
||||
symbol=str(rec.get("ts_code") or ""),
|
||||
trade_date=_to_date(rec.get("trade_date")) or date.min,
|
||||
source="tushare",
|
||||
close=_to_decimal(rec.get("close")),
|
||||
turnover_rate=_to_decimal(rec.get("turnover_rate")),
|
||||
volume_ratio=_to_decimal(rec.get("volume_ratio")),
|
||||
pe=_to_decimal(rec.get("pe")),
|
||||
pe_ttm=_to_decimal(rec.get("pe_ttm")),
|
||||
pb=_to_decimal(rec.get("pb")),
|
||||
ps=_to_decimal(rec.get("ps")),
|
||||
ps_ttm=_to_decimal(rec.get("ps_ttm")),
|
||||
dv_ratio=_to_decimal(rec.get("dv_ratio")),
|
||||
dv_ttm=_to_decimal(rec.get("dv_ttm")),
|
||||
total_share=_to_decimal(rec.get("total_share")),
|
||||
float_share=_to_decimal(rec.get("float_share")),
|
||||
free_share=_to_decimal(rec.get("free_share")),
|
||||
total_mv=_to_decimal(rec.get("total_mv")),
|
||||
circ_mv=_to_decimal(rec.get("circ_mv")),
|
||||
)
|
||||
)
|
||||
return rows
|
||||
|
||||
@staticmethod
|
||||
def normalize_name_history(records: list[dict[str, Any]]) -> list[StockNameHistory]:
|
||||
"""namechange → StockNameHistory(名称生效区间)。
|
||||
|
||||
注意:`namechange` 的区间是**完整历史**(一行一个名称生效段),
|
||||
`end_date` 为 NaN 表示「至今有效」;`change_reason` 为 ST/*ST/撤销ST 等。
|
||||
"""
|
||||
rows: list[StockNameHistory] = []
|
||||
for rec in records:
|
||||
symbol = _to_opt_str(rec.get("ts_code"))
|
||||
start = _to_date(rec.get("start_date"))
|
||||
name = _to_opt_str(rec.get("name"))
|
||||
if not symbol or not start or not name or not _SYMBOL_RE.match(symbol):
|
||||
continue
|
||||
rows.append(
|
||||
StockNameHistory(
|
||||
symbol=symbol,
|
||||
name=name,
|
||||
start_date=start,
|
||||
end_date=_to_date(rec.get("end_date")),
|
||||
ann_date=_to_date(rec.get("ann_date")),
|
||||
change_reason=_to_opt_str(rec.get("change_reason")),
|
||||
source="tushare",
|
||||
)
|
||||
)
|
||||
return rows
|
||||
|
||||
@staticmethod
|
||||
def normalize_financial(records: list[dict[str, Any]]) -> list[FinancialIndicator]:
|
||||
rows: list[FinancialIndicator] = []
|
||||
for rec in records:
|
||||
rows.append(
|
||||
FinancialIndicator(
|
||||
symbol=str(rec.get("ts_code") or ""),
|
||||
report_date=_to_date(rec.get("end_date")) or date.min,
|
||||
announce_date=_to_date(rec.get("ann_date")) or date.min,
|
||||
eps=_to_decimal(rec.get("eps")),
|
||||
roe=_to_decimal(rec.get("roe")),
|
||||
net_profit=_to_decimal(rec.get("n_income_attr_p")),
|
||||
gross_margin=_to_decimal(rec.get("grossprofit_margin")),
|
||||
)
|
||||
)
|
||||
return rows
|
||||
|
||||
# ---- 接口调用 ----
|
||||
|
||||
def get_stock_basic(self, list_status: str = "L") -> list[Stock]:
|
||||
"""股票基础信息(list_status: L=上市 / D=退市 / P=暂停上市)。
|
||||
|
||||
tushare `stock_basic` 不带 list_status 时**只返回在市股票**,因此
|
||||
`delist_date` 恒为空、退市股整体缺失 → 回测存在幸存者偏差。
|
||||
需要退市股时必须显式传 "D"(实测 2019-12 之后退市 230 只)。
|
||||
"""
|
||||
records = self._call(
|
||||
"stock_basic",
|
||||
list_status=list_status,
|
||||
fields="ts_code,symbol,name,area,industry,market,exchange,list_date,delist_date,status",
|
||||
)
|
||||
# 代码规范过滤:tushare 退市表含极少数非本地代码规范的记录
|
||||
# (实测 'T600018.SH' = 上港集箱(退),2006 年退市,T 前缀表示转入三板),
|
||||
# 直接归一化会因 symbol 正则校验失败而**中断整个列表** —— 跳过并如实告警,
|
||||
# 不做静默丢弃(AGENT.md §24)。
|
||||
kept, skipped = [], []
|
||||
for rec in records:
|
||||
code = str(rec.get("ts_code") or rec.get("symbol") or "")
|
||||
(kept if _SYMBOL_RE.match(code) else skipped).append(rec)
|
||||
if skipped:
|
||||
logger.warning(
|
||||
"tushare.stock_basic(list_status=%s) 跳过 %d 条不符合本地代码规范的记录:%s",
|
||||
list_status,
|
||||
len(skipped),
|
||||
[r.get("ts_code") for r in skipped[:5]],
|
||||
)
|
||||
return self.normalize_stock(kept, default_status=list_status)
|
||||
|
||||
def get_trade_cal(self, start: date, end: date) -> list[TradingCalendar]:
|
||||
records = self._call(
|
||||
"trade_cal",
|
||||
exchange="SSE",
|
||||
start_date=start.strftime(_TS_DATE),
|
||||
end_date=end.strftime(_TS_DATE),
|
||||
is_open="",
|
||||
)
|
||||
return self.normalize_calendar(records)
|
||||
|
||||
def get_daily(self, symbol: str, start: date, end: date) -> list[DailyBar]:
|
||||
records = self._call(
|
||||
"daily",
|
||||
ts_code=symbol,
|
||||
start_date=start.strftime(_TS_DATE),
|
||||
end_date=end.strftime(_TS_DATE),
|
||||
)
|
||||
return self.normalize_daily(records)
|
||||
|
||||
def get_adjust_factor(self, symbol: str, start: date, end: date) -> list[AdjustFactor]:
|
||||
records = self._call(
|
||||
"adj_factor",
|
||||
ts_code=symbol,
|
||||
start_date=start.strftime(_TS_DATE),
|
||||
end_date=end.strftime(_TS_DATE),
|
||||
)
|
||||
return self.normalize_adj_factor(records)
|
||||
|
||||
def get_financial(
|
||||
self,
|
||||
symbol: str,
|
||||
start: date | None = None,
|
||||
end: date | None = None,
|
||||
) -> list[FinancialIndicator]:
|
||||
"""fina_indicator:报告期窗口 + 100 条/请求上限自动分页。
|
||||
|
||||
Tushare 单次请求最多返回 100 条(超出按最新 100 条截断),因此
|
||||
全量历史必须按报告期窗口回卷分页,否则老报告期会被静默丢弃。
|
||||
"""
|
||||
lo = start or date(1990, 1, 1)
|
||||
hi = end or date.today()
|
||||
raw: list[dict[str, Any]] = []
|
||||
while lo <= hi:
|
||||
batch = self._call(
|
||||
"fina_indicator",
|
||||
ts_code=symbol,
|
||||
start_date=lo.strftime(_TS_DATE),
|
||||
end_date=hi.strftime(_TS_DATE),
|
||||
)
|
||||
raw += batch
|
||||
if len(batch) < 100:
|
||||
break
|
||||
ends = [
|
||||
datetime.strptime(str(r["end_date"])[:8], _TS_DATE).date()
|
||||
for r in batch
|
||||
if r.get("end_date")
|
||||
]
|
||||
if not ends:
|
||||
break
|
||||
next_hi = min(ends) - timedelta(days=1)
|
||||
if next_hi < lo: # 无进展保护(边界簇被截断等极端情况)
|
||||
break
|
||||
hi = next_hi
|
||||
return self.normalize_financial(raw)
|
||||
|
||||
# ---- 内部 ----
|
||||
|
||||
_RATE_LIMIT_MARKERS = ("频率超限", "每分钟", "frequenc", "too many")
|
||||
|
||||
def get_index_weight(self, index_code: str) -> list[IndexWeight]:
|
||||
"""指数成分(Tushare index_weight 全历史,ts_code 过滤)。"""
|
||||
records = self._call("index_weight", ts_code=index_code)
|
||||
return self.normalize_index_weight(records, index_code_fallback=index_code)
|
||||
|
||||
# Tushare 单次接口返回上限(实测 daily_basic 全市场单日 3700~5600 行、
|
||||
# namechange 2020+ 区间 4031 行):取满即告警,避免静默截断。
|
||||
MAX_ROWS_PER_CALL = 6000
|
||||
|
||||
# namechange 单次请求上限同样约 6000 行;实测 2020+ 区间仅 4031 行,
|
||||
# 但全历史(1990 起)会超限 —— 由 Syncer 按年分片调用,避免静默截断。
|
||||
_NAMECHANGE_FIELDS = "ts_code,name,start_date,end_date,ann_date,change_reason"
|
||||
|
||||
def get_name_changes(self, start: date, end: date) -> list[StockNameHistory]:
|
||||
"""区间内全市场名称变更(Tushare namechange,按公告/生效区间批量取)。"""
|
||||
records = self._call(
|
||||
"namechange",
|
||||
start_date=start.strftime(_TS_DATE),
|
||||
end_date=end.strftime(_TS_DATE),
|
||||
fields=self._NAMECHANGE_FIELDS,
|
||||
)
|
||||
if len(records) >= self.MAX_ROWS_PER_CALL:
|
||||
logger.warning(
|
||||
"tushare.namechange(%s~%s) 返回 %d 行,可能触及单次上限被截断,"
|
||||
"请缩小区间后重跑",
|
||||
start,
|
||||
end,
|
||||
len(records),
|
||||
)
|
||||
return self.normalize_name_history(records)
|
||||
|
||||
# daily_basic 单次请求上限 6000 行(全市场一日约 3700~5600 行),按交易日调用即可
|
||||
_DAILY_BASIC_FIELDS = (
|
||||
"ts_code,trade_date,close,turnover_rate,volume_ratio,pe,pe_ttm,pb,ps,ps_ttm,"
|
||||
"dv_ratio,dv_ttm,total_share,float_share,free_share,total_mv,circ_mv"
|
||||
)
|
||||
|
||||
def get_daily_basic(self, trade_date: date) -> list[DailyBasic]:
|
||||
"""单交易日全市场每日指标(daily_basic)。
|
||||
|
||||
注意:Tushare 单次 6000 行上限 —— 全市场单日实测 3700~5600 行
|
||||
(2020 年约 3700,2026 年约 5560),当前安全;但若未来上市公司数
|
||||
逼近 6000,需要按 ts_code 分片。此处对「恰好取满 6000 行」做告警,
|
||||
避免静默截断(AGENTS §7 数据可追溯)。
|
||||
"""
|
||||
records = self._call(
|
||||
"daily_basic",
|
||||
trade_date=trade_date.strftime(_TS_DATE),
|
||||
fields=self._DAILY_BASIC_FIELDS,
|
||||
)
|
||||
if len(records) >= self.MAX_ROWS_PER_CALL:
|
||||
logger.warning(
|
||||
"daily_basic %s 返回 %d 行(达到 %d 行上限),可能被截断,需按 ts_code 分片",
|
||||
trade_date,
|
||||
len(records),
|
||||
self.MAX_ROWS_PER_CALL,
|
||||
)
|
||||
return self.normalize_daily_basic(records)
|
||||
|
||||
def _call(self, api: str, **kwargs) -> list[dict[str, Any]]:
|
||||
"""带限速退避的调用:频率超限按指数退避(最长 _rate_limit_wait)等待后重试。"""
|
||||
last_error: Exception | None = None
|
||||
for attempt in range(self._max_retries):
|
||||
try:
|
||||
fn = getattr(self._pro, api)
|
||||
result = fn(**kwargs)
|
||||
if result is None:
|
||||
return []
|
||||
if hasattr(result, "to_dict"):
|
||||
return result.to_dict("records")
|
||||
if isinstance(result, list):
|
||||
return result
|
||||
return []
|
||||
except Exception as exc: # noqa: BLE001 —— tushare 异常无统一类型,逐一归类
|
||||
last_error = exc
|
||||
msg = str(exc)
|
||||
if "权限" in msg or "积分" in msg or "token" in msg.lower():
|
||||
raise DataSourceAuthenticationError(msg) from exc
|
||||
if any(marker in msg for marker in self._RATE_LIMIT_MARKERS):
|
||||
wait = min(self._rate_limit_wait, 2 ** (attempt + 1))
|
||||
logger.warning("tushare.%s 频率超限,退避 %.1fs 后重试", api, wait)
|
||||
time.sleep(wait)
|
||||
raise DataSourceError(
|
||||
f"tushare.{api} 重试 {self._max_retries} 次仍失败: {last_error}"
|
||||
) from last_error
|
||||
|
||||
|
||||
def _build_pro(token: str):
|
||||
if not token:
|
||||
raise DataSourceAuthenticationError(
|
||||
"缺少 TUSHARE_TOKEN:请 cp .env.example .env 并填入 Tushare Pro token"
|
||||
)
|
||||
try:
|
||||
ts = importlib.import_module("tushare")
|
||||
except ImportError as exc: # pragma: no cover —— 环境相关
|
||||
raise DataSourceError(
|
||||
"未安装 tushare 客户端:cd backend && uv sync --extra datasource-tushare"
|
||||
) from exc
|
||||
return ts.pro_api(token)
|
||||
@@ -11,6 +11,7 @@ from logging.config import fileConfig
|
||||
|
||||
from alembic import context
|
||||
from app.core.config import get_settings
|
||||
from app.infrastructure.persistence.sqlalchemy import models as _models # noqa: F401 —— 注册全部表
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from sqlalchemy import engine_from_config, pool
|
||||
|
||||
@@ -19,7 +20,11 @@ config = context.config
|
||||
if config.config_file_name is not None:
|
||||
fileConfig(config.config_file_name)
|
||||
|
||||
config.set_main_option("sqlalchemy.url", get_settings().database_url)
|
||||
# alembic.ini 中显式 sqlalchemy.url 优先(测试/运维可注入);否则用应用配置
|
||||
_db_url = config.get_main_option("sqlalchemy.url")
|
||||
if not _db_url:
|
||||
_db_url = get_settings().database_url
|
||||
config.set_main_option("sqlalchemy.url", _db_url)
|
||||
|
||||
target_metadata = Base.metadata
|
||||
|
||||
|
||||
+67
@@ -0,0 +1,67 @@
|
||||
"""add daily_basic table
|
||||
|
||||
Revision ID: 22d7380706f7
|
||||
Revises: f5e0d1c2b3a4
|
||||
Create Date: 2026-09-19 15:53:37.029097
|
||||
|
||||
每日指标表(Tushare daily_basic):估值 / 股息率 / 市值,幂等键 (symbol, trade_date)。
|
||||
供高股息等横截面选股因子使用(研究侧按 trade_date <= as_of 取用,无未来函数)。
|
||||
|
||||
说明:本文件由 `alembic revision --autogenerate` 生成后**手工裁剪**。
|
||||
自动生成时同时检出了 `index_weight` 的 drop —— 那是 `IndexWeightModel`
|
||||
未在 `models/__init__.py` 注册导致的假差异(表实际存在),
|
||||
已在同一提交中补上注册,此处只保留 daily_basic 的建表语句。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = '22d7380706f7'
|
||||
down_revision: str | None = 'f5e0d1c2b3a4'
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
'daily_basic',
|
||||
sa.Column(
|
||||
'id',
|
||||
sa.BigInteger().with_variant(sa.Integer(), 'sqlite'),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column('symbol', sa.String(length=12), nullable=False),
|
||||
sa.Column('trade_date', sa.Date(), nullable=False),
|
||||
sa.Column('source', sa.String(length=16), server_default='tushare', nullable=False),
|
||||
sa.Column('close', sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column('turnover_rate', sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column('volume_ratio', sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column('pe', sa.Numeric(precision=16, scale=4), nullable=True),
|
||||
sa.Column('pe_ttm', sa.Numeric(precision=16, scale=4), nullable=True),
|
||||
sa.Column('pb', sa.Numeric(precision=16, scale=4), nullable=True),
|
||||
sa.Column('ps', sa.Numeric(precision=16, scale=4), nullable=True),
|
||||
sa.Column('ps_ttm', sa.Numeric(precision=16, scale=4), nullable=True),
|
||||
sa.Column('dv_ratio', sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column('dv_ttm', sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column('total_share', sa.Numeric(precision=24, scale=4), nullable=True),
|
||||
sa.Column('float_share', sa.Numeric(precision=24, scale=4), nullable=True),
|
||||
sa.Column('free_share', sa.Numeric(precision=24, scale=4), nullable=True),
|
||||
sa.Column('total_mv', sa.Numeric(precision=24, scale=4), nullable=True),
|
||||
sa.Column('circ_mv', sa.Numeric(precision=24, scale=4), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('symbol', 'trade_date', name='uq_basic_symbol_date'),
|
||||
)
|
||||
with op.batch_alter_table('daily_basic', schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f('ix_daily_basic_symbol'), ['symbol'], unique=False)
|
||||
batch_op.create_index(batch_op.f('ix_daily_basic_trade_date'), ['trade_date'], unique=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.batch_alter_table('daily_basic', schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f('ix_daily_basic_trade_date'))
|
||||
batch_op.drop_index(batch_op.f('ix_daily_basic_symbol'))
|
||||
op.drop_table('daily_basic')
|
||||
+65
@@ -0,0 +1,65 @@
|
||||
"""phase4 jobs and experiments
|
||||
|
||||
Revision ID: 53113c80257f
|
||||
Revises: e4d188250fb2
|
||||
Create Date: 2026-09-06 17:16:05.976236
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "53113c80257f"
|
||||
down_revision: str | None = "e4d188250fb2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_table(
|
||||
"experiment",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("kind", sa.String(length=16), nullable=False),
|
||||
sa.Column("spec_json", sa.Text(), nullable=False),
|
||||
sa.Column("result_json", sa.Text(), nullable=False),
|
||||
sa.Column("summary_text", sa.String(length=200), nullable=True),
|
||||
sa.Column("code_version", sa.String(length=40), nullable=True),
|
||||
sa.Column("data_version", sa.String(length=40), nullable=True),
|
||||
sa.Column("job_id", sa.String(length=32), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_table(
|
||||
"job",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("kind", sa.String(length=16), nullable=False),
|
||||
sa.Column("status", sa.String(length=12), nullable=False),
|
||||
sa.Column("stage", sa.String(length=24), nullable=True),
|
||||
sa.Column("spec_json", sa.Text(), nullable=False),
|
||||
sa.Column("error", sa.Text(), nullable=True),
|
||||
sa.Column("result_json", sa.Text(), nullable=True),
|
||||
sa.Column("experiment_id", sa.String(length=32), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("started_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("finished_at", sa.DateTime(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
with op.batch_alter_table("job", schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f("ix_job_status"), ["status"], unique=False)
|
||||
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
with op.batch_alter_table("job", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_job_status"))
|
||||
|
||||
op.drop_table("job")
|
||||
op.drop_table("experiment")
|
||||
# ### end Alembic commands ###
|
||||
+47
@@ -0,0 +1,47 @@
|
||||
"""research result_json → MEDIUMTEXT(长回测结果落库)
|
||||
|
||||
Revision ID: 7b1c4e9a52d8
|
||||
Revises: 22d7380706f7
|
||||
Create Date: 2026-09-19
|
||||
|
||||
背景:全市场多年回测结果(净值/回撤曲线 + 逐笔成交 + Signal↔Fill + 个股收益曲线)
|
||||
实测约 1.2MB,MySQL `TEXT`(64KB)会报 1406 Data too long → Job 归档失败
|
||||
(2026-09 高股息案例实测:回测本身成功,落库失败)。
|
||||
|
||||
本迁移把 job.result_json / experiment.result_json 放宽到 MEDIUMTEXT(16MB)。
|
||||
SQLite 等方言不区分 TEXT 长度,迁移里做方言判断,非 MySQL 直接跳过。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from alembic import op
|
||||
from sqlalchemy import text
|
||||
|
||||
revision = "7b1c4e9a52d8"
|
||||
down_revision = "22d7380706f7"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
_COLUMNS = (("job", "result_json", True), ("experiment", "result_json", False))
|
||||
|
||||
|
||||
def _dialect() -> str:
|
||||
return op.get_bind().dialect.name
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if _dialect() != "mysql":
|
||||
return # SQLite/TEXT 无长度限制,无需变更
|
||||
for table, column, nullable in _COLUMNS:
|
||||
null_clause = "NULL" if nullable else "NOT NULL"
|
||||
op.execute(
|
||||
text(f"ALTER TABLE {table} MODIFY COLUMN {column} MEDIUMTEXT {null_clause}")
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if _dialect() != "mysql":
|
||||
return
|
||||
for table, column, nullable in _COLUMNS:
|
||||
null_clause = "NULL" if nullable else "NOT NULL"
|
||||
op.execute(text(f"ALTER TABLE {table} MODIFY COLUMN {column} TEXT {null_clause}"))
|
||||
+37
@@ -0,0 +1,37 @@
|
||||
"""stock_daily source/adjust 来源与口径标记
|
||||
|
||||
Revision ID: 91c4e27a03fb
|
||||
Revises: 53113c80257f
|
||||
Create Date: 2026-09-06
|
||||
|
||||
新浪兜底行带 source=sina / adjust=qfq 标记;现有 648 万行回填默认
|
||||
tushare / none(SQLite ADD COLUMN 带常量默认值,不重写现有数据)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "91c4e27a03fb"
|
||||
down_revision: str | None = "53113c80257f"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"stock_daily",
|
||||
sa.Column("source", sa.String(length=16), nullable=False, server_default="tushare"),
|
||||
)
|
||||
op.add_column(
|
||||
"stock_daily",
|
||||
sa.Column("adjust", sa.String(length=8), nullable=False, server_default="none"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("stock_daily", "adjust")
|
||||
op.drop_column("stock_daily", "source")
|
||||
+64
@@ -0,0 +1,64 @@
|
||||
"""add stock_name_history table
|
||||
|
||||
Revision ID: a3f8c21d9b47
|
||||
Revises: 7b1c4e9a52d8
|
||||
Create Date: 2026-09-19 21:40:12.000000
|
||||
|
||||
股票名称变更历史(Tushare namechange):时点 ST / 风险警示判定的依据。
|
||||
|
||||
**为什么需要**:`stock.name` 是最新名称快照,用它做 `exclude_st` 会把
|
||||
「曾为高股息、后来才变 ST/退市」的标的在整段历史里都排除 —— 而那正是
|
||||
「股息陷阱」样本。实测对照:同一 spec 仅改 exclude_st,收益 +35.71% → +32.01%,
|
||||
即约 3.70pp 的收益被名称快照口径隐藏。
|
||||
|
||||
幂等键 (symbol, start_date)。时点查询:
|
||||
`start_date <= as_of AND (end_date IS NULL OR end_date >= as_of)`。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = 'a3f8c21d9b47'
|
||||
down_revision: str | None = '7b1c4e9a52d8'
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
'stock_name_history',
|
||||
sa.Column(
|
||||
'id',
|
||||
sa.BigInteger().with_variant(sa.Integer(), 'sqlite'),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column('symbol', sa.String(length=12), nullable=False),
|
||||
sa.Column('name', sa.String(length=64), nullable=False),
|
||||
sa.Column('start_date', sa.Date(), nullable=False),
|
||||
sa.Column('end_date', sa.Date(), nullable=True),
|
||||
sa.Column('ann_date', sa.Date(), nullable=True),
|
||||
sa.Column('change_reason', sa.String(length=32), nullable=True),
|
||||
sa.Column('source', sa.String(length=16), server_default='tushare', nullable=False),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('symbol', 'start_date', name='uq_name_symbol_start'),
|
||||
)
|
||||
with op.batch_alter_table('stock_name_history', schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f('ix_stock_name_history_symbol'), ['symbol'], unique=False)
|
||||
batch_op.create_index(
|
||||
batch_op.f('ix_stock_name_history_start_date'), ['start_date'], unique=False
|
||||
)
|
||||
batch_op.create_index(
|
||||
batch_op.f('ix_stock_name_history_end_date'), ['end_date'], unique=False
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.batch_alter_table('stock_name_history', schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f('ix_stock_name_history_end_date'))
|
||||
batch_op.drop_index(batch_op.f('ix_stock_name_history_start_date'))
|
||||
batch_op.drop_index(batch_op.f('ix_stock_name_history_symbol'))
|
||||
op.drop_table('stock_name_history')
|
||||
+62
@@ -0,0 +1,62 @@
|
||||
"""selection_snapshot / selection_result 表(M6.3 选股落库)
|
||||
|
||||
Revision ID: a6c91d4e7f20
|
||||
Revises: d3f6c9a21b04
|
||||
Create Date: 2026-09-08
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "a6c91d4e7f20"
|
||||
down_revision: str | None = "d3f6c9a21b04"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"selection_snapshot",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("as_of", sa.Date(), nullable=False),
|
||||
sa.Column("method", sa.String(length=16), nullable=False),
|
||||
sa.Column("query_json", sa.Text(), nullable=False),
|
||||
sa.Column("statistics_json", sa.Text(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_selection_snapshot_as_of", "selection_snapshot", ["as_of"])
|
||||
op.create_index("ix_selection_snapshot_created_at", "selection_snapshot", ["created_at"])
|
||||
|
||||
op.create_table(
|
||||
"selection_result",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("selection_id", sa.String(length=32), nullable=False),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("rank", sa.Integer(), nullable=False),
|
||||
sa.Column("score", sa.Numeric(precision=14, scale=6), nullable=False),
|
||||
sa.Column("factor_values_json", sa.Text(), nullable=True),
|
||||
sa.Column("filter_status_json", sa.Text(), nullable=True),
|
||||
sa.Column("reason_json", sa.Text(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_selection_result_selection_id", "selection_result", ["selection_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_selection_result_selection_id", table_name="selection_result")
|
||||
op.drop_table("selection_result")
|
||||
op.drop_index("ix_selection_snapshot_created_at", table_name="selection_snapshot")
|
||||
op.drop_index("ix_selection_snapshot_as_of", table_name="selection_snapshot")
|
||||
op.drop_table("selection_snapshot")
|
||||
+48
@@ -0,0 +1,48 @@
|
||||
"""factor_definition:支持参数化因子实例(2026-10)
|
||||
|
||||
两处改动:
|
||||
1. `name` 64 → 128:参数化实例把参数写进名字
|
||||
(`momentum(window=90,direction=higher_is_better)`),64 位不够留余量。
|
||||
2. 新增 `enabled`:唯一由人配置的字段 —— 是否出现在因子下拉/字段库里。
|
||||
内置实例的开关注仍由代码注册表收敛;停用**不影响**已引用它的策略/归档解析,
|
||||
历史口径不能被开关改义。
|
||||
|
||||
SQLite 不支持直接改列类型,因此用 batch_alter_table(与本仓库既有迁移一致)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "a7c1e4b90f21"
|
||||
down_revision: str | None = "d6e7f8a9b0c1"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
with op.batch_alter_table("factor_definition", schema=None) as batch_op:
|
||||
batch_op.alter_column(
|
||||
"name",
|
||||
existing_type=sa.String(length=64),
|
||||
type_=sa.String(length=128),
|
||||
existing_nullable=False,
|
||||
)
|
||||
op.add_column(
|
||||
"factor_definition",
|
||||
sa.Column("enabled", sa.Boolean(), nullable=False, server_default="1"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("factor_definition", "enabled")
|
||||
with op.batch_alter_table("factor_definition", schema=None) as batch_op:
|
||||
batch_op.alter_column(
|
||||
"name",
|
||||
existing_type=sa.String(length=128),
|
||||
type_=sa.String(length=64),
|
||||
existing_nullable=False,
|
||||
)
|
||||
+113
@@ -0,0 +1,113 @@
|
||||
"""公共配置 + 回测组合表,并把存量策略收敛为「选股条件组合」(2026-09 重构)
|
||||
|
||||
Revision ID: b4c5d6e7f8a9
|
||||
Revises: a3f8c21d9b47
|
||||
Create Date: 2026-09-30
|
||||
|
||||
背景:把原来「一个策略 = 全套参数」拆成三件事 ——
|
||||
1. global_config:费率/印花税/滑点/最低佣金/复权口径/基准(全局唯一一行);
|
||||
2. selection_strategy(复用 strategy 表):只剩股票池 + 因子 + 过滤条件;
|
||||
3. backtest_combo:引用若干选股策略 + 回测参数(资金/持仓数/持仓天数区间/调仓时机/区间)。
|
||||
|
||||
本迁移:
|
||||
- 新建 global_config(并插入默认行)与 backtest_combo 两张表;
|
||||
- 把 strategy.config_json 里**已废弃的回测执行参数键**剥掉(selection / rebalance /
|
||||
costs / portfolio / price_adjustment / *_interval_months),只留 universe/factors/conditions,
|
||||
并把 spec_type 标为 selection。旧数据不丢(归档里的 ResearchSpec 快照原样保留只读)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "b4c5d6e7f8a9"
|
||||
down_revision: str | None = "a3f8c21d9b47"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
# 选股策略不再承载的回测执行参数键(读出时也会被仓储丢弃,这里在存储侧也清掉)
|
||||
_LEGACY_KEYS = (
|
||||
"selection",
|
||||
"rebalance",
|
||||
"costs",
|
||||
"portfolio",
|
||||
"price_adjustment",
|
||||
"selection_interval_months",
|
||||
"rebalance_interval_months",
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ---- 1. 公共配置(单例) ----
|
||||
op.create_table(
|
||||
"global_config",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("commission_rate", sa.Numeric(10, 6), nullable=False, server_default="0.0003"),
|
||||
sa.Column("stamp_tax_rate", sa.Numeric(10, 6), nullable=False, server_default="0.0005"),
|
||||
sa.Column("slippage_rate", sa.Numeric(10, 6), nullable=False, server_default="0.001"),
|
||||
sa.Column("min_commission", sa.Numeric(10, 4), nullable=False, server_default="5"),
|
||||
sa.Column("price_adjustment", sa.String(length=8), nullable=False, server_default="hfq"),
|
||||
sa.Column("benchmark", sa.String(length=16), nullable=False, server_default="000300.SH"),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
# 插入默认行(高股息场景默认后复权 hfq、最低佣金 5 元)
|
||||
op.execute(
|
||||
"INSERT INTO global_config (id, commission_rate, stamp_tax_rate, slippage_rate, "
|
||||
"min_commission, price_adjustment, benchmark) VALUES "
|
||||
"('default', 0.0003, 0.0005, 0.001, 5, 'hfq', '000300.SH')"
|
||||
)
|
||||
|
||||
# ---- 2. 回测组合 ----
|
||||
op.create_table(
|
||||
"backtest_combo",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("name", sa.String(length=64), nullable=False),
|
||||
sa.Column("description", sa.String(length=300), nullable=False, server_default=""),
|
||||
sa.Column("strategy_ids_json", sa.Text(), nullable=False),
|
||||
sa.Column("initial_capital", sa.Numeric(20, 2), nullable=False, server_default="1000000"),
|
||||
sa.Column("hold_count", sa.Integer(), nullable=False, server_default="20"),
|
||||
sa.Column("hold_min_days", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("hold_max_days", sa.Integer(), nullable=True),
|
||||
sa.Column("rebalance_freq", sa.String(length=12), nullable=False, server_default="monthly"),
|
||||
sa.Column("start_date", sa.Date(), nullable=False),
|
||||
sa.Column("end_date", sa.Date(), nullable=False),
|
||||
sa.Column("version", sa.String(length=16), nullable=False, server_default="1"),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("name", name="uq_backtest_combo_name"),
|
||||
)
|
||||
|
||||
# ---- 3. 存量策略 → 选股条件组合(剥掉回测执行参数键) ----
|
||||
conn = op.get_bind()
|
||||
rows = conn.execute(sa.text("SELECT id, config_json FROM strategy")).fetchall()
|
||||
for row_id, cfg_text in rows:
|
||||
try:
|
||||
data = json.loads(cfg_text) if cfg_text else {}
|
||||
except json.JSONDecodeError:
|
||||
continue # 损坏行不动它(读出时仓储也会容错)
|
||||
changed = False
|
||||
for key in _LEGACY_KEYS:
|
||||
if key in data:
|
||||
data.pop(key)
|
||||
changed = True
|
||||
# spec_type 收敛为 selection(旧值多为 backtest)
|
||||
if data.get("spec_type") != "selection":
|
||||
data["spec_type"] = "selection"
|
||||
changed = True
|
||||
if changed:
|
||||
conn.execute(
|
||||
sa.text("UPDATE strategy SET config_json = :cfg, spec_type = 'selection' WHERE id = :id"),
|
||||
{"cfg": json.dumps(data, ensure_ascii=False), "id": row_id},
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# 回滚:删两张新表。strategy.config_json 被剥掉的键无法精确还原
|
||||
# (原始值未备份),故 downgrade 仅撤表结构,不承诺恢复旧策略的完整 config。
|
||||
op.drop_table("backtest_combo")
|
||||
op.drop_table("global_config")
|
||||
+40
@@ -0,0 +1,40 @@
|
||||
"""factor_definition 表(M7.1 因子定义入库)
|
||||
|
||||
Revision ID: b7f2a5e81c33
|
||||
Revises: a6c91d4e7f20
|
||||
Create Date: 2026-09-09
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "b7f2a5e81c33"
|
||||
down_revision: str | None = "a6c91d4e7f20"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"factor_definition",
|
||||
sa.Column("name", sa.String(length=64), nullable=False),
|
||||
sa.Column("description", sa.String(length=500), nullable=False),
|
||||
sa.Column("formula", sa.String(length=500), nullable=False),
|
||||
sa.Column("brief", sa.String(length=500), nullable=False),
|
||||
sa.Column("frequency", sa.String(length=16), nullable=False),
|
||||
sa.Column("lookback", sa.Integer(), nullable=False),
|
||||
sa.Column("direction", sa.String(length=32), nullable=False),
|
||||
sa.Column("requires_json", sa.Text(), nullable=False),
|
||||
sa.Column("version", sa.String(length=16), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("name"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("factor_definition")
|
||||
+37
@@ -0,0 +1,37 @@
|
||||
"""factor_composite 表(M7.2b 因子组合保存/复用)
|
||||
|
||||
Revision ID: c3e9a0d1f4b5
|
||||
Revises: b7f2a5e81c33
|
||||
Create Date: 2026-09-09
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "c3e9a0d1f4b5"
|
||||
down_revision: str | None = "b7f2a5e81c33"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"factor_composite",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("name", sa.String(length=64), nullable=False),
|
||||
sa.Column("method", sa.String(length=16), nullable=False),
|
||||
sa.Column("description", sa.String(length=300), nullable=False),
|
||||
sa.Column("components_json", sa.Text(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("name", name="uq_factor_composite_name"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("factor_composite")
|
||||
+155
@@ -0,0 +1,155 @@
|
||||
"""重算存量选股策略的过时 description(2026-09 重构收尾)
|
||||
|
||||
Revision ID: c5d6e7f8a9b0
|
||||
Revises: b4c5d6e7f8a9
|
||||
Create Date: 2026-10-01
|
||||
|
||||
背景:b4c5d6e7f8a9 把一个策略的 config_json 里回测执行参数剥掉了,但**没有**重算
|
||||
`strategy.description`。旧描述是重构前由 describe_strategy 从「全套参数」自动生成的,
|
||||
于是策略库里会出现这种自相矛盾的说明:
|
||||
|
||||
「…每 6 个月重新择股、每 6 个月调仓,后复权口径、按调仓日收盘价成交
|
||||
(含佣金 0.03%/印花税 0.05%/滑点 0.1%)。」
|
||||
|
||||
而选股策略现在**不再持有**调仓/成本/复权,这些由「回测组合 + 公共配置」在回测时决定。
|
||||
本迁移用当前口径的 describe_strategy(纯函数,无 IO/DB)重算这些陈旧说明。
|
||||
|
||||
安全性 —— 只改「可证明是旧自动生成」的行,不碰人工撰写的说明:
|
||||
1. description 为空/纯空白(API 保存契约要求必须有说明,空值必然是历史遗留)→ 补全;
|
||||
2. description 含旧自动文案独有的回测执行词(佣金/印花税/滑点/调仓/择股/复权口径/
|
||||
收盘价成交/最低佣金/初始资金)→ 重算。新口径的说明**绝不会**出现这些词
|
||||
(见 strategy_doc._describe_selection_only),因此命中即旧自动文案。
|
||||
其余行原样保留(kept)。无法解析/校验失败的行跳过并打印告警,绝不静默改写。
|
||||
|
||||
为什么在迁移里 import 应用代码:说明文本的唯一事实来源就是 `describe_strategy`
|
||||
(AGENT.md §24:不许另写一份近似文案)。自己复制一份文案逻辑才是真正的漂移风险。
|
||||
代价是该迁移的产物依赖当时的代码版本 —— 对「一次性回填存量说明」这个用途可以接受,
|
||||
且新库 upgrade 时 strategy 表为空、不受影响。
|
||||
|
||||
downgrade 仅回滚结构层面:**不恢复**被重算的旧说明(原文未备份),因此不可逆。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "c5d6e7f8a9b0"
|
||||
down_revision: str | None = "b4c5d6e7f8a9"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
# 与 strategy.description 列宽一致(StrategyModel.description = String(300))
|
||||
_DESCRIPTION_MAX_CHARS = 300
|
||||
|
||||
# 选股策略不再承载的键(与 b4c5d6e7f8a9 一致;历史行可能仍残留)
|
||||
_LEGACY_KEYS = (
|
||||
"selection",
|
||||
"rebalance",
|
||||
"costs",
|
||||
"portfolio",
|
||||
"price_adjustment",
|
||||
"selection_interval_months",
|
||||
"rebalance_interval_months",
|
||||
)
|
||||
|
||||
# 由 DB 列承载、不应从 config_json 再喂给实体的键
|
||||
_COLUMN_KEYS = ("id", "name", "description", "spec_type", "version")
|
||||
|
||||
# 旧「全套参数」自动文案独有的回测执行词 —— 新口径说明不会出现(命中即认定陈旧)
|
||||
_LEGACY_MARKERS = (
|
||||
"佣金",
|
||||
"印花税",
|
||||
"滑点",
|
||||
"调仓",
|
||||
"择股",
|
||||
"复权口径",
|
||||
"收盘价成交",
|
||||
"最低佣金",
|
||||
"初始资金",
|
||||
)
|
||||
|
||||
_LEGACY_MARKER_RE = re.compile("|".join(_LEGACY_MARKERS))
|
||||
|
||||
|
||||
def _truncate(text: str) -> str:
|
||||
"""与 API 的说明补全同口径:超列宽按字符截断并显式加省略号。"""
|
||||
if len(text) <= _DESCRIPTION_MAX_CHARS:
|
||||
return text
|
||||
return text[: _DESCRIPTION_MAX_CHARS - 1] + "…"
|
||||
|
||||
|
||||
def _derive_summary(name: str, description: str, data: dict) -> str | None:
|
||||
"""按当前口径重算一句话说明;无法构造实体时返回 None(调用方跳过并告警)。"""
|
||||
# 延迟 import:保持迁移模块导入轻量,且让 alembic env 先完成自身引导。
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.quant.strategy_doc import describe_strategy
|
||||
|
||||
payload = dict(data)
|
||||
for key in _COLUMN_KEYS + _LEGACY_KEYS:
|
||||
payload.pop(key, None)
|
||||
try:
|
||||
st = SelectionStrategy(name=name, description=description, **payload)
|
||||
except Exception as exc: # noqa: BLE001 —— 逐行容错:坏行跳过并告警,不阻断整次迁移
|
||||
print(f"[refresh-strategy-docs] 跳过无法解析的策略 {name!r}: {exc}", flush=True)
|
||||
return None
|
||||
return describe_strategy(st).summary
|
||||
|
||||
|
||||
def _is_stale(name: str, description: str) -> bool:
|
||||
return (not (description or "").strip()) or bool(_LEGACY_MARKER_RE.search(description or ""))
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
rows = conn.execute(
|
||||
sa.text("SELECT id, name, description, config_json FROM strategy")
|
||||
).fetchall()
|
||||
|
||||
rewritten = kept = broken = 0
|
||||
for row_id, name, description, cfg_text in rows:
|
||||
try:
|
||||
data = json.loads(cfg_text) if cfg_text else {}
|
||||
except json.JSONDecodeError:
|
||||
broken += 1
|
||||
print(f"[refresh-strategy-docs] 跳过 config_json 损坏的策略 {row_id}", flush=True)
|
||||
continue
|
||||
if not isinstance(data, dict):
|
||||
broken += 1
|
||||
print(f"[refresh-strategy-docs] 跳过 config_json 非对象的策略 {row_id}", flush=True)
|
||||
continue
|
||||
if not _is_stale(name, description):
|
||||
kept += 1 # 人工撰写的说明:不动它
|
||||
continue
|
||||
summary = _derive_summary(name, description or "", data)
|
||||
if summary is None:
|
||||
broken += 1
|
||||
continue
|
||||
summary = _truncate(summary)
|
||||
if summary == (description or ""):
|
||||
kept += 1
|
||||
continue
|
||||
conn.execute(
|
||||
sa.text("UPDATE strategy SET description = :desc WHERE id = :id"),
|
||||
{"desc": summary, "id": row_id},
|
||||
)
|
||||
rewritten += 1
|
||||
|
||||
print(
|
||||
f"[refresh-strategy-docs] 重算 {rewritten} 条陈旧/空说明,"
|
||||
f"保留 {kept} 条,跳过 {broken} 条异常行(共 {len(rows)} 条)",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# 旧说明原文未备份,无法还原:回滚只表示「结构层面无事可做」。
|
||||
# 显式空实现(而非 pass 无说明),避免读者误以为会恢复文案。
|
||||
print(
|
||||
"[refresh-strategy-docs] downgrade:被重算的说明不可还原(原文未备份),不执行任何写操作",
|
||||
flush=True,
|
||||
)
|
||||
+33
@@ -0,0 +1,33 @@
|
||||
"""financial_indicator 增加 source 来源标记
|
||||
|
||||
Revision ID: d3f6c9a21b04
|
||||
Revises: 91c4e27a03fb
|
||||
Create Date: 2026-09-08
|
||||
|
||||
新浪校验兜底导入的财务行带 source=sina(字段可能不全),与 Tushare
|
||||
首选行区分;现有行回填默认 tushare(SQLite ADD COLUMN 带常量默认值,
|
||||
不重写现有数据)。AGENT.md §7 数据来源可追溯。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "d3f6c9a21b04"
|
||||
down_revision: str | None = "91c4e27a03fb"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"financial_indicator",
|
||||
sa.Column("source", sa.String(length=16), nullable=False, server_default="tushare"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("financial_indicator", "source")
|
||||
+45
@@ -0,0 +1,45 @@
|
||||
"""condition_field 表(2026-10 字段库:过滤条件字段目录入库)
|
||||
|
||||
Revision ID: d6e7f8a9b0c1
|
||||
Revises: c5d6e7f8a9b0
|
||||
Create Date: 2026-10-01
|
||||
|
||||
背景:策略库的过滤条件此前只能手填字段名(dv_ratio / static.industry …),
|
||||
用户看不到含义、写错也不报错(未知字段求值恒为 None,条件永远不通过)。
|
||||
本表存放字段库目录:内置字段由 quant/condition_fields.py 注册表在 API 首次读取时
|
||||
seed(只补不删,不覆盖用户改过的文案),自定义字段与停用状态也落在本表。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "d6e7f8a9b0c1"
|
||||
down_revision: str | None = "c5d6e7f8a9b0"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"condition_field",
|
||||
sa.Column("name", sa.String(length=64), nullable=False),
|
||||
sa.Column("label", sa.String(length=64), nullable=False),
|
||||
sa.Column("description", sa.String(length=500), nullable=False),
|
||||
sa.Column("kind", sa.String(length=8), nullable=False),
|
||||
sa.Column("group_name", sa.String(length=32), nullable=False),
|
||||
sa.Column("unit", sa.String(length=16), nullable=False),
|
||||
sa.Column("source", sa.String(length=8), nullable=False),
|
||||
sa.Column("enabled", sa.Boolean(), nullable=False),
|
||||
sa.Column("sort_order", sa.Integer(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("name"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("condition_field")
|
||||
+62
@@ -0,0 +1,62 @@
|
||||
"""signal_snapshot / signal_event 表(M8.1 交易信号)
|
||||
|
||||
Revision ID: d8e0b2f3c4d5
|
||||
Revises: c3e9a0d1f4b5
|
||||
Create Date: 2026-09-09
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "d8e0b2f3c4d5"
|
||||
down_revision: str | None = "c3e9a0d1f4b5"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"signal_snapshot",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("as_of", sa.Date(), nullable=False),
|
||||
sa.Column("query_json", sa.Text(), nullable=False),
|
||||
sa.Column("rules_json", sa.Text(), nullable=False),
|
||||
sa.Column("statistics_json", sa.Text(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_signal_snapshot_as_of", "signal_snapshot", ["as_of"])
|
||||
op.create_index("ix_signal_snapshot_created_at", "signal_snapshot", ["created_at"])
|
||||
|
||||
op.create_table(
|
||||
"signal_event",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("signal_id", sa.String(length=32), nullable=False),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("signal_date", sa.Date(), nullable=False),
|
||||
sa.Column("signal_type", sa.String(length=8), nullable=False),
|
||||
sa.Column("score", sa.Numeric(precision=14, scale=6), nullable=True),
|
||||
sa.Column("price", sa.Numeric(precision=14, scale=4), nullable=True),
|
||||
sa.Column("reason_json", sa.Text(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_signal_event_signal_id", "signal_event", ["signal_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_signal_event_signal_id", table_name="signal_event")
|
||||
op.drop_table("signal_event")
|
||||
op.drop_index("ix_signal_snapshot_created_at", table_name="signal_snapshot")
|
||||
op.drop_index("ix_signal_snapshot_as_of", table_name="signal_snapshot")
|
||||
op.drop_table("signal_snapshot")
|
||||
+38
@@ -0,0 +1,38 @@
|
||||
"""strategy 表(M8.3 策略持久化)
|
||||
|
||||
Revision ID: e1f2a3b4c5d6
|
||||
Revises: d8e0b2f3c4d5
|
||||
Create Date: 2026-09-09
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "e1f2a3b4c5d6"
|
||||
down_revision: str | None = "d8e0b2f3c4d5"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"strategy",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("name", sa.String(length=64), nullable=False),
|
||||
sa.Column("description", sa.String(length=300), nullable=False),
|
||||
sa.Column("spec_type", sa.String(length=16), nullable=False),
|
||||
sa.Column("config_json", sa.Text(), nullable=False),
|
||||
sa.Column("version", sa.String(length=16), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("name", name="uq_strategy_name"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("strategy")
|
||||
+178
@@ -0,0 +1,178 @@
|
||||
"""phase1 market data tables
|
||||
|
||||
Revision ID: e4d188250fb2
|
||||
Revises:
|
||||
Create Date: 2026-09-06 16:58:13.904265
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "e4d188250fb2"
|
||||
down_revision: str | None = None
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_table(
|
||||
"adjust_factor",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("factor", sa.Numeric(precision=20, scale=6), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("symbol", "trade_date", name="uq_adj_symbol_date"),
|
||||
)
|
||||
with op.batch_alter_table("adjust_factor", schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f("ix_adjust_factor_symbol"), ["symbol"], unique=False)
|
||||
batch_op.create_index(
|
||||
batch_op.f("ix_adjust_factor_trade_date"), ["trade_date"], unique=False
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"financial_indicator",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("report_date", sa.Date(), nullable=False),
|
||||
sa.Column("announce_date", sa.Date(), nullable=False),
|
||||
sa.Column("eps", sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column("roe", sa.Numeric(precision=10, scale=4), nullable=True),
|
||||
sa.Column("total_revenue", sa.Numeric(precision=24, scale=2), nullable=True),
|
||||
sa.Column("net_profit", sa.Numeric(precision=24, scale=2), nullable=True),
|
||||
sa.Column("gross_margin", sa.Numeric(precision=10, scale=4), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("symbol", "report_date", "announce_date", name="uq_fin_sym_rep_ann"),
|
||||
)
|
||||
with op.batch_alter_table("financial_indicator", schema=None) as batch_op:
|
||||
batch_op.create_index(
|
||||
batch_op.f("ix_financial_indicator_announce_date"), ["announce_date"], unique=False
|
||||
)
|
||||
batch_op.create_index(
|
||||
batch_op.f("ix_financial_indicator_report_date"), ["report_date"], unique=False
|
||||
)
|
||||
batch_op.create_index(batch_op.f("ix_financial_indicator_symbol"), ["symbol"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"stock",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("name", sa.String(length=64), nullable=False),
|
||||
sa.Column("industry", sa.String(length=64), nullable=True),
|
||||
sa.Column("area", sa.String(length=32), nullable=True),
|
||||
sa.Column("market", sa.String(length=16), nullable=True),
|
||||
sa.Column("exchange", sa.String(length=8), nullable=True),
|
||||
sa.Column("list_date", sa.Date(), nullable=False),
|
||||
sa.Column("delist_date", sa.Date(), nullable=True),
|
||||
sa.Column("status", sa.String(length=8), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
with op.batch_alter_table("stock", schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f("ix_stock_symbol"), ["symbol"], unique=True)
|
||||
|
||||
op.create_table(
|
||||
"stock_daily",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("open", sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column("high", sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column("low", sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column("close", sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column("volume", sa.Numeric(precision=24, scale=2), nullable=True),
|
||||
sa.Column("amount", sa.Numeric(precision=24, scale=2), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("symbol", "trade_date", name="uq_daily_symbol_date"),
|
||||
)
|
||||
with op.batch_alter_table("stock_daily", schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f("ix_stock_daily_symbol"), ["symbol"], unique=False)
|
||||
batch_op.create_index(batch_op.f("ix_stock_daily_trade_date"), ["trade_date"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"sync_log",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("source", sa.String(length=16), nullable=False),
|
||||
sa.Column("api", sa.String(length=32), nullable=False),
|
||||
sa.Column("request_time", sa.DateTime(), nullable=False),
|
||||
sa.Column("success", sa.Boolean(), nullable=False),
|
||||
sa.Column("failure_reason", sa.String(length=500), nullable=True),
|
||||
sa.Column("row_count", sa.Integer(), nullable=False),
|
||||
sa.Column("data_start", sa.Date(), nullable=True),
|
||||
sa.Column("data_end", sa.Date(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
with op.batch_alter_table("sync_log", schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f("ix_sync_log_source"), ["source"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"trading_calendar",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("calendar_date", sa.Date(), nullable=False),
|
||||
sa.Column("is_open", sa.Boolean(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
with op.batch_alter_table("trading_calendar", schema=None) as batch_op:
|
||||
batch_op.create_index(
|
||||
batch_op.f("ix_trading_calendar_calendar_date"), ["calendar_date"], unique=True
|
||||
)
|
||||
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
with op.batch_alter_table("trading_calendar", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_trading_calendar_calendar_date"))
|
||||
|
||||
op.drop_table("trading_calendar")
|
||||
with op.batch_alter_table("sync_log", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_sync_log_source"))
|
||||
|
||||
op.drop_table("sync_log")
|
||||
with op.batch_alter_table("stock_daily", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_stock_daily_trade_date"))
|
||||
batch_op.drop_index(batch_op.f("ix_stock_daily_symbol"))
|
||||
|
||||
op.drop_table("stock_daily")
|
||||
with op.batch_alter_table("stock", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_stock_symbol"))
|
||||
|
||||
op.drop_table("stock")
|
||||
with op.batch_alter_table("financial_indicator", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_financial_indicator_symbol"))
|
||||
batch_op.drop_index(batch_op.f("ix_financial_indicator_report_date"))
|
||||
batch_op.drop_index(batch_op.f("ix_financial_indicator_announce_date"))
|
||||
|
||||
op.drop_table("financial_indicator")
|
||||
with op.batch_alter_table("adjust_factor", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_adjust_factor_trade_date"))
|
||||
batch_op.drop_index(batch_op.f("ix_adjust_factor_symbol"))
|
||||
|
||||
op.drop_table("adjust_factor")
|
||||
# ### end Alembic commands ###
|
||||
+46
@@ -0,0 +1,46 @@
|
||||
"""index_weight 表(B1 指数历史成分)
|
||||
|
||||
Revision ID: f5e0d1c2b3a4
|
||||
Revises: e1f2a3b4c5d6
|
||||
Create Date: 2026-09-09
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "f5e0d1c2b3a4"
|
||||
down_revision: str | None = "e1f2a3b4c5d6"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"index_weight",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("index_code", sa.String(length=12), nullable=False),
|
||||
sa.Column("index_name", sa.String(length=64), nullable=True),
|
||||
sa.Column("trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("weight", sa.Numeric(precision=10, scale=6), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("index_code", "trade_date", "symbol", name="uq_idxw_code_date_sym"),
|
||||
)
|
||||
op.create_index("ix_index_weight_index_code", "index_weight", ["index_code"])
|
||||
op.create_index("ix_index_weight_trade_date", "index_weight", ["trade_date"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_index_weight_trade_date", table_name="index_weight")
|
||||
op.drop_index("ix_index_weight_index_code", table_name="index_weight")
|
||||
op.drop_table("index_weight")
|
||||
@@ -0,0 +1,129 @@
|
||||
"""数据集快照指纹:`experiment.data_version` 的真实取值来源。
|
||||
|
||||
归档的用途是「往复查看 / 复现」,因此除了代码版本(code_version)还必须记录
|
||||
**当时的数据口径**:数据同步到哪一天、各表规模多大。此前该字段从未被写入
|
||||
(永远 NULL),本模块补上。
|
||||
|
||||
诚实性与成本(AGENT.md §7 不静默 / §24 不假装支持):
|
||||
- 交易日 `MAX(trade_date)` 走索引,实测 ~0.2ms,**真实值**;
|
||||
- 大表(stock_daily / adjust_factor / daily_basic,各 800 万行量级)的全表
|
||||
`COUNT(*)` 实测单次 ~1.2s,三次合计 ~3.6s —— 不允许出现在请求路径上;
|
||||
故 MySQL 下改用 `information_schema.TABLES.TABLE_ROWS`(单次查询实测 ~0.5ms),
|
||||
它是 InnoDB 统计缓存的**近似值**(实测 stock_daily 7688126 vs 真实 COUNT(*)
|
||||
8052698,偏差 ~4.6%),因此字符串里用 `≈` 明确标注为近似,绝不冒充精确计数。
|
||||
- 非 MySQL 方言(测试用的 SQLite 等)数据量小,直接 `COUNT(*)` 得到精确值,
|
||||
不带 `≈` 标记;
|
||||
- 任何一段取不到(连接失败 / 表不存在 / 权限不足)都**降级**:要么丢弃该段,
|
||||
要么整体返回 `unavailable`,绝不编造数字。
|
||||
|
||||
字符串格式(≤40 字符 —— experiment.data_version 列是 varchar(40),本文件不新增
|
||||
迁移,故必须塞得下;段超长时从右往左丢弃低优先级段,再退化到 `d<date>`):
|
||||
|
||||
d<YYYYMMDD> stock_daily 的最大交易日(`d-` 表示该表无数据 / 取不到交易日)
|
||||
n<count> stock_daily 行数
|
||||
a<count> adjust_factor 行数
|
||||
b<count> daily_basic 行数
|
||||
|
||||
行数段写法:`≈<N>k` = MySQL 近似值(k = 千行,四舍五入);`<N>` = 精确值。
|
||||
示例:`d20260904;n≈8053k;a≈8189k;b≈7718k`
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from sqlalchemy import func, select, text
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import (
|
||||
AdjustFactorModel,
|
||||
DailyBasicModel,
|
||||
StockDailyModel,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# experiment.data_version 列宽(varchar(40),见 models/jobs.py + 既有迁移)
|
||||
DATA_VERSION_MAX_LEN = 40
|
||||
|
||||
# 段定义:(段名, model),顺序即优先级(越靠前越重要)
|
||||
_TABLES: tuple[tuple[str, type], ...] = (
|
||||
("n", StockDailyModel),
|
||||
("a", AdjustFactorModel),
|
||||
("b", DailyBasicModel),
|
||||
)
|
||||
|
||||
_FALLBACK = "unavailable"
|
||||
|
||||
|
||||
def _fmt_rows(rows: int | None, *, approx: bool) -> str | None:
|
||||
"""行数段:近似值用「≈Nk」(千行),精确值用原样数字。"""
|
||||
if rows is None:
|
||||
return None
|
||||
if not approx:
|
||||
return str(rows)
|
||||
return f"≈{round(rows / 1000)}k"
|
||||
|
||||
|
||||
def _latest_trade_date(session: Session) -> str:
|
||||
"""stock_daily 最大交易日 → `YYYYMMDD`;无数据 / 取不到 → `-`。"""
|
||||
try:
|
||||
day = session.execute(select(func.max(StockDailyModel.trade_date))).scalar()
|
||||
except Exception as exc: # noqa: BLE001 —— 指纹取不到必须降级,不影响归档主体
|
||||
logger.warning("data_version: MAX(trade_date) 取不到:%s: %s", type(exc).__name__, exc)
|
||||
return "-"
|
||||
return day.strftime("%Y%m%d") if day is not None else "-"
|
||||
|
||||
|
||||
def _approx_rows_mysql(session: Session) -> dict[str, int]:
|
||||
"""MySQL:一次 information_schema 查询拿全部表的近似行数(实测 ~0.5ms)。"""
|
||||
try:
|
||||
rows = session.execute(
|
||||
text(
|
||||
"SELECT TABLE_NAME, TABLE_ROWS FROM information_schema.TABLES "
|
||||
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME IN :names"
|
||||
).bindparams(names=tuple(m.__tablename__ for _, m in _TABLES))
|
||||
).all()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("data_version: information_schema 查询失败:%s: %s", type(exc).__name__, exc)
|
||||
return {}
|
||||
return {str(name): int(n) for name, n in rows if n is not None}
|
||||
|
||||
|
||||
def _exact_rows(session: Session, model: type) -> int | None:
|
||||
"""其它方言(SQLite 等,数据量小):精确 COUNT(*)。"""
|
||||
try:
|
||||
return int(session.execute(select(func.count()).select_from(model)).scalar() or 0)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("data_version: COUNT(*) 失败:%s: %s", type(exc).__name__, exc)
|
||||
return None
|
||||
|
||||
|
||||
def compute_data_version(session: Session) -> str:
|
||||
"""计算数据集快照指纹(格式与取舍见模块 docstring)。失败降级为 `unavailable`。"""
|
||||
try:
|
||||
dialect = session.get_bind().dialect.name
|
||||
except Exception: # noqa: BLE001
|
||||
dialect = ""
|
||||
|
||||
approx = dialect == "mysql"
|
||||
approx_rows = _approx_rows_mysql(session) if approx else {}
|
||||
|
||||
# 首段恒为交易日段:`d-` 本身也是真实信息(表为空 / 取不到交易日),保留
|
||||
segments: list[str] = [f"d{_latest_trade_date(session)}"]
|
||||
for name, model in _TABLES:
|
||||
rows = (
|
||||
approx_rows.get(model.__tablename__) if approx else _exact_rows(session, model)
|
||||
)
|
||||
seg = _fmt_rows(rows, approx=approx)
|
||||
if seg is not None:
|
||||
segments.append(f"{name}{seg}")
|
||||
|
||||
# 列宽护栏:超 40 字符时从右往左丢弃低优先级段(保留的段仍是真实值,绝不截断数字)
|
||||
while len(segments) > 1 and len(";".join(segments)) > DATA_VERSION_MAX_LEN:
|
||||
segments.pop()
|
||||
|
||||
if len(segments) == 1 and segments[0] == "d-":
|
||||
# 三张表连行数都读不到、交易日也没有:如实标记「不可用」,不编造
|
||||
return _FALLBACK
|
||||
return ";".join(segments)
|
||||
@@ -3,3 +3,45 @@
|
||||
新增表流程(AGENT.md §12):Model → Alembic Migration → Test。
|
||||
模型统一继承 infra.persistence.sqlalchemy.base.Base。
|
||||
"""
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.models.combo import ( # noqa: F401
|
||||
BacktestComboModel,
|
||||
GlobalConfigModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.composite import ( # noqa: F401
|
||||
FactorCompositeModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.condition_field import ( # noqa: F401
|
||||
ConditionFieldModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.factor import ( # noqa: F401
|
||||
FactorDefinitionModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.index import ( # noqa: F401
|
||||
IndexWeightModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.jobs import ( # noqa: F401
|
||||
ExperimentModel,
|
||||
JobModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import ( # noqa: F401
|
||||
AdjustFactorModel,
|
||||
DailyBasicModel,
|
||||
FinancialIndicatorModel,
|
||||
StockDailyModel,
|
||||
StockModel,
|
||||
StockNameHistoryModel,
|
||||
SyncLogModel,
|
||||
TradingCalendarModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.selection import ( # noqa: F401
|
||||
SelectionResultModel,
|
||||
SelectionSnapshotModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.signal import ( # noqa: F401
|
||||
SignalEventModel,
|
||||
SignalSnapshotModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.strategy import ( # noqa: F401
|
||||
StrategyModel,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
"""公共配置 + 回测组合表(2026-09 重构)。
|
||||
|
||||
- `global_config`:全局唯一一行(id="default"),存费率/滑点/最低佣金/复权口径/基准。
|
||||
- `backtest_combo`:回测组合,引用若干选股策略(strategy_ids JSON)+ 回测参数
|
||||
(资金/持仓数/持仓天数区间/调仓时机/区间)。费率与复权不在此表 —— 运行时从
|
||||
global_config 快照进归档的 config_snapshot,保证可复现。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
from sqlalchemy import Date, DateTime, Integer, Numeric, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
|
||||
class GlobalConfigModel(Base):
|
||||
__tablename__ = "global_config"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True) # 恒为 "default"
|
||||
commission_rate: Mapped[float] = mapped_column(Numeric(10, 6), default=0.0003)
|
||||
stamp_tax_rate: Mapped[float] = mapped_column(Numeric(10, 6), default=0.0005)
|
||||
slippage_rate: Mapped[float] = mapped_column(Numeric(10, 6), default=0.001)
|
||||
min_commission: Mapped[float] = mapped_column(Numeric(10, 4), default=5.0)
|
||||
price_adjustment: Mapped[str] = mapped_column(String(8), default="hfq")
|
||||
benchmark: Mapped[str] = mapped_column(String(16), default="000300.SH")
|
||||
updated_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
|
||||
|
||||
|
||||
class BacktestComboModel(Base):
|
||||
__tablename__ = "backtest_combo"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
name: Mapped[str] = mapped_column(String(64), unique=True)
|
||||
description: Mapped[str] = mapped_column(String(300), default="")
|
||||
strategy_ids_json: Mapped[str] = mapped_column(Text) # JSON list[str]
|
||||
initial_capital: Mapped[float] = mapped_column(Numeric(20, 2), default=1_000_000.0)
|
||||
hold_count: Mapped[int] = mapped_column(Integer, default=20)
|
||||
hold_min_days: Mapped[int] = mapped_column(Integer, default=0)
|
||||
hold_max_days: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
rebalance_freq: Mapped[str] = mapped_column(String(12), default="monthly")
|
||||
start_date: Mapped[date] = mapped_column(Date)
|
||||
end_date: Mapped[date] = mapped_column(Date)
|
||||
version: Mapped[str] = mapped_column(String(16), default="1")
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -0,0 +1,25 @@
|
||||
"""因子组合表(M7.2b)。
|
||||
|
||||
factor_composite:可保存/复用的因子组合定义;components 以 JSON 存
|
||||
(组件方向冗余快照自因子注册表,计算时以注册表为准)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
|
||||
class FactorCompositeModel(Base):
|
||||
__tablename__ = "factor_composite"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
name: Mapped[str] = mapped_column(String(64), unique=True)
|
||||
method: Mapped[str] = mapped_column(String(16), default="fixed")
|
||||
description: Mapped[str] = mapped_column(String(300), default="")
|
||||
components_json: Mapped[str] = mapped_column(Text)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -0,0 +1,31 @@
|
||||
"""字段库表(2026-10)。
|
||||
|
||||
condition_field:过滤条件字段的目录契约源(name 主键幂等)。
|
||||
内置字段由 quant/condition_fields.py 的注册表在 API 首次读取时 seed(只补不删),
|
||||
自定义字段与用户改过的文案都落在这张表里。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, Integer, String
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
|
||||
class ConditionFieldModel(Base):
|
||||
__tablename__ = "condition_field"
|
||||
|
||||
name: Mapped[str] = mapped_column(String(64), primary_key=True)
|
||||
label: Mapped[str] = mapped_column(String(64), default="")
|
||||
description: Mapped[str] = mapped_column(String(500), default="")
|
||||
kind: Mapped[str] = mapped_column(String(8), default="num")
|
||||
group_name: Mapped[str] = mapped_column(String(32), default="行情")
|
||||
unit: Mapped[str] = mapped_column(String(16), default="")
|
||||
source: Mapped[str] = mapped_column(String(8), default="builtin")
|
||||
enabled: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
sort_order: Mapped[int] = mapped_column(Integer, default=100)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -0,0 +1,33 @@
|
||||
"""因子目录表(M7.1;2026-10 支持参数化实例)。
|
||||
|
||||
factor_definition:因子名(主键,参数化实例的参数就写在名字里)→ 元数据;requires 以 JSON 存。
|
||||
`enabled` 是唯一由人配置的字段(是否出现在因子下拉里);内置实例的开关注由代码注册表
|
||||
收敛(见 application/services/factor_catalog.py),停用只影响「能不能被选中」,
|
||||
不影响已引用它的策略/归档解析 —— 历史不能被开关改义。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, Integer, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
|
||||
class FactorDefinitionModel(Base):
|
||||
__tablename__ = "factor_definition"
|
||||
|
||||
# 128:参数化实例的名字把参数写全(如 momentum(window=90,direction=lower_is_better))
|
||||
name: Mapped[str] = mapped_column(String(128), primary_key=True)
|
||||
description: Mapped[str] = mapped_column(String(500), default="")
|
||||
formula: Mapped[str] = mapped_column(String(500), default="")
|
||||
brief: Mapped[str] = mapped_column(String(500), default="")
|
||||
frequency: Mapped[str] = mapped_column(String(16), default="daily")
|
||||
lookback: Mapped[int] = mapped_column(Integer, default=20)
|
||||
direction: Mapped[str] = mapped_column(String(32), default="higher_is_better")
|
||||
requires_json: Mapped[str] = mapped_column(Text, default="[]")
|
||||
version: Mapped[str] = mapped_column(String(16), default="1")
|
||||
enabled: Mapped[bool] = mapped_column(Boolean, default=True, server_default="1")
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -0,0 +1,31 @@
|
||||
"""指数历史成分表(B1)。
|
||||
|
||||
index_weight:指数成分快照(index_code, trade_date, symbol 唯一)。
|
||||
历史成分查询(members_at)取 <= as_of 最近一期快照 —— Survivorship-free Universe。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
from sqlalchemy import BigInteger, Date, Integer, Numeric, String, UniqueConstraint
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
PK_INT = BigInteger().with_variant(Integer, "sqlite")
|
||||
|
||||
|
||||
class IndexWeightModel(Base):
|
||||
__tablename__ = "index_weight"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("index_code", "trade_date", "symbol", name="uq_idxw_code_date_sym"),
|
||||
)
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
index_code: Mapped[str] = mapped_column(String(12), index=True)
|
||||
index_name: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
trade_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
symbol: Mapped[str] = mapped_column(String(12))
|
||||
weight: Mapped[Decimal | None] = mapped_column(Numeric(10, 6), nullable=True)
|
||||
@@ -0,0 +1,51 @@
|
||||
"""Phase 4:Job(异步任务)与 Experiment(实验归档)表。
|
||||
|
||||
`result_json` 使用 MEDIUMTEXT(MySQL 上限 16MB):全市场多年回测的结果含
|
||||
净值/回撤曲线、逐笔成交、Signal↔Fill 记录与个股收益曲线,实测可达数 MB,
|
||||
MySQL `TEXT`(64KB)会直接报 1406 Data too long 导致 Job 归档失败
|
||||
(2026-09 高股息案例实测:约 1.2MB → 落库失败)。
|
||||
SQLite 不区分 TEXT 长度,故模型层统一用 MEDIUMTEXT.with_variant 保持跨库可用。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, String, Text
|
||||
from sqlalchemy.dialects.mysql import MEDIUMTEXT
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
# 长 JSON 列类型:MySQL 用 MEDIUMTEXT(16MB),其它方言退化为 TEXT(SQLite 无长度限制)
|
||||
_LONG_JSON = Text().with_variant(MEDIUMTEXT(), "mysql")
|
||||
|
||||
|
||||
class JobModel(Base):
|
||||
__tablename__ = "job"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
kind: Mapped[str] = mapped_column(String(16))
|
||||
status: Mapped[str] = mapped_column(String(12), index=True)
|
||||
stage: Mapped[str | None] = mapped_column(String(24), nullable=True)
|
||||
spec_json: Mapped[str] = mapped_column(Text)
|
||||
error: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
result_json: Mapped[str | None] = mapped_column(_LONG_JSON, nullable=True)
|
||||
experiment_id: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
started_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
|
||||
finished_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
|
||||
|
||||
|
||||
class ExperimentModel(Base):
|
||||
__tablename__ = "experiment"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
kind: Mapped[str] = mapped_column(String(16))
|
||||
spec_json: Mapped[str] = mapped_column(Text)
|
||||
result_json: Mapped[str] = mapped_column(_LONG_JSON)
|
||||
summary_text: Mapped[str | None] = mapped_column(String(200), nullable=True)
|
||||
code_version: Mapped[str | None] = mapped_column(String(40), nullable=True)
|
||||
data_version: Mapped[str | None] = mapped_column(String(40), nullable=True)
|
||||
job_id: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -0,0 +1,171 @@
|
||||
"""Phase 1 市场数据表模型(SQLAlchemy 2.x 声明式)。
|
||||
|
||||
列名与 domain.entities.market 字段一一对应,便于 Repository 双向映射。
|
||||
Decimal 字段用 Numeric:SQLite 以浮点近似存储,未来 MySQL 下精确。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
|
||||
from sqlalchemy import (
|
||||
BigInteger,
|
||||
Boolean,
|
||||
Date,
|
||||
DateTime,
|
||||
Integer,
|
||||
Numeric,
|
||||
String,
|
||||
UniqueConstraint,
|
||||
)
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
# SQLite 只对 INTEGER PRIMARY KEY 自增;MySQL 下用 BIGINT
|
||||
PK_INT = BigInteger().with_variant(Integer, "sqlite")
|
||||
|
||||
SYMBOL_LEN = 12
|
||||
|
||||
|
||||
class StockModel(Base):
|
||||
__tablename__ = "stock"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), unique=True, index=True)
|
||||
name: Mapped[str] = mapped_column(String(64))
|
||||
industry: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
area: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
market: Mapped[str | None] = mapped_column(String(16), nullable=True)
|
||||
exchange: Mapped[str | None] = mapped_column(String(8), nullable=True)
|
||||
list_date: Mapped[date] = mapped_column(Date)
|
||||
delist_date: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
status: Mapped[str] = mapped_column(String(8), default="L")
|
||||
|
||||
|
||||
class TradingCalendarModel(Base):
|
||||
__tablename__ = "trading_calendar"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
calendar_date: Mapped[date] = mapped_column(Date, unique=True, index=True)
|
||||
is_open: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
|
||||
|
||||
class StockDailyModel(Base):
|
||||
"""不复权日线。"""
|
||||
|
||||
__tablename__ = "stock_daily"
|
||||
__table_args__ = (UniqueConstraint("symbol", "trade_date", name="uq_daily_symbol_date"),)
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), index=True)
|
||||
trade_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
source: Mapped[str] = mapped_column(String(16), default="tushare", server_default="tushare")
|
||||
adjust: Mapped[str] = mapped_column(String(8), default="none", server_default="none")
|
||||
open: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
high: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
low: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
close: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
volume: Mapped[Decimal | None] = mapped_column(Numeric(24, 2), nullable=True)
|
||||
amount: Mapped[Decimal | None] = mapped_column(Numeric(24, 2), nullable=True)
|
||||
|
||||
|
||||
class AdjustFactorModel(Base):
|
||||
__tablename__ = "adjust_factor"
|
||||
__table_args__ = (UniqueConstraint("symbol", "trade_date", name="uq_adj_symbol_date"),)
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), index=True)
|
||||
trade_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
factor: Mapped[Decimal] = mapped_column(Numeric(20, 6))
|
||||
|
||||
|
||||
class DailyBasicModel(Base):
|
||||
"""每日指标快照(Tushare daily_basic)—— 估值 / 股息率 / 市值。
|
||||
|
||||
幂等键 (symbol, trade_date):同一交易日同一股票唯一一行。
|
||||
dv_ratio/dv_ttm 为时点值,研究侧按 trade_date <= as_of 取用(无未来函数)。
|
||||
"""
|
||||
|
||||
__tablename__ = "daily_basic"
|
||||
__table_args__ = (UniqueConstraint("symbol", "trade_date", name="uq_basic_symbol_date"),)
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), index=True)
|
||||
trade_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
source: Mapped[str] = mapped_column(String(16), default="tushare", server_default="tushare")
|
||||
close: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
turnover_rate: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
volume_ratio: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
pe: Mapped[Decimal | None] = mapped_column(Numeric(16, 4), nullable=True)
|
||||
pe_ttm: Mapped[Decimal | None] = mapped_column(Numeric(16, 4), nullable=True)
|
||||
pb: Mapped[Decimal | None] = mapped_column(Numeric(16, 4), nullable=True)
|
||||
ps: Mapped[Decimal | None] = mapped_column(Numeric(16, 4), nullable=True)
|
||||
ps_ttm: Mapped[Decimal | None] = mapped_column(Numeric(16, 4), nullable=True)
|
||||
dv_ratio: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
dv_ttm: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
total_share: Mapped[Decimal | None] = mapped_column(Numeric(24, 4), nullable=True)
|
||||
float_share: Mapped[Decimal | None] = mapped_column(Numeric(24, 4), nullable=True)
|
||||
free_share: Mapped[Decimal | None] = mapped_column(Numeric(24, 4), nullable=True)
|
||||
total_mv: Mapped[Decimal | None] = mapped_column(Numeric(24, 4), nullable=True)
|
||||
circ_mv: Mapped[Decimal | None] = mapped_column(Numeric(24, 4), nullable=True)
|
||||
|
||||
|
||||
class FinancialIndicatorModel(Base):
|
||||
"""财务指标快照 —— report_date(报告期) 与 announce_date(公告日) 并存。"""
|
||||
|
||||
__tablename__ = "financial_indicator"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("symbol", "report_date", "announce_date", name="uq_fin_sym_rep_ann"),
|
||||
)
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), index=True)
|
||||
report_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
announce_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
source: Mapped[str] = mapped_column(
|
||||
String(16), default="tushare", server_default="tushare"
|
||||
)
|
||||
eps: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
roe: Mapped[Decimal | None] = mapped_column(Numeric(10, 4), nullable=True)
|
||||
total_revenue: Mapped[Decimal | None] = mapped_column(Numeric(24, 2), nullable=True)
|
||||
net_profit: Mapped[Decimal | None] = mapped_column(Numeric(24, 2), nullable=True)
|
||||
gross_margin: Mapped[Decimal | None] = mapped_column(Numeric(10, 4), nullable=True)
|
||||
|
||||
|
||||
class SyncLogModel(Base):
|
||||
__tablename__ = "sync_log"
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
source: Mapped[str] = mapped_column(String(16), index=True)
|
||||
api: Mapped[str] = mapped_column(String(32))
|
||||
request_time: Mapped[datetime] = mapped_column(DateTime)
|
||||
success: Mapped[bool] = mapped_column(Boolean)
|
||||
failure_reason: Mapped[str | None] = mapped_column(String(500), nullable=True)
|
||||
row_count: Mapped[int] = mapped_column(default=0)
|
||||
data_start: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
data_end: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
|
||||
|
||||
class StockNameHistoryModel(Base):
|
||||
"""股票名称变更历史(Tushare namechange)—— 时点 ST / 风险警示判定的依据。
|
||||
|
||||
幂等键 (symbol, start_date):同一股票同一名称生效起点唯一一行。
|
||||
查询语义:`name` 在 [start_date, end_date] 内有效;`end_date` 为空表示至今有效。
|
||||
时点取值:`start_date <= as_of AND (end_date IS NULL OR end_date >= as_of)`。
|
||||
"""
|
||||
|
||||
__tablename__ = "stock_name_history"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("symbol", "start_date", name="uq_name_symbol_start"),
|
||||
)
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), index=True)
|
||||
name: Mapped[str] = mapped_column(String(64))
|
||||
start_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
end_date: Mapped[date | None] = mapped_column(Date, nullable=True, index=True)
|
||||
ann_date: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
change_reason: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
source: Mapped[str] = mapped_column(String(16), default="tushare", server_default="tushare")
|
||||
@@ -0,0 +1,42 @@
|
||||
"""选股持久化表(M6.3)。
|
||||
|
||||
selection_snapshot:一次选股运行的查询与统计快照(复现/历史查询用)
|
||||
selection_result:逐候选行(symbol/rank/score + 可解释字段 JSON),
|
||||
回答 v2 §8「某日为什么选这只股票」。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
from sqlalchemy import BigInteger, Date, DateTime, Integer, Numeric, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
PK_INT = BigInteger().with_variant(Integer, "sqlite")
|
||||
|
||||
|
||||
class SelectionSnapshotModel(Base):
|
||||
__tablename__ = "selection_snapshot"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
as_of: Mapped[date] = mapped_column(Date, index=True)
|
||||
method: Mapped[str] = mapped_column(String(16))
|
||||
query_json: Mapped[str] = mapped_column(Text)
|
||||
statistics_json: Mapped[str] = mapped_column(Text)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, index=True)
|
||||
|
||||
|
||||
class SelectionResultModel(Base):
|
||||
__tablename__ = "selection_result"
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
selection_id: Mapped[str] = mapped_column(String(32), index=True)
|
||||
symbol: Mapped[str] = mapped_column(String(12))
|
||||
rank: Mapped[int] = mapped_column(Integer)
|
||||
score: Mapped[float] = mapped_column(Numeric(14, 6))
|
||||
factor_values_json: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
filter_status_json: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
reason_json: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -0,0 +1,41 @@
|
||||
"""交易信号表(M8.1)。
|
||||
|
||||
signal_snapshot:一次信号运行的查询/规则/统计快照
|
||||
signal_event:逐信号(signal_type + score/price + trigger_reason JSON)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
from sqlalchemy import BigInteger, Date, DateTime, Integer, Numeric, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
PK_INT = BigInteger().with_variant(Integer, "sqlite")
|
||||
|
||||
|
||||
class SignalSnapshotModel(Base):
|
||||
__tablename__ = "signal_snapshot"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
as_of: Mapped[date] = mapped_column(Date, index=True)
|
||||
query_json: Mapped[str] = mapped_column(Text)
|
||||
rules_json: Mapped[str] = mapped_column(Text)
|
||||
statistics_json: Mapped[str] = mapped_column(Text)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, index=True)
|
||||
|
||||
|
||||
class SignalEventModel(Base):
|
||||
__tablename__ = "signal_event"
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
signal_id: Mapped[str] = mapped_column(String(32), index=True)
|
||||
symbol: Mapped[str] = mapped_column(String(12))
|
||||
signal_date: Mapped[date] = mapped_column(Date)
|
||||
signal_type: Mapped[str] = mapped_column(String(8))
|
||||
score: Mapped[float | None] = mapped_column(Numeric(14, 6), nullable=True)
|
||||
price: Mapped[float | None] = mapped_column(Numeric(14, 4), nullable=True)
|
||||
reason_json: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -0,0 +1,23 @@
|
||||
"""strategy 表(M8.3)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
|
||||
class StrategyModel(Base):
|
||||
__tablename__ = "strategy"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
name: Mapped[str] = mapped_column(String(64), unique=True)
|
||||
description: Mapped[str] = mapped_column(String(300), default="")
|
||||
# 2026-09 重构后策略库只存选股策略,新行一律 selection(历史行的 backtest 由数据迁移收敛)
|
||||
spec_type: Mapped[str] = mapped_column(String(16), default="selection")
|
||||
config_json: Mapped[str] = mapped_column(Text)
|
||||
version: Mapped[str] = mapped_column(String(16), default="1")
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user