Compare commits

...
21 Commits
Author SHA1 Message Date
Simon 314bfc159f docs: 同步 README/USAGE 至 M6–M8(选股系统主线 + MySQL + 新 API/Agent 工具/限制)
- README:架构分层示意(Selection/Signal/Portfolio/Composite Engine)、MySQL 默认配置、
  核心能力表、里程碑 M0–M8
- USAGE:概览模块与页面清单、引擎分层 §6.0、API 表新增 /api/selections /api/signals
  /api/strategies /api/composites、Agent 10 工具、已知限制(M8.4 延后/全市场选股同步耗时)
2026-09-09 06:29:51 +08:00
Simon db147c4232 docs: 记录 M6–M8 执行状态(ROADMAP 里程碑行 + DEV_PLAN 章节) 2026-09-09 00:42:17 +08:00
Simon e5a23f176d feat(web): M8.6b 回测页做实 + 首页入口更新(零新后端 API)
- backtest 页:成本参数表单(手续费/印花税/滑点 %)+ 初始资金 → ResearchSpec.costs/
  initial_capital;结果新增年度收益卡与成交明细表(买/卖日、买/卖价、收益、累计换手)
- Dashboard 入口卡加入「股票筛选」「交易信号」(6 张)
- lib/types.ts:ResearchSpec 增加 price_adjustment/costs/portfolio/initial_capital,
  Trade 补 entry_price/exit_price
- 页面 200 冒烟 + tsc --noEmit 通过
2026-09-09 00:41:49 +08:00
Simon e1ac23fa25 feat(web): M8.6a 交易信号页(/signals)+ 导航与类型
- app/signals/page.tsx:因子+买卖排名阈值+时点+白名单 → BUY/WATCH/SELL 表格
  (类型 Pill 着色、得分/价格/触发理由)+ 最近信号记录回看
- app-shell 导航加入「交易信号」;lib/types.ts 增加 Signal* 类型与 universe.symbols
- 真实 MySQL 端到端:/api/signals 返回可解释信号(跌破 MA60 的 SELL 警示等);
  /signals 页面 200;tsc --noEmit 通过
2026-09-09 00:40:22 +08:00
Simon 8b2f8ac35c feat(agent): M8.5 Agent 工具补齐(screen_stocks/explain_selection/generate_signals/create_strategy)
- tools_impl 新增 4 工具(Agent 共 10 个):
  screen_stocks(因子评分选股,symbols 白名单防全市场长任务)、
  explain_selection(读回选股结果并解释理由)、
  generate_signals(BUY/WATCH/SELL + 规则)、create_strategy(命名策略入库)
- 全部经白名单 Tool + Repository/Session,无 shell/写删权限扩张
- tests/test_agent_selection_tools.py 5 例(选股/策略保存+重名/解释 404/注册表);全量 pytest 通过
2026-09-09 00:39:42 +08:00
Simon 9d25d466e5 feat(strategy): M8.3 策略模型 + /api/strategies(命名配置资产,可展开为 ResearchSpec)
- StrategyDefinition:universe/factors/selection/rebalance/costs/portfolio +
  price_adjustment(除 period 外完整策略定义);to_research_spec(period) 展开为标准 Spec
- strategy 表(migration e1f2a3b4c5d6,MySQL 已应用;name 唯一)+ StrategyRepository
- /api/strategies:POST/GET/DELETE + POST /{id}/expand(period+initial_capital → ResearchSpec)
- tests/test_strategies.py(repo CRUD/同名/expand、API CRUD/400/404);全量 pytest 通过
2026-09-09 00:38:09 +08:00
Simon 692bdb3be5 feat(portfolio): M8.2 Portfolio Engine 模块化(等权收敛 + 约束显式标注)
- research.PortfolioSpec(weighting=equal;max_position_pct/max_industry_weight_pct 预留)
  + ResearchSpec.portfolio;config_snapshot 自动记录组合配置
- quant/portfolio.py:equal_weight_budget(与既有等权回测语义一致,行为收敛到本模块)+
  unimplemented_notes(设置约束即在结果中显式标注未建模,禁止假装支持)
- TopKBacktestRunner 预算与 unimplemented 改用 portfolio 模块;默认配置数值不变
  (一致性/quant 引擎回归通过);tests 补约束标注与 config_snapshot;全量 pytest 通过
2026-09-09 00:36:38 +08:00
Simon ba52edc2d6 feat(signal): M8.1 交易信号引擎(规则 + signal_event 落库 + /api/signals)
- SignalRules(买入 rank 阈值/趋势 MA/动量 + 卖出区间/破位警示)+ SignalEvent
  (BUY/WATCH/SELL,score/price/trigger_reason 可解释)+ SignalResult/Meta
- quant/signal.generate_signals:与选股同一评分引擎取全市场 rank,按规则分类输出
- signal_snapshot/signal_event 表(migration d8e0b2f3c4d5,MySQL 已应用)+ Repo
- SignalService + POST /api/signals(同步+落库)、GET 详情/列表
- tests/test_signals.py(引擎分类/排序/破位不 BUY、service、API 提交读回);全量 pytest 通过
2026-09-09 00:35:37 +08:00
Simon ef09d5b419 feat(quant): M7.3 研究行情口径显式化(默认不复权 none,可切 qfq)
- DailyBarRepository.get_range_many / stream_range_many_columns 增加 adjust 参数
  (默认 'none')→ SQL 层过滤口径,消除 stock_daily 混 source/adjust 污染因子的风险
- ResearchSpec / SelectionQuery 增加 price_adjustment(none|qfq),随 config_snapshot
  落库可溯源;ResearchService._load_daily 与 SelectionService 装配按口径取数
- tests/test_price_adjustment.py:repo 读取按 adjust 过滤(none/qfq 各自命中)、
  spec 默认与字段记录;全量 pytest 通过
2026-09-09 00:32:55 +08:00
Simon 4fa2bb748e feat(composite): M7.2b 因子组合落库 + /api/composites CRUD
- factor_composite 表(migration c3e9a0d1f4b5,MySQL 已应用;name 唯一)
- CompositeDefinition/Component 实体 + CompositeRepository Protocol + SQLAlchemy 实现
- /api/composites:POST(注册表自动填充组件 direction;未注册因子 400)、GET 列表/详情、DELETE
- tests/test_composites_api.py(repo CRUD/同名拒绝/删除、API 方向填充/404/400);全量 pytest 通过
2026-09-09 00:30:37 +08:00
Simon 273aee2772 refactor(quant): M7.2a Composite Engine 模块化(quant/composite.py)
- cross_sectional_zscore / composite_score / build_factor_panels 从 local_engine 迁入
  quant/composite.py;新增统一入口 build_score_panel(daily, factor_specs)
- local_engine re-export 保持旧引用兼容;selection/engine 的评分面板构建均指向
  composite —— 选股与回测的复合分实现收敛于一处
- 回归:quant/eval/research/selection 一致性/qlib 引擎测试全过;全量 pytest 通过
2026-09-09 00:29:23 +08:00
Simon 8f47b5b603 feat(factor): M7.1 因子定义入库 + /api/factors 读库(目录契约源)
- factor_definition 表(migration b7f2a5e81c33,MySQL 已应用):name 主键 + 元数据
  (formula/brief/frequency/lookback/direction/requires JSON/version)+ FactorDefinition
  entity(from_registry_def 由代码注册表构造)
- FactorRepository Protocol + SQLAlchemy 实现(幂等 upsert/list/get)
- /api/factors 改读 DB;目录为空自动 seed 注册表(幂等)—— 保留自定义因子登记能力
  (计算仍须代码注册,引用未注册因子照常 FactorError,防伪因子)
- tests/test_factor_catalog.py(repo 幂等/roundtrip/registry seed、API seed+字段齐全);
  test_api 的 client fixture 补 tmp sqlite session(factors 读库);全量 pytest 通过
2026-09-09 00:28:16 +08:00
Simon 0ab9038570 feat(web): M6.5 股票筛选页(/selection)—— 评分/条件双模式 + 当前/历史选股
- app/selection/page.tsx:因子评分(多因子+权重+TopN)与条件模式(结构化条件行
  编辑:static.*/因子/close>ma60/ROE 等,AND 语义);as_of 留空=最近交易日;
  结果表(rank/symbol/score/因子值/入选理由)+ 最近选股记录点击回看
- lib/types.ts:SelectionQuery/Condition/Candidate/Result/Meta/Run DTO
- app-shell 导航加入「股票筛选」;rail 脚注与 Dashboard 文案 SQLite→MySQL(M-DB 收尾)
- 端到端验证:真实 MySQL 上 POST /api/selections 返回候选并落库可读回;
  /selection 页面 200;后端全量 pytest + 前端 tsc --noEmit 通过
2026-09-09 00:25:27 +08:00
Simon 0d3e123de3 feat(selection): M6.4 回测与选股共用评分引擎(v2 §25 一致性锁定)
- quant/selection.score_panel_for_factors:复合分面板构建收敛为共享函数;
  LocalEngine.run_backtest 与 SelectionEngine.run_score_selection 均调它 ——
  消除「回测一套评分、选股另一套」的隐患
- tests/test_selection_backtest_consistency.py:对回测每个调仓日验证
  SelectionService.select(as_of=d, top_n) 候选 == 该日回测实际持仓(月调仓多时点),
  排序方向一致性亦验证;全量 pytest 通过
2026-09-09 00:22:08 +08:00
Simon c60dc78c88 feat(selection): M6.3 选股结果落库 + /api/selections(快照 + 逐候选行)
- domain/repositories/selection.py + SelectionMeta:选股 Repository Protocol
- models/selection.py + migration a6c91d4e7f20:selection_snapshot(查询/统计快照)+
  selection_result(逐候选:rank/score/factor_values/filter_status/reason JSON)
- repositories/selection_impl.py:save/get/list_recent(读回重建 SelectionResult)
- api/selections.py:POST /api/selections(同步执行+落库,返回 selection_id+result)、
  GET /{id}、GET 列表(as_of/method 过滤);deps 装配 SelectionService/Repo
- config._build_mysql_url:仅编码破坏 URL 结构的字符(修 Alembic configparser 遇 %21 崩)
- 迁移已在 MySQL 应用(alembic head a6c91d4e7f20);tests/test_selections_api.py 4 例
  提交→读回一致/404/列表过滤/condition;全量 pytest 通过
2026-09-09 00:20:42 +08:00
Simon 75c5472c31 feat(selection): M6.2 条件选股(method=condition + 财务可见性防护)
- quant/selection.run_condition_selection:结构化条件 AND 求值 —— 字段域
  static.*(行业/市场…)、技术列与派生量(close/volume/ma20/ma60)、已注册因子
  (momentum_60 等)、fundamental.*(announce_date<=as_of 的最新已公告财务值);
  条件支持 value 字面量与 ref 字段比较(如 close > ma60);结果带 filter_status/reason
- SelectionQuery 校验调整:condition 模式为纯过滤(不再强制 top_n/top_pct)
- FinancialRepository 新增 list_announced_many(批量防未来函数读取)+ SQLAlchemy 实现;
  SelectionService 注入 financial_repo 并按 announce_date 取每股最新一版
- tests/test_selection_condition.py:8 例(行业 in/ne、动量>0、close>ma60 ref、阈值、
  ROE 过滤且未来公告不可见、缺财务 repo 报错、更早 as_of 排除);全量 pytest 通过
2026-09-09 00:16:45 +08:00
Simon 25a1d9531a feat(selection): M6.1 Universe 选股范围执行器(规则化 + symbols 白名单 + 历史日语义)
- quant/universe.py:filter_stocks 从 quant/service 迁出并集中(ST/上市天数/退市过滤),
  as_of 当前/历史日语义由 delist/list_date 保证;exclude_suspended 依赖停牌表未建模,
  由上层显式标注(选股结果 unimplemented)
- research.UniverseSpec 增加 symbols 白名单(非空时仅白名单内参与,再叠加其余过滤)
- quant/service re-export filter_stocks(外部引用不变);SelectionService 已共用
- tests/test_universe.py:6 例覆盖当前/历史日、ST、上市天数、退市、白名单;全量 pytest 通过
2026-09-09 00:13:28 +08:00
Simon f3586adb25 feat(selection): M6.0 选股契约与评分引擎(SelectionQuery/Result + select(as_of))
- domain/entities/selection.py:SelectionQuery(universe+method+factors+top_n/top_pct/
  min_score+as_of+预热)与 SelectionResult/Candidate/Statistics(v2 §14.2/§21.1 DTO);
  ConditionSpec 字段就位供 M6.2 条件选股
- quant/selection.py:Selection Engine method=score —— 复合分(zscore×权重×方向)
  → TopN/Top% 截断;observation_date=<=as_of 最近交易日(防未来函数,v2 §9);
  候选带 factor_values 与 selection_reason(可解释)
- application/services/selection_service.py:选股用例(universe 过滤 → 装配 → 引擎)
- quant/service.py:抽取公共 load_daily_df 供研究/选股共用(行为不变)
- tests/test_selection.py:11 例 —— TopN/排序/理由、as_of 防未来函数、ST/上市天数/
  退市过滤、top_pct/min_score、空数据与查询校验;全量 pytest 通过
2026-09-09 00:12:28 +08:00
Simon 697ffc767b docs: DEV_PLAN_v2 以选股系统为下一阶段主线(用户定调)
- 新路线:M6 选股系统(Universe 选股范围 + Selection Engine A 条件/B 评分 +
  selection_result/snapshot 落库 + /api/selections + 回测共用引擎 + Web 选股页)
  → M7 因子层落地(原 M6 内容后移为支撑)→ M8 Signal/Portfolio/Strategy/Agent/模型选股
- 选股 MVP 复用现有数据与 9 因子/composite_score,不阻塞于因子入库
- 明确历史/当前 as_of 一致性、selection_reason 可解释、未来函数防护与范围控制
2026-09-09 00:06:31 +08:00
Simon 0ffd574f30 docs: 下阶段开发计划(架构 v2 落地 M6-M8)+ MySQL 迁移文档同步
- docs/DEV_PLAN_v2.md:基于 ARCHITECTURE_v2 与 M0-M5 现状的下一阶段计划
  (M-DB 迁移收尾 → M6 因子定义入库/复合因子/口径修复 → M7 Selection/Signal/
  Portfolio 引擎分层 → M8 Strategy 平台化/Web 做实/Agent 工具补齐;含本机 Redis
  127.0.0.1:6379 的接入触发点与执行顺序)
- ROADMAP.md:登记 M-DB 里程碑并指向 DEV_PLAN_v2
- USAGE.md / README.md:数据库描述由 SQLite 更新为 MySQL(config.yaml database.mysql)
2026-09-08 23:58:51 +08:00
Simon 6c2f198261 feat(db): SQLite 全量迁移至 MySQL(config.yaml 配置化 + 迁移脚本 + 一致性校验)
- config.yaml database.mysql:host/port/db/user/charset 明文可提交;密码经 password_env
  引用 .env 的 MYSQL_PASSWORD(AGENT.md §33 密钥不进 git)
- config.py _build_mysql_url:URL 优先级 DATABASE_URL env > database.mysql 段 > sqlite 兜底
- pyproject 引入 pymysql>=1.1
- tests/conftest.py 强制每进程 /tmp SQLite(测试绝不触 MySQL 开发库);test_config 覆盖
  mysql 组装/密码可选/sqlite 兜底分支
- scripts/migrate_sqlite_to_mysql.py:sqlite→mysql 一次性迁移工具(keyset 分页 + chunk
  多值 INSERT + 攒批 commit + 幂等续传 + 低配 MySQL 节流 --throttle-sec + --verify-only
  行数与抽样一致性校验);已用于 data/quant.db 约 1605 万行迁移并经校验一致
2026-09-08 23:58:48 +08:00
80 changed files with 5980 additions and 210 deletions
+7 -3
View File
@@ -9,12 +9,16 @@
TUSHARE_TOKEN= TUSHARE_TOKEN=
# ---- 数据库连接 ---- # ---- 数据库连接 ----
# 留空时使用默认 SQLite:<项目根>/data/quant.db(相对路径自动解析到项目根) # 默认库由 config.yaml database.mysql 决定(MySQL 192.168.1.10/qlib)。
# 连接优先级:本文件 DATABASE_URL > config.yaml database.mysql > SQLite 兜底。
DATABASE_URL= DATABASE_URL=
# SQLite 示例(显式指定): # SQLite 示例(显式指定):
# DATABASE_URL=sqlite:///./data/quant.db # DATABASE_URL=sqlite:///./data/quant.db
# MySQL 示例(未来切换,仅需改此值,业务代码不变): # MySQL 示例(手动覆盖 config.yaml 的 mysql 段时):
# DATABASE_URL=mysql+pymysql://user:password@127.0.0.1:3306/quant # DATABASE_URL=mysql+pymysql://user:password@127.0.0.1:3306/qlib
# MySQL 密码(config.yaml database.mysql.password_env 引用;host/port/db/user 在 config.yaml)
MYSQL_PASSWORD=
# ---- AI Agent / LLM(Phase 5)---- # ---- AI Agent / LLM(Phase 5)----
# 只需填写 API Key;URL 与模型名已在 config.yaml 的 agent.llm 中配置 # 只需填写 API Key;URL 与模型名已在 config.yaml 的 agent.llm 中配置
+41 -19
View File
@@ -1,8 +1,8 @@
# qlib-platform # 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)(架构文档),**改代码前必须阅读**。 详见 [AGENT.md](./AGENT.md)(开发约束)与 [docs/ARCHITECTURE.md](./docs/ARCHITECTURE.md)(架构文档),**改代码前必须阅读**。
```text ```text
Web 前端 (frontend/web) Web 前端 (frontend/web:总览 / 股票池 / 股票筛选 / 因子研究 / 因子组合 / 交易信号 / 选股回测 / 实验)
↓ REST / SSE ↓ REST / SSE(异步 Job 状态机)
FastAPI (backend) FastAPI (backend:业务对象 API,见下「核心能力」)
↓ Research Specification ↓ Research Specification(统一研究契约)
Application Service ── Quant Service ── Qlib Adapter ── Qlib Application Service(SelectionService / SignalService / ResearchService / Strategy …)
↓ ↓ ├── Selection Engine(条件选股 / 因子评分 / as_of 历史与当前一致)
Repository / DAO Parquet / SQLite ├── Signal Engine(BUY / WATCH / SELL + 理由)
├── Portfolio Engine(等权;约束预留并如实标注)
└── Quant Service ── Composite Engine ── Qlib Adapter ── Qlib
↓ ↓
Repository / DAO MySQL(默认)/ Parquet
``` ```
### 目录结构 ### 目录结构
@@ -35,7 +39,8 @@ qlib/
│ │ ├── application/ # 用例 / 应用服务(编排,不含框架细节) │ │ ├── application/ # 用例 / 应用服务(编排,不含框架细节)
│ │ ├── domain/ # 领域实体 + Repository Protocol │ │ ├── domain/ # 领域实体 + Repository Protocol
│ │ ├── infrastructure/ # SQLAlchemy / Alembic 等基础设施实现 │ │ ├── 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 访问) │ │ ├── agent/ # AI Research Agent(Phase 5,仅 Tool 访问)
│ │ └── core/ # 配置、通用组件 │ │ └── core/ # 配置、通用组件
│ └── tests/ │ └── tests/
@@ -75,21 +80,38 @@ uv run uvicorn app.main:app --reload --port 8000
# 交互文档:http://127.0.0.1:8000/docs # 交互文档: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 ```bash
cd backend cd backend
uv run alembic upgrade head # 首次运行会在 data/quant.db 建立版本表 uv run alembic upgrade head # 在当前配置指向的库上建表(默认 MySQL qlib)
uv run alembic revision --autogenerate -m "add xxx table" # 修改 Model 后生成迁移 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) | 股票筛选 | `POST /api/selections`(A 条件选股 / B 因子评分 TopN,`as_of` 当前/历史,结果可解释可复现);Web `/selection` |
- **Phase 4 Experiment**:所有研究自动可复现存档 | 因子目录 | `factor_definition` 落库;`GET /api/factors`;组合可保存复用 `POST /api/composites` |
- **Phase 5 AI Agent**:自然语言 → Research Plan → 受控 Tool → Experiment | 因子研究 | Job 异步 IC / RankIC / 分层测试(`/api/jobs`) |
| 交易信号 | `POST /api/signals`:评分排名 + 趋势规则 → BUY/WATCH/SELL + 理由;Web `/signals` |
| 选股回测 | `POST /api/backtests`(成本/涨跌停近似;与当前选股共用同一评分引擎,v2 §25 一致性) |
| 策略资产 | `strategy` 落库 + `/api/strategies`(命名配置,可展开为回测 spec) |
| Experiment | 每次研究自动归档 + 一键复跑 + SSE 进度 |
| AI Agent | `POST /api/agent/chat`,10 个受控 Tool(含 screen_stocks / explain_selection / generate_signals / create_strategy) |
## 里程碑(详见 [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 模型选股按需延后)
## 约定速查 ## 约定速查
+182
View File
@@ -12,10 +12,22 @@ from datetime import date
from app.agent.tools import Tool from app.agent.tools import Tool
from app.application.services.job_executor import default_factories, submit_and_run 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.research import ( from app.domain.entities.research import (
BacktestResult, BacktestResult,
FactorTestReport, FactorTestReport,
ResearchSpec, ResearchSpec,
UniverseSpec,
)
from app.domain.entities.selection import SelectionQuery
from app.domain.entities.signal import SignalRules
from app.domain.entities.strategy import StrategyDefinition
from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl import (
SqlAlchemySelectionRepository,
)
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
SqlAlchemyStrategyRepository,
) )
@@ -159,6 +171,121 @@ def build_tools(factories: dict | None = None) -> list[Tool]:
) )
return "\n".join(out) 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 = [x.strip() for x in str(_pick(args, "factors", "momentum_60")).split(",") if x.strip()]
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 = [x.strip() for x in str(_pick(args, "factors", "momentum_60")).split(",") if x.strip()]
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.strip(), "weight": 1.0}
for x in str(_pick(args, "factors", "momentum_60")).split(",")
if x.strip()
]
if not factors:
return "请提供至少一个 factors(逗号分隔)"
description = str(_pick(args, "description", "") or "")
st = StrategyDefinition(
name=name,
description=description,
universe=UniverseSpec(
exclude_st=bool(_pick(args, "exclude_st", True)), min_listing_days=0
),
factors=factors,
selection={"top_n": int(_pick(args, "top_n", 10) or 10)},
rebalance=str(_pick(args, "rebalance", "monthly")),
)
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]})"
return [ return [
Tool( Tool(
"search_stocks", "search_stocks",
@@ -236,6 +363,61 @@ def build_tools(factories: dict | None = None) -> list[Tool]:
}, },
compare_experiments, 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,
),
] ]
+69
View File
@@ -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}
+71
View File
@@ -10,15 +10,39 @@ from typing import Annotated
from fastapi import Depends from fastapi import Depends
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.application.services.selection_service import SelectionService
from app.application.services.signal_service import SignalService
from app.domain.repositories.composite import CompositeRepository
from app.domain.repositories.factor import FactorRepository
from app.domain.repositories.jobs import ExperimentRepository, JobRepository from app.domain.repositories.jobs import ExperimentRepository, JobRepository
from app.domain.repositories.market import ( from app.domain.repositories.market import (
DailyBarRepository, DailyBarRepository,
FinancialRepository,
StockRepository, 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.composite_impl import (
SqlAlchemyCompositeRepository,
)
from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
SqlAlchemyFactorRepository,
)
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository, SqlAlchemyDailyBarRepository,
SqlAlchemyFinancialRepository,
SqlAlchemyStockRepository, 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.infrastructure.persistence.sqlalchemy.session import get_session
from app.quant.engine import LocalEngine, QuantEngine from app.quant.engine import LocalEngine, QuantEngine
from app.quant.service import ResearchService from app.quant.service import ResearchService
@@ -34,6 +58,10 @@ def _daily_repo_factory(session: DbSession) -> DailyBarRepository:
return SqlAlchemyDailyBarRepository(session) return SqlAlchemyDailyBarRepository(session)
def _financial_repo_factory(session: DbSession) -> FinancialRepository:
return SqlAlchemyFinancialRepository(session)
def _engine_factory() -> QuantEngine: def _engine_factory() -> QuantEngine:
return LocalEngine() return LocalEngine()
@@ -46,10 +74,53 @@ def _service_factory(
return ResearchService(stock_repo, daily_repo, engine) return ResearchService(stock_repo, daily_repo, engine)
def _signal_service_factory(
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
) -> SignalService:
return SignalService(stock_repo, daily_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)],
) -> SelectionService:
return SelectionService(stock_repo, daily_repo, financial_repo)
def _selection_repo_factory(session: DbSession) -> SelectionRepository:
return SqlAlchemySelectionRepository(session)
def _factor_repo_factory(session: DbSession) -> FactorRepository:
return SqlAlchemyFactorRepository(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)
StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)] StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)] DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
EngineDep = Annotated[QuantEngine, Depends(_engine_factory)] EngineDep = Annotated[QuantEngine, Depends(_engine_factory)]
ResearchServiceDep = Annotated[ResearchService, Depends(_service_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)]
CompositeRepoDep = Annotated[CompositeRepository, Depends(_composite_repo_factory)]
SignalRepoDep = Annotated[SignalRepository, Depends(_signal_repo_factory)]
SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)]
StrategyRepoDep = Annotated[StrategyRepository, Depends(_strategy_repo_factory)]
def _job_repo_factory(session: DbSession): def _job_repo_factory(session: DbSession):
+13 -17
View File
@@ -1,26 +1,22 @@
"""因子目录 API:/api/factors。""" """因子目录 API:/api/factors(M7.1 起读 DB factor_definition)。
目录为空时自动从代码注册表 seed(幂等);随后可登记自定义因子元数据。
响应为 FactorDefinition 实体(含 requires 列表等)。
"""
from __future__ import annotations from __future__ import annotations
from fastapi import APIRouter from fastapi import APIRouter
from app.quant.factors import list_factors from app.api.deps import DbSession, FactorRepoDep
from app.application.services.factor_catalog import seed_registry_factors
from app.domain.entities.factor import FactorDefinition
router = APIRouter(prefix="/factors", tags=["factors"]) router = APIRouter(prefix="/factors", tags=["factors"])
@router.get("", summary="因子目录(含元数据)") @router.get("", summary="因子目录(含元数据,来自 factor_definition 表)")
def list_factor_catalog() -> list[dict]: def list_factor_catalog(factor_repo: FactorRepoDep, session: DbSession) -> list[FactorDefinition]:
return [ if not factor_repo.list():
{ seed_registry_factors(factor_repo, session) # 首次:从代码注册表 seed
"name": d.name, return factor_repo.list()
"description": d.description,
"brief": d.brief,
"formula": d.formula,
"frequency": d.frequency,
"lookback": d.lookback,
"direction": d.direction,
"requires": list(d.requires),
}
for d in list_factors()
]
+17 -1
View File
@@ -8,13 +8,29 @@ from __future__ import annotations
from fastapi import APIRouter from fastapi import APIRouter
from app.api import agent, experiments, factors, health, jobs, research, stocks from app.api import (
agent,
composites,
experiments,
factors,
health,
jobs,
research,
selections,
signals,
stocks,
strategies,
)
api_router = APIRouter() api_router = APIRouter()
api_router.include_router(health.router) api_router.include_router(health.router)
api_router.include_router(stocks.router) api_router.include_router(stocks.router)
api_router.include_router(factors.router) api_router.include_router(factors.router)
api_router.include_router(composites.router)
api_router.include_router(research.router) api_router.include_router(research.router)
api_router.include_router(selections.router)
api_router.include_router(signals.router)
api_router.include_router(strategies.router)
api_router.include_router(jobs.router) api_router.include_router(jobs.router)
api_router.include_router(experiments.router) api_router.include_router(experiments.router)
api_router.include_router(agent.router) api_router.include_router(agent.router)
+71
View File
@@ -0,0 +1,71 @@
"""选股 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
from typing import Annotated
from fastapi import APIRouter, HTTPException, Query
from pydantic import BaseModel
from app.api.deps import (
DbSession,
SelectionRepoDep,
SelectionServiceDep,
)
from app.application.services.job_executor import new_id
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("", 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)
+65
View File
@@ -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)
+80
View File
@@ -0,0 +1,80 @@
"""策略 API(M8.3):/api/strategies CRUD + 展开为 ResearchSpec。
POST /api/strategies 保存策略(name 唯一)
GET /api/strategies 列表
GET /api/strategies/{id}
DELETE /api/strategies/{id}
POST /api/strategies/{id}/expand body: {period:[start,end], initial_capital?} → ResearchSpec
"""
from __future__ import annotations
from datetime import date
from fastapi import APIRouter, HTTPException
from pydantic import BaseModel, Field
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 StrategyDefinition
router = APIRouter(prefix="/strategies", tags=["strategies"])
class ExpandRequest(BaseModel):
period: tuple[date, date]
initial_capital: float = Field(default=1_000_000.0, gt=0)
@router.post("", response_model=StrategyDefinition, summary="保存策略")
def create_strategy(
definition: StrategyDefinition,
strategy_repo: StrategyRepoDep,
session: DbSession,
) -> StrategyDefinition:
try:
saved = strategy_repo.save(definition.model_copy(update={"id": new_id("STG")}))
session.commit()
return saved
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@router.get("", response_model=list[StrategyDefinition], summary="策略列表")
def list_strategies(strategy_repo: StrategyRepoDep) -> list[StrategyDefinition]:
return strategy_repo.list()
@router.get("/{strategy_id}", response_model=StrategyDefinition, summary="读取策略")
def get_strategy(strategy_id: str, strategy_repo: StrategyRepoDep) -> StrategyDefinition:
row = strategy_repo.get(strategy_id)
if row is None:
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
return row
@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}
@router.post("/{strategy_id}/expand", response_model=ResearchSpec, summary="展开为研究 Spec")
def expand_strategy(
strategy_id: str,
req: ExpandRequest,
strategy_repo: StrategyRepoDep,
) -> ResearchSpec:
row = strategy_repo.get(strategy_id)
if row is None:
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
if req.period[0] >= req.period[1]:
raise HTTPException(status_code=400, detail="period 必须满足 start < end")
return row.to_research_spec(period=req.period, initial_capital=req.initial_capital)
@@ -0,0 +1,21 @@
"""因子目录用例:把代码注册表因子 seed 进 DB(M7.1)。
DB 为目录契约源;本服务在 /api/factors 首次读取为空时自动 seed(幂等),
后续代码新增因子也通过同一入口同步,保持目录与可计算因子一致。
"""
from __future__ import annotations
from app.domain.entities.factor import FactorDefinition
from app.domain.repositories.factor import FactorRepository
from app.quant.factors import list_factors
def seed_registry_factors(repo: FactorRepository, session) -> int:
"""把 quant/factors 注册表的元数据 upsert 进 factor_definition(幂等)。"""
defs = [FactorDefinition.from_registry_def(d) for d in list_factors()]
if not defs:
return 0
n = repo.upsert_many(defs)
session.commit()
return n
@@ -0,0 +1,116 @@
"""选股用例入口(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 filter_stocks, load_daily_df
_FUNDAMENTAL_PREFIX = "fundamental."
class SelectionService:
"""选股用例入口:select(query) → SelectionResult(当前或历史 as_of)。"""
def __init__(
self,
stock_repo: StockRepository,
daily_repo: DailyBarRepository,
financial_repo: FinancialRepository | None = None,
) -> None:
self._stock_repo = stock_repo
self._daily_repo = daily_repo
self._financial_repo = financial_repo
def select(self, query: SelectionQuery) -> SelectionResult:
as_of = query.as_of or date.today()
stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=as_of)
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))
daily = load_daily_df(
self._daily_repo,
symbols,
as_of - timedelta(days=query.warmup_days),
as_of,
columns,
adjust=query.price_adjustment,
)
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 _run(
self,
query: SelectionQuery,
daily: pd.DataFrame,
stocks: list,
as_of: date,
financial: dict[str, FinancialIndicator],
) -> SelectionResult:
if query.method == "score":
return run_score_selection(daily, query, as_of)
return run_condition_selection(daily, stocks, query, as_of, financial)
@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,39 @@
"""信号用例(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 filter_stocks, load_daily_df
from app.quant.signal import generate_signals
class SignalService:
def __init__(self, stock_repo: StockRepository, daily_repo: DailyBarRepository) -> None:
self._stock_repo = stock_repo
self._daily_repo = daily_repo
def signal(self, query: SelectionQuery, rules: SignalRules) -> SignalResult:
as_of = query.as_of or date.today()
stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=as_of)
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)
+48 -1
View File
@@ -99,6 +99,50 @@ def _normalize_sqlite_url(url: str) -> str:
return f"{_SQLITE_PREFIX}{(PROJECT_ROOT / rest).resolve()}" 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}"
@lru_cache @lru_cache
def get_settings() -> Settings: def get_settings() -> Settings:
"""加载并缓存 Settings。config / env 路径可通过参数覆盖以便测试。""" """加载并缓存 Settings。config / env 路径可通过参数覆盖以便测试。"""
@@ -115,7 +159,10 @@ def get_settings() -> Settings:
llm_model_env = agent_llm.get("model_env") or "LLM_MODEL" llm_model_env = agent_llm.get("model_env") or "LLM_MODEL"
default_url = "sqlite:///./data/quant.db" 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
migrations_rel = _deep(cfg, "database.migrations_dir") or ( migrations_rel = _deep(cfg, "database.migrations_dir") or (
"app/infrastructure/persistence/migrations" "app/infrastructure/persistence/migrations"
+33
View File
@@ -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
+39
View File
@@ -0,0 +1,39 @@
"""因子目录领域实体(M7.1:因子元数据 DB 化,v2 §11)。
DB 是因子目录的契约源:元数据(含自定义因子登记)入库;
计算执行仍由代码注册表(quant/factors.py)提供 —— 登记但未注册计算的因子
在 score/condition 中引用时仍抛 FactorError(防静默伪因子)。
"""
from __future__ import annotations
from datetime import datetime
from pydantic import BaseModel, Field
class FactorDefinition(BaseModel):
name: str = Field(min_length=1, max_length=64)
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
@classmethod
def from_registry_def(cls, d) -> FactorDefinition:
"""由 quant/factors.FactorDef(dataclass)构造目录实体(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),
)
+29 -2
View File
@@ -16,12 +16,20 @@ from pydantic import BaseModel, Field, field_validator, model_validator
class UniverseSpec(BaseModel): class UniverseSpec(BaseModel):
"""股票池口径。MVP:市场 + 过滤条件;指数成分等 Phase 3 扩展。""" """股票池口径。MVP:市场 + 过滤条件;指数成分等 Phase 3 扩展。
market: str = Field(default="CN_A", description="CN_A / CN_B / ...") symbols 白名单:非空时仅这些股票参与(再叠加其余过滤);供自选池/测试使用。
market 目前为预留字段(stock.market 存储主板/创业板/科创板等中文枚举,过滤未启用)。
"""
market: str = Field(default="CN_A", description="CN_A / CN_B / ...(预留)")
exclude_st: bool = True exclude_st: bool = True
exclude_suspended: bool = True exclude_suspended: bool = True
min_listing_days: int = Field(default=250, ge=0, description="上市至少 N 个自然日") min_listing_days: int = Field(default=250, ge=0, description="上市至少 N 个自然日")
symbols: list[str] = Field(
default_factory=list,
description="白名单(可选):非空时仅这些 symbol 参与选股/回测",
)
class FactorSpec(BaseModel): class FactorSpec(BaseModel):
@@ -37,6 +45,20 @@ class SelectionSpec(BaseModel):
top_n: int = Field(default=30, ge=1, le=1000) top_n: int = Field(default=30, ge=1, le=1000)
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): class CostSpec(BaseModel):
"""交易成本模型(单边比例)。 """交易成本模型(单边比例)。
@@ -54,11 +76,16 @@ class ResearchSpec(BaseModel):
type: str = Field(default="backtest", pattern="^(factor_test|backtest)$") type: str = Field(default="backtest", pattern="^(factor_test|backtest)$")
universe: UniverseSpec = UniverseSpec() universe: UniverseSpec = UniverseSpec()
price_adjustment: str = Field(
default="none", pattern="^(none|qfq)$",
description="研究行情口径:none 不复权(默认)/ qfq 前复权(result 与 config_snapshot 中显式)",
)
factors: list[FactorSpec] = Field(min_length=1) factors: list[FactorSpec] = Field(min_length=1)
selection: SelectionSpec = SelectionSpec() selection: SelectionSpec = SelectionSpec()
rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$") rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$")
period: tuple[date, date] period: tuple[date, date]
costs: CostSpec = CostSpec() costs: CostSpec = CostSpec()
portfolio: PortfolioSpec = PortfolioSpec()
initial_capital: float = Field(default=1_000_000.0, gt=0) initial_capital: float = Field(default=1_000_000.0, gt=0)
@field_validator("period") @field_validator("period")
+140
View File
@@ -0,0 +1,140 @@
"""选股系统领域对象(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 FactorSpec, UniverseSpec
class SelectionQuery(BaseModel):
"""一次选股查询(v2 §14.2 Selection 输入)。"""
universe: UniverseSpec = UniverseSpec()
price_adjustment: str = Field(
default="none", pattern="^(none|qfq)$",
description="行情口径:none 不复权(默认)/ qfq 前复权(结果 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
class ConditionSpec(BaseModel):
"""结构化选股条件(M6.2 使用)。
field 域:
- static.*:股票基础字段(industry / market / area / exchange / status…)
- 行情/技术字段:close / ma20 / ma60 / volume 及全部已注册因子名(momentum_60 等)
- fundamental.*:财务字段(eps / roe / total_revenue / net_profit / gross_margin),
仅取 announce_date <= as_of 的最新已公告值(防未来函数)
右操作数取 value(字面量)或 ref(另一字段名),二者二选一。
"""
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 SelectionCandidate(BaseModel):
"""单只候选股(v2 §14.3/§21.1)。"""
symbol: str
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
+57
View File
@@ -0,0 +1,57 @@
"""交易信号领域实体(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
+53
View File
@@ -0,0 +1,53 @@
"""策略领域实体(M8.3,v2 §17/§5.3)。
Strategy = 完整策略定义(universe + factors + selection + rebalance + costs +
portfolio,除回测区间 period 外),保存为命名资产;回测时补 period 展开为
ResearchSpec(v2 §18 Research Specification 为统一契约,策略是其持久化形态)。
"""
from __future__ import annotations
from datetime import date, datetime
from pydantic import BaseModel, Field
from app.domain.entities.research import (
CostSpec,
FactorSpec,
PortfolioSpec,
SelectionSpec,
UniverseSpec,
)
class StrategyDefinition(BaseModel):
id: str = ""
name: str = Field(min_length=1, max_length=64)
description: str = ""
spec_type: str = Field(default="backtest", pattern="^(backtest|factor_test)$")
universe: UniverseSpec = UniverseSpec()
price_adjustment: str = Field(default="none", pattern="^(none|qfq)$")
factors: list[FactorSpec] = Field(min_length=1)
selection: SelectionSpec = SelectionSpec()
rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$")
costs: CostSpec = CostSpec()
portfolio: PortfolioSpec = PortfolioSpec()
version: str = "1"
created_at: datetime | None = None
def to_research_spec(self, period: tuple[date, date], initial_capital: float | None = None):
"""补全回测区间/资金后展开为标准 ResearchSpec(可在 Job/回测执行)。"""
from app.domain.entities.research import ResearchSpec
return ResearchSpec(
type=self.spec_type,
universe=self.universe,
price_adjustment=self.price_adjustment,
factors=self.factors,
selection=self.selection,
rebalance=self.rebalance,
period=period,
costs=self.costs,
portfolio=self.portfolio,
initial_capital=initial_capital if initial_capital else 1_000_000.0,
)
@@ -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。"""
+16
View File
@@ -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: ...
+20 -2
View File
@@ -42,8 +42,14 @@ class DailyBarRepository(Protocol):
def get_range(self, symbol: str, start: date, end: date) -> list[DailyBar]: ... def get_range(self, symbol: str, start: date, end: date) -> list[DailyBar]: ...
def get_range_many(self, symbols: Sequence[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( def stream_range_many_columns(
self, self,
@@ -51,6 +57,7 @@ class DailyBarRepository(Protocol):
start: date, start: date,
end: date, end: date,
columns: Sequence[str], columns: Sequence[str],
adjust: str = "none",
) -> Iterator[tuple]: ) -> Iterator[tuple]:
"""流式(分批 yield)返回 symbol, trade_date(iso str), 数值列(float) 元组。 """流式(分批 yield)返回 symbol, trade_date(iso str), 数值列(float) 元组。
@@ -86,6 +93,17 @@ class FinancialRepository(Protocol):
) -> list[FinancialIndicator]: ) -> list[FinancialIndicator]:
"""只返回 announce_date <= as_of_date 的记录 —— 未来函数红线。""" """只返回 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): class SyncLogRepository(Protocol):
def add(self, log: SyncLog) -> SyncLog: ... def add(self, log: SyncLog) -> SyncLog: ...
@@ -0,0 +1,29 @@
"""选股 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.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 过滤)。"""
+19
View File
@@ -0,0 +1,19 @@
"""Signal Repository Protocol(M8.1 落库)。"""
from __future__ import annotations
from datetime import date
from typing import Protocol
from app.domain.entities.signal import 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]: ...
@@ -0,0 +1,20 @@
"""策略 Repository Protocol(M8.3)。"""
from __future__ import annotations
from typing import Protocol
from app.domain.entities.strategy import StrategyDefinition
class StrategyRepository(Protocol):
def save(self, definition: StrategyDefinition) -> StrategyDefinition:
"""新建(name 冲突抛 ValueError)。"""
def get(self, strategy_id: str) -> StrategyDefinition | None: ...
def get_by_name(self, name: str) -> StrategyDefinition | None: ...
def list(self) -> list[StrategyDefinition]: ...
def delete(self, strategy_id: str) -> bool: ...
@@ -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")
@@ -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")
@@ -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")
@@ -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")
@@ -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")
@@ -4,6 +4,12 @@
模型统一继承 infra.persistence.sqlalchemy.base.Base。 模型统一继承 infra.persistence.sqlalchemy.base.Base。
""" """
from app.infrastructure.persistence.sqlalchemy.models.composite import ( # noqa: F401
FactorCompositeModel,
)
from app.infrastructure.persistence.sqlalchemy.models.factor import ( # noqa: F401
FactorDefinitionModel,
)
from app.infrastructure.persistence.sqlalchemy.models.jobs import ( # noqa: F401 from app.infrastructure.persistence.sqlalchemy.models.jobs import ( # noqa: F401
ExperimentModel, ExperimentModel,
JobModel, JobModel,
@@ -16,3 +22,14 @@ from app.infrastructure.persistence.sqlalchemy.models.market import ( # noqa: F
SyncLogModel, SyncLogModel,
TradingCalendarModel, 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,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,28 @@
"""因子目录表(M7.1)。
factor_definition:因子元数据契约源(name 主键幂等);requires 以 JSON 存。
"""
from __future__ import annotations
from datetime import datetime
from sqlalchemy import 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"
name: Mapped[str] = mapped_column(String(64), 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")
created_at: Mapped[datetime] = mapped_column(DateTime)
@@ -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,22 @@
"""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="")
spec_type: Mapped[str] = mapped_column(String(16), default="backtest")
config_json: Mapped[str] = mapped_column(Text)
version: Mapped[str] = mapped_column(String(16), default="1")
created_at: Mapped[datetime] = mapped_column(DateTime)
@@ -0,0 +1,80 @@
"""因子组合 Repository 的 SQLAlchemy 实现(M7.2b)。"""
from __future__ import annotations
import json
from datetime import datetime
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.domain.entities.composite import CompositeComponent, CompositeDefinition
from app.infrastructure.persistence.sqlalchemy.models.composite import FactorCompositeModel
class SqlAlchemyCompositeRepository:
def __init__(self, session: Session) -> None:
self._session = session
def save(self, definition: CompositeDefinition) -> CompositeDefinition:
if not definition.id:
raise ValueError("需要 id(由调用方生成)")
exists = self._session.get(FactorCompositeModel, definition.id)
dup = self._session.scalar(
select(FactorCompositeModel)
.where(FactorCompositeModel.name == definition.name)
.limit(1)
)
if dup is not None and dup.id != definition.id:
raise ValueError(f"组合名已存在:{definition.name}")
now = definition.created_at or datetime.now()
if exists is None:
self._session.add(
FactorCompositeModel(
id=definition.id,
name=definition.name,
method=definition.method,
description=definition.description,
components_json=json.dumps(
[c.model_dump() for c in definition.components], ensure_ascii=False
),
created_at=now,
)
)
else:
exists.name = definition.name
exists.method = definition.method
exists.description = definition.description
exists.components_json = json.dumps(
[c.model_dump() for c in definition.components], ensure_ascii=False
)
self._session.flush()
return definition
def get(self, composite_id: str) -> CompositeDefinition | None:
row = self._session.get(FactorCompositeModel, composite_id)
return _to_entity(row) if row else None
def list(self) -> list[CompositeDefinition]:
rows = self._session.scalars(
select(FactorCompositeModel).order_by(FactorCompositeModel.name)
).all()
return [_to_entity(r) for r in rows]
def delete(self, composite_id: str) -> bool:
row = self._session.get(FactorCompositeModel, composite_id)
if row is None:
return False
self._session.delete(row)
return True
def _to_entity(row: FactorCompositeModel) -> CompositeDefinition:
return CompositeDefinition(
id=row.id,
name=row.name,
method=row.method,
description=row.description,
components=[CompositeComponent(**c) for c in json.loads(row.components_json)],
created_at=row.created_at,
)
@@ -0,0 +1,78 @@
"""因子目录 Repository 的 SQLAlchemy 实现(M7.1)。"""
from __future__ import annotations
import json
from datetime import datetime
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.domain.entities.factor import FactorDefinition
from app.infrastructure.persistence.sqlalchemy.models.factor import FactorDefinitionModel
def _to_entity(row: FactorDefinitionModel) -> FactorDefinition:
return FactorDefinition(
name=row.name,
description=row.description,
formula=row.formula,
brief=row.brief,
frequency=row.frequency,
lookback=row.lookback,
direction=row.direction,
requires=json.loads(row.requires_json or "[]"),
version=row.version,
created_at=row.created_at,
)
class SqlAlchemyFactorRepository:
def __init__(self, session: Session) -> None:
self._session = session
def upsert_many(self, definitions: list[FactorDefinition]) -> int:
if not definitions:
return 0
existing = {
r.name: r
for r in self._session.scalars(
select(FactorDefinitionModel).where(
FactorDefinitionModel.name.in_([d.name for d in definitions])
)
)
}
now = datetime.now()
for d in definitions:
row = existing.get(d.name)
if row is None:
self._session.add(
FactorDefinitionModel(
name=d.name,
description=d.description,
formula=d.formula,
brief=d.brief,
frequency=d.frequency,
lookback=d.lookback,
direction=d.direction,
requires_json=json.dumps(d.requires),
version=d.version,
created_at=now,
)
)
else:
for k, v in d.model_dump(exclude={"created_at"}).items():
if k == "requires":
v = json.dumps(v)
setattr(row, k, v)
return len(definitions)
def list(self) -> list[FactorDefinition]:
rows = self._session.scalars(
select(FactorDefinitionModel).order_by(FactorDefinitionModel.name)
).all()
return [_to_entity(r) for r in rows]
def get(self, name: str) -> FactorDefinition | None:
row = self._session.get(FactorDefinitionModel, name)
return _to_entity(row) if row else None
@@ -158,16 +158,24 @@ class SqlAlchemyDailyBarRepository:
).all() ).all()
return [DailyBar.model_validate(r, from_attributes=True) for r in rows] return [DailyBar.model_validate(r, from_attributes=True) for r in rows]
def get_range_many(self, symbols: Sequence[str], start: date, end: date) -> list[DailyBar]: def get_range_many(
rows = self._session.scalars( self,
symbols: Sequence[str],
start: date,
end: date,
adjust: str = "none",
) -> list[DailyBar]:
stmt = (
select(StockDailyModel) select(StockDailyModel)
.where( .where(
StockDailyModel.symbol.in_(list(symbols)), StockDailyModel.symbol.in_(list(symbols)),
StockDailyModel.trade_date >= start, StockDailyModel.trade_date >= start,
StockDailyModel.trade_date <= end, StockDailyModel.trade_date <= end,
StockDailyModel.adjust == adjust, # 研究主口径:不复权(v2 §8)
) )
.order_by(StockDailyModel.trade_date) .order_by(StockDailyModel.trade_date)
).all() )
rows = self._session.scalars(stmt).all()
return [DailyBar.model_validate(r, from_attributes=True) for r in rows] return [DailyBar.model_validate(r, from_attributes=True) for r in rows]
def stream_range_many_columns( def stream_range_many_columns(
@@ -176,6 +184,7 @@ class SqlAlchemyDailyBarRepository:
start: date, start: date,
end: date, end: date,
columns: Sequence[str], columns: Sequence[str],
adjust: str = "none",
) -> Iterator[tuple]: ) -> Iterator[tuple]:
"""流式返回 (symbol, trade_date_iso, *float_cols) 元组,分批拉取。 """流式返回 (symbol, trade_date_iso, *float_cols) 元组,分批拉取。
@@ -193,6 +202,7 @@ class SqlAlchemyDailyBarRepository:
StockDailyModel.symbol.in_(list(symbols)), StockDailyModel.symbol.in_(list(symbols)),
StockDailyModel.trade_date >= start, StockDailyModel.trade_date >= start,
StockDailyModel.trade_date <= end, StockDailyModel.trade_date <= end,
StockDailyModel.adjust == adjust,
) )
.order_by(StockDailyModel.symbol, StockDailyModel.trade_date) .order_by(StockDailyModel.symbol, StockDailyModel.trade_date)
.execution_options(yield_per=20000) .execution_options(yield_per=20000)
@@ -280,6 +290,28 @@ class SqlAlchemyFinancialRepository:
rows = self._session.scalars(stmt).all() rows = self._session.scalars(stmt).all()
return [FinancialIndicator.model_validate(r, from_attributes=True) for r in rows] return [FinancialIndicator.model_validate(r, from_attributes=True) for r in rows]
def list_announced_many(
self,
symbols: Sequence[str],
as_of_date: date,
) -> list[FinancialIndicator]:
"""批量:这些股票 announce_date <= as_of_date 的全部记录(防未来函数)。"""
if not symbols:
return []
rows = self._session.scalars(
select(FinancialIndicatorModel)
.where(
FinancialIndicatorModel.symbol.in_(list(symbols)),
FinancialIndicatorModel.announce_date <= as_of_date,
)
.order_by(
FinancialIndicatorModel.symbol,
FinancialIndicatorModel.announce_date,
FinancialIndicatorModel.report_date,
)
).all()
return [FinancialIndicator.model_validate(r, from_attributes=True) for r in rows]
class SqlAlchemySyncLogRepository: class SqlAlchemySyncLogRepository:
def __init__(self, session: Session) -> None: def __init__(self, session: Session) -> None:
@@ -0,0 +1,112 @@
"""选股 Repository 的 SQLAlchemy 实现(M6.3)。
save:snapshot + 候选逐行(同 session,由调用方 commit);
get:读回并重建 SelectionResult;list_recent:历史元数据。
"""
from __future__ import annotations
import json
from datetime import date, datetime
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.domain.entities.selection import (
SelectionCandidate,
SelectionMeta,
SelectionResult,
SelectionStatistics,
)
from app.infrastructure.persistence.sqlalchemy.models.selection import (
SelectionResultModel,
SelectionSnapshotModel,
)
class SqlAlchemySelectionRepository:
def __init__(self, session: Session) -> None:
self._session = session
def save(self, selection_id: str, result: SelectionResult) -> None:
self._session.add(
SelectionSnapshotModel(
id=selection_id,
as_of=result.as_of_date,
method=result.method,
query_json=json.dumps(result.config_snapshot, ensure_ascii=False),
statistics_json=json.dumps(result.statistics.model_dump(mode="json")),
created_at=datetime.now(),
)
)
now = datetime.now()
for c in result.candidates:
self._session.add(
SelectionResultModel(
selection_id=selection_id,
symbol=c.symbol,
rank=c.rank,
score=c.score,
factor_values_json=json.dumps(c.factor_values, ensure_ascii=False),
filter_status_json=json.dumps(c.filter_status, ensure_ascii=False),
reason_json=json.dumps(c.selection_reason, ensure_ascii=False),
created_at=now,
)
)
self._session.flush()
def get(self, selection_id: str) -> SelectionResult | None:
snap = self._session.get(SelectionSnapshotModel, selection_id)
if snap is None:
return None
rows = self._session.scalars(
select(SelectionResultModel)
.where(SelectionResultModel.selection_id == selection_id)
.order_by(SelectionResultModel.rank)
).all()
stats = SelectionStatistics.model_validate_json(snap.statistics_json)
candidates = [
SelectionCandidate(
symbol=r.symbol,
rank=r.rank,
score=float(r.score),
factor_values=json.loads(r.factor_values_json or "{}"),
filter_status=json.loads(r.filter_status_json or "[]"),
selection_reason=json.loads(r.reason_json or "[]"),
)
for r in rows
]
return SelectionResult(
as_of_date=snap.as_of,
method=snap.method,
statistics=stats,
candidates=candidates,
config_snapshot=json.loads(snap.query_json),
)
def list_recent(
self,
as_of: date | None = None,
method: str | None = None,
limit: int = 20,
) -> list[SelectionMeta]:
stmt = select(SelectionSnapshotModel).order_by(SelectionSnapshotModel.created_at.desc())
if as_of is not None:
stmt = stmt.where(SelectionSnapshotModel.as_of == as_of)
if method is not None:
stmt = stmt.where(SelectionSnapshotModel.method == method)
stmt = stmt.limit(limit)
metas: list[SelectionMeta] = []
for snap in self._session.scalars(stmt).all():
stats = SelectionStatistics.model_validate_json(snap.statistics_json)
metas.append(
SelectionMeta(
id=snap.id,
as_of=snap.as_of,
method=snap.method,
universe_size=stats.universe_size,
selected=stats.selected,
created_at=snap.created_at,
)
)
return metas
@@ -0,0 +1,108 @@
"""Signal Repository 的 SQLAlchemy 实现(M8.1)。"""
from __future__ import annotations
import json
from datetime import date, datetime
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.domain.entities.signal import (
SignalEvent,
SignalMeta,
SignalResult,
SignalRules,
SignalStatistics,
)
from app.infrastructure.persistence.sqlalchemy.models.signal import (
SignalEventModel,
SignalSnapshotModel,
)
class SqlAlchemySignalRepository:
def __init__(self, session: Session) -> None:
self._session = session
def save(self, signal_id: str, result: SignalResult) -> None:
self._session.add(
SignalSnapshotModel(
id=signal_id,
as_of=result.as_of_date,
query_json=json.dumps(
{
"as_of": result.as_of_date.isoformat(),
**{k: v for k, v in result.config_snapshot.items() if k != "as_of"},
},
ensure_ascii=False,
),
rules_json=json.dumps(result.rules.model_dump(mode="json"), ensure_ascii=False),
statistics_json=json.dumps(result.statistics.model_dump(mode="json")),
created_at=datetime.now(),
)
)
now = datetime.now()
for e in result.events:
self._session.add(
SignalEventModel(
signal_id=signal_id,
symbol=e.symbol,
signal_date=e.signal_date,
signal_type=e.signal_type,
score=e.score,
price=e.price,
reason_json=json.dumps(e.trigger_reason, ensure_ascii=False),
created_at=now,
)
)
self._session.flush()
def get(self, signal_id: str) -> SignalResult | None:
snap = self._session.get(SignalSnapshotModel, signal_id)
if snap is None:
return None
rows = self._session.scalars(
select(SignalEventModel)
.where(SignalEventModel.signal_id == signal_id)
.order_by(SignalEventModel.signal_type, SignalEventModel.symbol)
).all()
rules = SignalRules.model_validate_json(snap.rules_json)
stats = SignalStatistics.model_validate_json(snap.statistics_json)
return SignalResult(
as_of_date=snap.as_of,
rules=rules,
statistics=stats,
events=[
SignalEvent(
symbol=r.symbol,
signal_date=r.signal_date,
signal_type=r.signal_type,
score=float(r.score) if r.score is not None else None,
price=float(r.price) if r.price is not None else None,
trigger_reason=json.loads(r.reason_json or "[]"),
)
for r in rows
],
config_snapshot={"as_of": snap.as_of.isoformat()},
)
def list_recent(self, as_of: date | None = None, limit: int = 20) -> list[SignalMeta]:
stmt = select(SignalSnapshotModel).order_by(SignalSnapshotModel.created_at.desc())
if as_of is not None:
stmt = stmt.where(SignalSnapshotModel.as_of == as_of)
stmt = stmt.limit(limit)
metas: list[SignalMeta] = []
for snap in self._session.scalars(stmt).all():
stats = SignalStatistics.model_validate_json(snap.statistics_json)
metas.append(
SignalMeta(
id=snap.id,
as_of=snap.as_of,
buy=stats.buy,
watch=stats.watch,
sell=stats.sell,
created_at=snap.created_at,
)
)
return metas
@@ -0,0 +1,86 @@
"""策略 Repository 的 SQLAlchemy 实现(M8.3)。
config 以 JSON 存(StrategyDefinition.model_dump);读取时重建实体。
"""
from __future__ import annotations
import json
from datetime import datetime
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.domain.entities.strategy import StrategyDefinition
from app.infrastructure.persistence.sqlalchemy.models.strategy import StrategyModel
class SqlAlchemyStrategyRepository:
def __init__(self, session: Session) -> None:
self._session = session
def save(self, definition: StrategyDefinition) -> StrategyDefinition:
if not definition.id:
raise ValueError("需要 id(由调用方生成)")
dup = self._session.scalar(
select(StrategyModel).where(StrategyModel.name == definition.name).limit(1)
)
if dup is not None and dup.id != definition.id:
raise ValueError(f"策略名已存在:{definition.name}")
now = definition.created_at or datetime.now()
row = self._session.get(StrategyModel, definition.id)
config_json = json.dumps(definition.model_dump(exclude={"id", "created_at"}), ensure_ascii=False)
if row is None:
self._session.add(
StrategyModel(
id=definition.id,
name=definition.name,
description=definition.description,
spec_type=definition.spec_type,
config_json=config_json,
version=definition.version,
created_at=now,
)
)
else:
row.name = definition.name
row.description = definition.description
row.spec_type = definition.spec_type
row.config_json = config_json
row.version = definition.version
self._session.flush()
return definition
def get(self, strategy_id: str) -> StrategyDefinition | None:
row = self._session.get(StrategyModel, strategy_id)
return _to_entity(row) if row else None
def get_by_name(self, name: str) -> StrategyDefinition | None:
row = self._session.scalar(
select(StrategyModel).where(StrategyModel.name == name).limit(1)
)
return _to_entity(row) if row else None
def list(self) -> list[StrategyDefinition]:
rows = self._session.scalars(
select(StrategyModel).order_by(StrategyModel.name)
).all()
return [_to_entity(r) for r in rows]
def delete(self, strategy_id: str) -> bool:
row = self._session.get(StrategyModel, strategy_id)
if row is None:
return False
self._session.delete(row)
return True
def _to_entity(row: StrategyModel) -> StrategyDefinition:
data = json.loads(row.config_json)
# 列字段由 DB 行回填,避免与 config_json 重复
for key in ("name", "version", "description", "spec_type"):
data.pop(key, None)
return StrategyDefinition(
id=row.id, name=row.name, version=row.version,
created_at=row.created_at, **data,
)
+73
View File
@@ -0,0 +1,73 @@
"""Composite Factor Engine(M7.2,ARCHITECTURE_v2 §13)。
把「多因子 → 加权复合分面板」独立成模块,供:
- 选股(SelectionEngine.run_score_selection)
- 回测(LocalEngine.run_backtest)
共用同一实现(v2 §25 一致性)。
method 先落地 fixed(截面 zscore × 方向 × 权重 求和);Rank/Z/IC 加权留接口。
"""
from __future__ import annotations
import math
import pandas as pd
from app.quant.factors import FactorDef, compute_factor
def cross_sectional_zscore(panel: pd.DataFrame) -> pd.DataFrame:
"""截面 z-score。
候选不足 2 只(如单股票池)时退化为 0:无比较基准,但保留为可候选值;
全列缺失才为 NaN(该日不可选股)。
"""
def _row_z(row: pd.Series) -> pd.Series:
valid = row.dropna()
if len(valid) == 0:
return pd.Series(float("nan"), index=row.index)
if len(valid) == 1:
return pd.Series(0.0, index=row.index)
mu, sd = valid.mean(), valid.std()
if sd == 0 or math.isnan(sd):
return pd.Series(0.0, index=row.index)
return (row - mu) / sd
return panel.apply(_row_z, axis=1)
def composite_score(panels: list[tuple[str, pd.DataFrame, float, str]]) -> pd.DataFrame:
"""按 (name, panel, weight, direction) 计算加权复合 zscore。
direction="lower_is_better" 的因子取负号后相加(统一为「得分高者优先」)。
"""
total = None
for _name, panel, weight, direction in panels:
z = cross_sectional_zscore(panel)
if direction == "lower_is_better":
z = -z
contribution = z * weight
total = contribution if total is None else total.add(contribution, fill_value=0)
assert total is not None
return total
def build_factor_panels(
daily: pd.DataFrame, factor_specs
) -> list[tuple[str, pd.DataFrame, float, str]]:
"""按 spec.factors 计算面板与权重(因子不存在即报错)。"""
panels: list[tuple[str, pd.DataFrame, float, str]] = []
for fs in factor_specs:
defn: FactorDef
defn, panel = compute_factor(fs.name, daily)
panels.append((fs.name, panel, fs.weight, defn.direction))
return panels
def build_score_panel(daily: pd.DataFrame, factor_specs) -> pd.DataFrame:
"""因子加权复合分面板(index=trade_date, columns=symbol)。
回测与选股共用的统一入口 —— 保证 v2 §25/§27 一致性。
"""
return composite_score(build_factor_panels(daily, factor_specs))
+4 -8
View File
@@ -11,12 +11,8 @@ import pandas as pd
from app.domain.entities.research import BacktestResult, FactorTestReport, ResearchSpec from app.domain.entities.research import BacktestResult, FactorTestReport, ResearchSpec
from app.quant.factors import FactorError, get_factor from app.quant.factors import FactorError, get_factor
from app.quant.local_engine import ( from app.quant.local_engine import TopKBacktestRunner, run_spec_factor_test
TopKBacktestRunner, from app.quant.selection import score_panel_for_factors
build_factor_panels,
composite_score,
run_spec_factor_test,
)
# LocalEngine 路径恒需 close(TopK 收盘撮合 / 前瞻收益) # LocalEngine 路径恒需 close(TopK 收盘撮合 / 前瞻收益)
_CLOSE = {"close"} _CLOSE = {"close"}
@@ -64,7 +60,7 @@ class LocalEngine:
return report return report
def run_backtest(self, daily: pd.DataFrame, spec: ResearchSpec) -> BacktestResult: def run_backtest(self, daily: pd.DataFrame, spec: ResearchSpec) -> BacktestResult:
panels = build_factor_panels(daily, spec.factors) # 评分面板与选股共用同一构建(v2 §25:回测与当前选股同引擎)
score = composite_score(panels) score = score_panel_for_factors(daily, spec.factors)
close = daily.pivot(index="trade_date", columns="symbol", values="close").sort_index() close = daily.pivot(index="trade_date", columns="symbol", values="close").sort_index()
return TopKBacktestRunner(spec, score, close).run() return TopKBacktestRunner(spec, score, close).run()
+8 -54
View File
@@ -26,8 +26,13 @@ from app.domain.entities.research import (
Trade, Trade,
YearlyReturn, YearlyReturn,
) )
from app.quant.composite import ( # noqa: F401 —— re-export(模块化后旧引用仍可用)
build_factor_panels,
composite_score,
cross_sectional_zscore,
)
from app.quant.evaluation import run_factor_test from app.quant.evaluation import run_factor_test
from app.quant.factors import FactorDef, compute_factor from app.quant.portfolio import equal_weight_budget, unimplemented_notes
TRADING_DAYS = 252 TRADING_DAYS = 252
_DEFAULT_UNIMPLEMENTED = [ _DEFAULT_UNIMPLEMENTED = [
@@ -46,45 +51,6 @@ def _limit_up_ratio(symbol: str) -> float:
return 1.099 return 1.099
def cross_sectional_zscore(panel: pd.DataFrame) -> pd.DataFrame:
"""截面 z-score。
候选不足 2 只(如单股票池)时退化为 0:无比较基准,但保留为可候选值;
全列缺失才为 NaN(该日不可选股)。
"""
def _row_z(row: pd.Series) -> pd.Series:
valid = row.dropna()
if len(valid) == 0:
return pd.Series(float("nan"), index=row.index)
if len(valid) == 1:
return pd.Series(0.0, index=row.index)
mu, sd = valid.mean(), valid.std()
if sd == 0 or math.isnan(sd):
return pd.Series(0.0, index=row.index)
return (row - mu) / sd
return panel.apply(_row_z, axis=1)
def composite_score(
panels: list[tuple[str, pd.DataFrame, float, str]],
) -> pd.DataFrame:
"""按 (name, panel, weight, direction) 计算加权复合 zscore。
direction="lower_is_better" 的因子取负号后相加(统一为「得分高者优先」)。
"""
total = None
for _name, panel, weight, direction in panels:
z = cross_sectional_zscore(panel)
if direction == "lower_is_better":
z = -z
contribution = z * weight
total = contribution if total is None else total.add(contribution, fill_value=0)
assert total is not None
return total
def rebalance_dates(index: pd.Index, rebalance: str, start: date) -> list[pd.Timestamp]: def rebalance_dates(index: pd.Index, rebalance: str, start: date) -> list[pd.Timestamp]:
"""按频率取首个交易日(>= start)。""" """按频率取首个交易日(>= start)。"""
periods = index.to_period("M" if rebalance == "monthly" else "W") periods = index.to_period("M" if rebalance == "monthly" else "W")
@@ -204,7 +170,7 @@ class TopKBacktestRunner:
targets.append(s) targets.append(s)
if targets: if targets:
budget = cash / len(targets) budget = equal_weight_budget(cash, len(targets))
for s in targets: for s in targets:
c = float(close_d[s]) c = float(close_d[s])
price_in = c * (1 + self.costs.slippage_rate) price_in = c * (1 + self.costs.slippage_rate)
@@ -298,23 +264,11 @@ class TopKBacktestRunner:
positions=positions, positions=positions,
trades=trades, trades=trades,
turnover_pct=round(sum(notional) / max(init, 1) * 100, 2), turnover_pct=round(sum(notional) / max(init, 1) * 100, 2),
unimplemented=list(_DEFAULT_UNIMPLEMENTED), unimplemented=list(_DEFAULT_UNIMPLEMENTED) + unimplemented_notes(self.spec.portfolio),
config_snapshot=self.spec.model_dump(mode="json"), config_snapshot=self.spec.model_dump(mode="json"),
) )
def build_factor_panels(
daily: pd.DataFrame, factor_specs
) -> list[tuple[str, pd.DataFrame, float, str]]:
"""按 spec.factors 计算面板与权重(因子不存在即报错)。"""
panels: list[tuple[str, pd.DataFrame, float, str]] = []
for fs in factor_specs:
defn: FactorDef
defn, panel = compute_factor(fs.name, daily)
panels.append((fs.name, panel, fs.weight, defn.direction))
return panels
def run_spec_factor_test( def run_spec_factor_test(
daily: pd.DataFrame, daily: pd.DataFrame,
spec: ResearchSpec, spec: ResearchSpec,
+31
View File
@@ -0,0 +1,31 @@
"""Portfolio Engine(v2 §16)—— 组合构建模块(M8.2)。
MVP:等权资金拆分(与既有 TopK 回测等权语义一致,行为收敛到本模块);
单股/行业上限等约束为预留字段,未建模时由回测器写入 unimplemented
(禁止假装支持,AGENT.md §24)。
"""
from __future__ import annotations
from app.domain.entities.research import PortfolioSpec
def equal_weight_budget(cash: float, target_count: int) -> float:
"""等权单标的预算:现金均分(target_count>0)。"""
if target_count <= 0:
return 0.0
return cash / target_count
def unimplemented_notes(portfolio: PortfolioSpec) -> list[str]:
"""组合层未建模项说明(默认空;设置约束即显式标注)。"""
notes: list[str] = []
if portfolio.max_position_pct is not None:
notes.append(
f"最大单股权重 {portfolio.max_position_pct:.0%} 约束未建模(Portfolio v1 仅等权)"
)
if portfolio.max_industry_weight_pct is not None:
notes.append(
f"最大行业权重 {portfolio.max_industry_weight_pct:.0%} 约束未建模(Portfolio v1 仅等权)"
)
return notes
+375
View File
@@ -0,0 +1,375 @@
"""Selection Engine(ARCHITECTURE_v2 §14)—— 纯 pandas 执行层。
当前实现 method=score:因子加权复合分 → TopN/Top% 截断,输出 SelectionResult。
M6.2 在同一模块加入 method=condition(结构化条件选股)。
未来函数纪律:面板只在 <= observation_date 的数据上计算;observation_date 是
<= as_of 的最近可用交易日(as_of 显式传入即历史选股,None 则到数据最新)。
data 长表由 Service 装配(已按 universe 过滤 symbol、含预热窗口)。
"""
from __future__ import annotations
from datetime import date
import pandas as pd
from app.domain.entities.market import FinancialIndicator
from app.domain.entities.selection import (
SelectionCandidate,
SelectionQuery,
SelectionResult,
SelectionStatistics,
)
from app.quant.composite import build_score_panel
from app.quant.factors import FactorError, compute_factor, get_factor
_UNIMPLEMENTED_DEFAULT = [
"exclude_suspended 依赖停牌数据,当前未建模(结果可能包含停牌股)",
]
def score_panel_for_factors(daily: pd.DataFrame, factor_specs) -> pd.DataFrame:
"""因子加权复合分面板(index=trade_date, columns=symbol)。
回测(LocalEngine)与选股(run_score_selection)共用同一构建 ——
保证 v2 §25/§27「历史回测与当前选股使用同一套引擎」的一致性。
"""
return build_score_panel(daily, factor_specs) # 未知因子在此抛 FactorError
def resolve_observation_date(daily: pd.DataFrame, as_of: date | None) -> pd.Timestamp | None:
"""<= as_of 的最近可用交易日;as_of=None 取数据最新一日。"""
if daily.empty:
return None
dates = pd.to_datetime(daily["trade_date"])
if as_of is None:
return dates.max()
avail = dates[dates <= pd.Timestamp(as_of)]
return avail.max() if len(avail) else None
def factor_columns(query: SelectionQuery) -> set[str]:
"""score 模式所需行情数值列(数据装配裁剪用)。"""
needed = {"close"}
for fs in query.factors:
try:
defn, _fn = get_factor(fs.name)
except FactorError:
continue # 未知因子由执行期统一报错(score_selection 中 build_factor_panels)
needed.update(defn.requires)
return needed
def run_score_selection(
daily: pd.DataFrame,
query: SelectionQuery,
as_of: date | None,
) -> SelectionResult:
"""因子评分选股(v2 §14.1B):复合分 → 排序 → TopN/Top%。"""
if query.method != "score":
raise ValueError(f"run_score_selection 需要 method=score,当前 {query.method}")
obs = resolve_observation_date(daily, as_of)
if obs is None:
resolved = as_of or date.today()
return SelectionResult(
as_of_date=resolved,
method=query.method,
statistics=SelectionStatistics(),
candidates=[],
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
config_snapshot=query.model_dump(mode="json"),
)
resolved = obs.date()
# 只允许使用 <= obs 的数据(面板计算在截断后数据上进行)
view = daily[pd.to_datetime(daily["trade_date"]) <= obs]
if view.empty:
return SelectionResult(
as_of_date=resolved,
method=query.method,
statistics=SelectionStatistics(),
candidates=[],
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
config_snapshot=query.model_dump(mode="json"),
)
score = score_panel_for_factors(view, query.factors).loc[obs].dropna().sort_values(
ascending=False
)
# 每因子在 obs 行的原始值(factor_values 供展示与解释;与 build_factor_panels 同数据)
raw: dict[str, pd.Series] = {}
for fs in query.factors:
_defn, panel = compute_factor(fs.name, view)
if obs in panel.index:
raw[fs.name] = panel.loc[obs]
candidates_df = _truncate(score, query)
evaluated = int(len(score)) # score 已 dropna,长度即有分股票数
candidates: list[SelectionCandidate] = []
for rank, (sym, sc) in enumerate(candidates_df.items(), start=1):
factor_values = {
name: _to_float(series.get(sym))
for name, series in raw.items()
if isinstance(series, pd.Series)
}
factor_values = {k: v for k, v in factor_values.items() if v is not None}
candidates.append(
SelectionCandidate(
symbol=sym,
rank=rank,
score=round(float(sc), 6),
factor_values=factor_values,
selection_reason=_score_reason(query, sym, raw),
)
)
return SelectionResult(
as_of_date=resolved,
method=query.method,
statistics=SelectionStatistics(
universe_size=_symbol_count(view),
evaluated=evaluated,
selected=len(candidates),
),
candidates=candidates,
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
config_snapshot=query.model_dump(mode="json"),
)
def _truncate(score: pd.Series, query: SelectionQuery) -> pd.Series:
"""按 top_n / top_pct / min_score 截断(入参已按分数降序)。"""
s = score
if query.min_score is not None:
s = s[s >= query.min_score]
if query.top_pct is not None:
n = max(int(round(len(s) * query.top_pct)), 1)
s = s.head(n)
elif query.top_n is not None:
s = s.head(query.top_n)
return s
def _score_reason(query: SelectionQuery, symbol: str, raw: dict[str, pd.Series]) -> list[str]:
"""生成可读的入选理由:列每个因子的观测值与权重。"""
reasons: list[str] = []
for fs in query.factors:
try:
defn, _fn = get_factor(fs.name)
except FactorError:
continue
series = raw.get(fs.name)
val = _to_float(series.get(symbol)) if isinstance(series, pd.Series) else None
if val is None:
continue
good = defn.direction == "higher_is_better"
reasons.append(
f"{fs.name}={val:.4f}(权重 {fs.weight},{'越高越好' if good else '越低越好'})"
)
return reasons
def _symbol_count(daily: pd.DataFrame) -> int:
return int(daily["symbol"].nunique()) if not daily.empty and "symbol" in daily else 0
# ---------- method=condition:结构化条件选股(M6.2) ----------
# 技术字段:预计算派生量 + 行情原列(原列需在装配列中才可用)
_TECH_DERIVED = ("ma20", "ma60")
_STATIC_PREFIX = "static."
_FUNDAMENTAL_PREFIX = "fundamental."
def condition_needed_columns(query: SelectionQuery) -> set[str]:
"""条件引用的行情列(fundamental/static 走元数据与财务表,不需要行情列)。"""
needed = {"close"}
names = [c.field for c in query.conditions] + [
c.ref for c in query.conditions if c.ref and not c.ref.startswith(_FUNDAMENTAL_PREFIX)
]
for f in names:
if not f or f.startswith((_STATIC_PREFIX, _FUNDAMENTAL_PREFIX)):
continue
if f in {"open", "high", "low", "close", "volume", "amount", *_TECH_DERIVED}:
if f not in _TECH_DERIVED:
needed.add(f)
continue
try: # 其余按已注册因子处理
defn, _fn = get_factor(f)
except FactorError:
raise ValueError(
f"条件字段未知:{f}(可用: 行情列/ma20/ma60/已注册因子/static.*/fundamental.*)"
) from None
needed.update(defn.requires)
return needed
def run_condition_selection(
daily: pd.DataFrame,
stocks: list,
query: SelectionQuery,
as_of: date | None,
financial: dict[str, FinancialIndicator] | None = None,
) -> SelectionResult:
"""条件选股(v2 §14.1A):全部条件 AND 通过者入选(无排序;truncation 不适用)。
fields 域:static.*(股票基础)、close/volume/amount/ma20/ma60/已注册因子(行情)、
fundamental.*(announce_date <= as_of 的最新已公告财务值 —— 防未来函数由 Service 取数保证)。
"""
if query.method != "condition":
raise ValueError(f"run_condition_selection 需要 method=condition,当前 {query.method}")
obs = resolve_observation_date(daily, as_of)
resolved = (obs.date() if obs is not None else as_of) or date.today()
if obs is None or daily.empty:
return SelectionResult(
as_of_date=resolved, method=query.method,
statistics=SelectionStatistics(), candidates=[],
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
config_snapshot=query.model_dump(mode="json"),
)
view = daily[pd.to_datetime(daily["trade_date"]) <= obs]
close = view.pivot(index="trade_date", columns="symbol", values="close").sort_index()
close.index = pd.to_datetime(close.index)
# 技术字段面板(obs 行)
tech: dict[str, pd.Series] = {}
for col in ("close", "open", "high", "low", "volume", "amount"):
if col in view.columns and col != "close":
panel = view.pivot(index="trade_date", columns="symbol", values=col).sort_index()
panel.index = pd.to_datetime(panel.index)
tech[col] = panel.loc[obs]
tech["close"] = close.loc[obs]
tech["ma20"] = close.rolling(20).mean().loc[obs]
tech["ma60"] = close.rolling(60).mean().loc[obs]
# 因子字段按需计算
for cond in query.conditions:
for f in (cond.field, cond.ref):
if f is None or f.startswith((_STATIC_PREFIX, _FUNDAMENTAL_PREFIX)) or f in tech:
continue
if f in _TECH_DERIVED or f in ("close", "open", "high", "low", "volume", "amount"):
continue
try:
_defn, panel = compute_factor(f, view)
except FactorError:
continue # 已在 condition_needed_columns 报错;此处防御
if obs in panel.index:
tech[f] = panel.loc[obs]
statics = {s.symbol: s.model_dump() for s in stocks}
candidates: list[SelectionCandidate] = []
passed_symbols: list[str] = []
for sym in sorted(statics):
statuses: list[str] = []
all_ok = True
for cond in query.conditions:
ok = _eval_condition(cond, sym, statics, tech, financial or {})
statuses.append(f"{cond.field} {cond.op} {cond.ref or cond.value}: {'通过' if ok else '未通过'}")
all_ok = all_ok and ok
if all_ok:
passed_symbols.append(sym)
candidates.append(
SelectionCandidate(
symbol=sym,
rank=0, # 占位,末尾统一编号
score=1.0,
filter_status=statuses,
selection_reason=[f"通过全部 {len(query.conditions)} 条条件"],
)
)
for rank, c in enumerate(candidates, start=1):
c.rank = rank
return SelectionResult(
as_of_date=resolved,
method=query.method,
statistics=SelectionStatistics(
universe_size=len(statics),
evaluated=len(statics),
selected=len(candidates),
),
candidates=candidates,
unimplemented=list(_UNIMPLEMENTED_DEFAULT) + [
"条件选股为纯过滤(AND),未排序/未截断;如需排序请在 factors 中提供评分",
],
config_snapshot=query.model_dump(mode="json"),
)
def _eval_condition(
cond,
sym: str,
statics: dict,
tech: dict[str, pd.Series],
financial: dict,
) -> bool:
"""求值单条条件:value 与 ref 二选一;left 与 right 同为 field 或 field vs 字面量。"""
left = _field_value(cond.field, sym, statics, tech, financial)
if cond.ref is not None:
right = _field_value(cond.ref, sym, statics, tech, financial)
else:
right = cond.value
return _compare(left, right, cond.op)
def _field_value(field, sym, statics, tech, financial):
if field.startswith(_STATIC_PREFIX):
return statics.get(sym, {}).get(field[len(_STATIC_PREFIX):])
if field.startswith(_FUNDAMENTAL_PREFIX):
fin = financial.get(sym)
return getattr(fin, field[len(_FUNDAMENTAL_PREFIX):], None) if fin else None
series = tech.get(field)
if series is None:
return None
v = series.get(sym)
return None if v is None or (isinstance(v, float) and v != v) else v # NaN → None
def _compare(left, right, op: str) -> bool:
"""混合比较:None 视为不可用 → 除 ne 外不通过;数值/字符串分别处理。"""
if op == "ne":
return left != right
if left is None or right is None:
return False
try:
if isinstance(left, (int, float)) or isinstance(right, (int, float)):
return _num_cmp(float(left), float(right), op)
except (TypeError, ValueError):
pass
# 字符串/其它:支持 eq/ne/in/not_in
if op == "eq":
return left == right
if op == "in":
return left in right
if op == "not_in":
return left not in right
if op in ("gt", "gte", "lt", "lte"):
return _num_cmp(left, right, op) # 尝试数值,字符串会 ValueError → False
return False
def _num_cmp(a: float, b: float, op: str) -> bool:
if op == "gt":
return a > b
if op == "gte":
return a >= b
if op == "lt":
return a < b
if op == "lte":
return a <= b
if op == "eq":
return a == b
return a != b
def _to_float(v) -> float | None:
if v is None:
return None
try:
f = float(v)
except (TypeError, ValueError):
return None
if f != f: # NaN
return None
return f
+42 -41
View File
@@ -15,41 +15,22 @@ from datetime import date, timedelta
import pandas as pd import pandas as pd
from app.domain.entities.market import Stock
from app.domain.entities.research import ( from app.domain.entities.research import (
BacktestResult, BacktestResult,
FactorTestReport, FactorTestReport,
ResearchSpec, ResearchSpec,
UniverseSpec,
) )
from app.domain.repositories.market import ( from app.domain.repositories.market import (
DailyBarRepository, DailyBarRepository,
StockRepository, StockRepository,
) )
from app.quant.engine import QuantEngine from app.quant.engine import QuantEngine
from app.quant.universe import filter_stocks # noqa: F401 —— 选股/回测共用范围过滤
# 流式路径每攒多少行落一个 DataFrame 分片(控制 concat 峰值) # 流式路径每攒多少行落一个 DataFrame 分片(控制 concat 峰值)
_FRAME_CHUNK_ROWS = 50_000 _FRAME_CHUNK_ROWS = 50_000
def filter_stocks(stocks: list[Stock], universe: UniverseSpec, as_of: date) -> list[Stock]:
"""按股票池口径过滤(名称含 ST 判定 —— 名称快照为当日口径,属历史可追溯数据)。"""
out: list[Stock] = []
for s in stocks:
if s.delist_date is not None and s.delist_date < as_of:
continue
if universe.exclude_st and s.name and "ST" in s.name.upper():
continue
if (
universe.min_listing_days
and s.list_date
and (as_of - s.list_date).days < universe.min_listing_days
):
continue
out.append(s)
return out
def bars_to_daily_df(bars) -> pd.DataFrame: def bars_to_daily_df(bars) -> pd.DataFrame:
"""DailyBar 列表 → 引擎长表 DataFrame(symbol/trade_date/ohlc/volume/amount)。 """DailyBar 列表 → 引擎长表 DataFrame(symbol/trade_date/ohlc/volume/amount)。
@@ -88,6 +69,39 @@ def _frame_from_stream(rows: Iterable[tuple], columns: list[str]) -> pd.DataFram
return df return df
def load_daily_df(
daily_repo,
symbols: list[str],
start: date,
end: date,
columns: list[str],
adjust: str = "none",
) -> pd.DataFrame:
"""从 Repository 装配行情长表(供研究/选股共用)。
优先走流式列裁剪(stream_range_many_columns,SQL 侧转 REAL、分批),
失败或实现缺失时回退 get_range_many / 逐只 get_range。
"""
if not symbols:
return pd.DataFrame()
streamer = getattr(daily_repo, "stream_range_many_columns", None)
if streamer is not None:
try:
return _frame_from_stream(
streamer(symbols, start, end, sorted(columns), adjust=adjust), sorted(columns)
)
except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现)
pass
get_many = getattr(daily_repo, "get_range_many", None)
if get_many is not None:
bars = list(get_many(symbols, start, end, adjust=adjust))
else: # 兜底:逐只查询
bars = []
for sym in symbols:
bars.extend(daily_repo.get_range(sym, start, end))
return bars_to_daily_df(bars)
class ResearchService: class ResearchService:
"""研究用例入口(因子测试 / 回测)。依赖注入 Repository 与引擎。""" """研究用例入口(因子测试 / 回测)。依赖注入 Repository 与引擎。"""
@@ -120,26 +134,13 @@ class ResearchService:
# 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量) # 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量)
data_start = start - timedelta(days=300) data_start = start - timedelta(days=300)
stocks = filter_stocks(self._stock_repo.list(), spec.universe, as_of=start) stocks = filter_stocks(self._stock_repo.list(), spec.universe, as_of=start)
if not stocks:
return pd.DataFrame()
symbols = [s.symbol for s in stocks]
# 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV) # 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV)
required = self._engine.required_columns(spec) required = self._engine.required_columns(spec)
streamer = getattr(self._daily_repo, "stream_range_many_columns", None) return load_daily_df(
if streamer is not None: self._daily_repo,
try: [s.symbol for s in stocks],
return _frame_from_stream( data_start,
streamer(symbols, data_start, end, sorted(required)), sorted(required) end,
) sorted(required),
except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现) adjust=spec.price_adjustment,
pass )
# 旧路径:逐实体(供内存 / Fake 仓储等实现使用)
get_many = getattr(self._daily_repo, "get_range_many", None)
if get_many is not None:
bars = list(get_many(symbols, data_start, end))
else: # 兜底:逐只查询
bars = []
for s in stocks:
bars.extend(self._daily_repo.get_range(s.symbol, data_start, end))
return bars_to_daily_df(bars)
+135
View File
@@ -0,0 +1,135 @@
"""Signal Engine(v2 §15)—— 纯 pandas 执行。
输入:行情长表(<=as_of)+ SelectionQuery(评分因子)+ SignalRules。
流程:复合分 → 全市场 rank → 按规则判定 BUY / WATCH / SELL,
事件带 trigger_reason 与 score/price(可解释)。回测的买入逻辑(TopK+可买过滤)
与这里的 BUY 建议同源于同一评分引擎(v2 §25 一致)。
"""
from __future__ import annotations
from datetime import date
import pandas as pd
from app.domain.entities.selection import SelectionQuery
from app.domain.entities.signal import (
SignalEvent,
SignalResult,
SignalRules,
SignalStatistics,
)
from app.quant.composite import build_score_panel
from app.quant.selection import resolve_observation_date
def generate_signals(
daily: pd.DataFrame,
query: SelectionQuery,
rules: SignalRules,
as_of: date | None,
) -> SignalResult:
if not query.factors:
raise ValueError("signal 需要评分因子(SelectionQuery.factors)")
obs = resolve_observation_date(daily, as_of)
resolved = (obs.date() if obs is not None else as_of) or date.today()
if obs is None or daily.empty:
return _empty(query, rules, resolved)
view = daily[pd.to_datetime(daily["trade_date"]) <= obs]
score = build_score_panel(view, query.factors).loc[obs].dropna().sort_values(ascending=False)
close = view.pivot(index="trade_date", columns="symbol", values="close").sort_index()
close.index = pd.to_datetime(close.index)
c_d = close.loc[obs]
ma_d = close.rolling(rules.trend_ma).mean().loc[obs]
mom20_d = (close / close.shift(20) - 1.0).loc[obs]
events: list[SignalEvent] = []
stats = SignalStatistics(universe_size=int(len(score)))
for rank, (sym, sc) in enumerate(score.items(), start=1):
if rank > rules.max_output_rank:
break
c = _num(c_d.get(sym))
ma = _num(ma_d.get(sym))
mom = _num(mom20_d.get(sym))
price = float(c) if c is not None else None
reason: list[str] = []
event_type = "WATCH"
trend_ok = c is not None and ma is not None and c > ma
momentum_ok = mom is not None and mom > 0
if rank <= rules.buy_rank_threshold and (
not rules.buy_require_trend or trend_ok
) and (not rules.buy_require_momentum or momentum_ok):
event_type = "BUY"
reason = [f"综合分排名第 {rank}(≤买入阈值 {rules.buy_rank_threshold})"]
if rules.buy_require_trend:
reason.append(f"close > MA{rules.trend_ma}(趋势向上)")
if rules.buy_require_momentum:
reason.append("close > 20 日前收盘(动量为正)")
elif rank <= rules.buy_rank_threshold:
event_type = "WATCH"
reason = [f"综合分排名第 {rank}(买入区间)"]
if rules.buy_require_trend and not trend_ok:
reason.append(f"但 close < MA{rules.trend_ma}(趋势未确认)")
elif rank <= rules.sell_rank_threshold:
# 观望带
if rules.sell_on_trend_break and c is not None and ma is not None and c < ma:
event_type = "SELL"
reason = [f"跌破 MA{rules.trend_ma}(持仓者应卖出/减仓),rank={rank}"]
else:
event_type = "WATCH"
reason = [f"rank={rank}(买入区间外、卖出区间内:观望)"]
else:
event_type = "SELL"
reason = [f"综合分排名第 {rank}(>卖出阈值 {rules.sell_rank_threshold},持仓者应卖出)"]
if c is not None and c < 0:
continue # 防御负价
events.append(
SignalEvent(
symbol=sym,
signal_date=resolved,
signal_type=event_type,
score=round(float(sc), 6),
price=round(price, 4) if price is not None else None,
trigger_reason=reason,
)
)
if event_type == "BUY":
stats.buy += 1
elif event_type == "SELL":
stats.sell += 1
else:
stats.watch += 1
return SignalResult(
as_of_date=resolved,
rules=rules,
statistics=stats,
events=events,
config_snapshot={
"as_of": resolved.isoformat(),
"factors": [f.model_dump() for f in query.factors],
"rules": rules.model_dump(mode="json"),
},
)
def _num(v) -> float | None:
if v is None:
return None
try:
f = float(v)
except (TypeError, ValueError):
return None
return None if f != f else f # NaN → None
def _empty(query: SelectionQuery, rules: SignalRules, resolved: date) -> SignalResult:
return SignalResult(
as_of_date=resolved,
rules=rules,
statistics=SignalStatistics(),
events=[],
config_snapshot={"as_of": resolved.isoformat(), "rules": rules.model_dump(mode="json")},
)
+42
View File
@@ -0,0 +1,42 @@
"""Universe:选股/回测的股票范围执行器(ARCHITECTURE_v2 §14/§20 Universe 输入)。
把 ResearchService.filter_stocks 的语义规则化并集中于此:
- 当前日与历史日(as_of)都必须正确:退市股(delist < as_of)、上市时间(list_date)
- exclude_st 按**当前名称快照**含 ST 判定(历史可追溯数据;历史改名无法回溯,属近似,
见结果 unimplemented 说明)
- exclude_suspended 依赖停牌数据表(尚未建模),此处不做剔除,由上层显式标注
- symbols 白名单:非空时仅这些 symbol 参与(自选池 / 测试用)
"""
from __future__ import annotations
from collections.abc import Sequence
from datetime import date
from app.domain.entities.market import Stock
from app.domain.entities.research import UniverseSpec
def filter_stocks(
stocks: Sequence[Stock],
universe: UniverseSpec,
as_of: date,
) -> list[Stock]:
"""按股票池口径过滤,返回 as_of 时点应纳入的股票列表。"""
symbols = set(universe.symbols) if universe.symbols else None
out: list[Stock] = []
for s in stocks:
if symbols is not None and s.symbol not in symbols:
continue
if s.delist_date is not None and s.delist_date < as_of:
continue
if universe.exclude_st and s.name and "ST" in s.name.upper():
continue
if (
universe.min_listing_days
and s.list_date
and (as_of - s.list_date).days < universe.min_listing_days
):
continue
out.append(s)
return out
+1
View File
@@ -14,6 +14,7 @@ dependencies = [
# Qlib 研究引擎:PyPI 无 aarch64 wheel(见 docs/ROADMAP.md §2 备注),故从源码 git 安装并固定 commit。 # Qlib 研究引擎:PyPI 无 aarch64 wheel(见 docs/ROADMAP.md §2 备注),故从源码 git 安装并固定 commit。
# 网络下载困难时使用 HTTP 代理(见根目录 AGENT.md §0:192.168.1.160:3128)。 # 网络下载困难时使用 HTTP 代理(见根目录 AGENT.md §0:192.168.1.160:3128)。
"pyqlib @ git+https://github.com/microsoft/qlib.git@79633dd", "pyqlib @ git+https://github.com/microsoft/qlib.git@79633dd",
"pymysql>=1.1",
] ]
[project.optional-dependencies] [project.optional-dependencies]
+7 -1
View File
@@ -1,8 +1,10 @@
"""全局测试配置。 """全局测试配置。
- 强制 SQLite 临时文件数据库:默认库已是 MySQL(config.yaml database.mysql),
测试绝不允许触碰开发 MySQL(qlib@192.168.1.10);该 env 在 app 模块
首次导入前生效,SessionLocal/engine 会按它构造。
- JOB_MODE=local:单测/CI 里 Job 在本进程执行(不 spawn 子进程、不依赖安装态) - JOB_MODE=local:单测/CI 里 Job 在本进程执行(不 spawn 子进程、不依赖安装态)
- QLIB_SKIP_STALE_JOB_CLEANUP=1:TestClient 启动 lifespan 不清理残留 Job - QLIB_SKIP_STALE_JOB_CLEANUP=1:TestClient 启动 lifespan 不清理残留 Job
(避免连接/写入开发库 data/quant.db)
必须在任何 app / config 模块首次导入前生效,故放在 conftest 模块顶层。 必须在任何 app / config 模块首次导入前生效,故放在 conftest 模块顶层。
""" """
@@ -11,6 +13,10 @@ from __future__ import annotations
import os import os
# 测试数据库:每进程唯一 /tmp sqlite 文件(隔离、不落项目目录、不触 MySQL)
_TEST_DB = f"/tmp/qlib-pytest-{os.getpid()}.db"
os.environ["DATABASE_URL"] = f"sqlite:///{_TEST_DB}"
# 强制(不是 setdefault):测试绝不 spawn 研究子进程 / 不触碰开发库 # 强制(不是 setdefault):测试绝不 spawn 研究子进程 / 不触碰开发库
os.environ["JOB_MODE"] = "local" os.environ["JOB_MODE"] = "local"
os.environ["QLIB_SKIP_STALE_JOB_CLEANUP"] = "1" os.environ["QLIB_SKIP_STALE_JOB_CLEANUP"] = "1"
@@ -0,0 +1,94 @@
"""M8.5 Agent 新工具测试:screen_stocks / explain_selection / create_strategy。
使用 tmp SQLite(真实 SQLAlchemy repo)+ build_tools(factories) 直调工具。
"""
from __future__ import annotations
from datetime import date
import pytest
from app.agent.tools_impl import build_tools
from app.domain.entities.market import Stock
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
SqlAlchemyStockRepository,
)
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
_SYMS = ["600000.SH", "600001.SH", "600002.SH", "600003.SH", "600004.SH"]
@pytest.fixture()
def tools(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'agent.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
df = synthetic_daily({s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)}, n=320)
with Session() as session:
SqlAlchemyStockRepository(session).upsert_many(
[
Stock(symbol=s, name=f"测试股份{i}", list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)
]
)
SqlAlchemyDailyBarRepository(session).upsert_many(bars_dataframe_to_daily_bars(df))
session.commit()
factories = {
"session_factory": lambda: Session(),
"stock_repo_factory": lambda s: SqlAlchemyStockRepository(s),
"daily_repo_factory": lambda s: SqlAlchemyDailyBarRepository(s),
"experiment_repo_factory": lambda s: None,
"engine": None,
}
return build_tools(factories=factories)
def _invoke(tools, name: str, args: dict) -> str:
tool = next(t for t in tools if t.name == name)
return tool.invoke(args)
class TestAgentSelectionTools:
def test_screen_stocks(self, tools) -> None:
out = _invoke(
tools,
"screen_stocks",
{
"factors": "momentum_60",
"top_n": 3,
"as_of": "2024-12-31",
"symbols": ",".join(_SYMS),
},
)
assert "600000.SH" in out or "600001.SH" in out
assert "score=" in out
def test_screen_stocks_empty_scope(self, tools) -> None:
out = _invoke(tools, "screen_stocks", {"symbols": "600999.SH", "as_of": "2024-12-31"})
assert "无候选" in out or "600999" in out
def test_create_strategy_and_explain_flow(self, tools) -> None:
out = _invoke(
tools,
"create_strategy",
{"name": "Agent 策略", "factors": "momentum_60", "top_n": 5},
)
assert "策略已保存" in out
# 重名被拒(工具返回错误消息而非崩溃)
out2 = _invoke(tools, "create_strategy", {"name": "Agent 策略", "factors": "momentum_20"})
assert "已存在" in out2 or "失败" in out2
def test_explain_selection_missing(self, tools) -> None:
out = _invoke(tools, "explain_selection", {"selection_id": "SEL-NOPE"})
assert "不存在" in out
def test_tool_names_registered(self, tools) -> None:
names = {t.name for t in tools}
assert {"screen_stocks", "explain_selection", "generate_signals", "create_strategy"} <= names
+17 -2
View File
@@ -11,6 +11,7 @@ from datetime import date
import pytest import pytest
from app.api import deps from app.api import deps
from app.domain.entities.market import Stock from app.domain.entities.market import Stock
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.main import app from app.main import app
from app.quant.engine import LocalEngine from app.quant.engine import LocalEngine
from app.quant.service import ResearchService from app.quant.service import ResearchService
@@ -41,13 +42,13 @@ def _mem_stocks() -> list[Stock]:
@pytest.fixture() @pytest.fixture()
def client() -> TestClient: def client(tmp_path) -> TestClient:
drifts = {sym: 0.003 - 0.0015 * i for i, sym in enumerate(_SYMS)} drifts = {sym: 0.003 - 0.0015 * i for i, sym in enumerate(_SYMS)}
daily_df = synthetic_daily(drifts, n=300) daily_df = synthetic_daily(drifts, n=300)
bars = bars_dataframe_to_daily_bars(daily_df) bars = bars_dataframe_to_daily_bars(daily_df)
class _MemDailyRepo: class _MemDailyRepo:
def get_range_many(self, symbols, start, end): def get_range_many(self, symbols, start, end, adjust="none"):
out = [] out = []
for b in bars: for b in bars:
if b.symbol in symbols and start <= b.trade_date <= end: if b.symbol in symbols and start <= b.trade_date <= end:
@@ -60,6 +61,20 @@ def client() -> TestClient:
service = ResearchService(_MemStockRepo(_mem_stocks()), _MemDailyRepo(), LocalEngine()) service = ResearchService(_MemStockRepo(_mem_stocks()), _MemDailyRepo(), LocalEngine())
app.dependency_overrides[deps._stock_repo_factory] = lambda: _MemStockRepo(_mem_stocks()) # noqa: SLF001 app.dependency_overrides[deps._stock_repo_factory] = lambda: _MemStockRepo(_mem_stocks()) # noqa: SLF001
app.dependency_overrides[deps._service_factory] = lambda: service # noqa: SLF001 app.dependency_overrides[deps._service_factory] = lambda: service # noqa: SLF001
# /api/factors 自 M7.1 读 DB(factor_definition)→ 提供 tmp sqlite session
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
SessionLocal = sessionmaker(bind=engine, expire_on_commit=False)
def _session_override():
with SessionLocal() as s:
yield s
app.dependency_overrides[deps.get_session] = _session_override
with TestClient(app) as c: with TestClient(app) as c:
yield c yield c
app.dependency_overrides.clear() app.dependency_overrides.clear()
+125
View File
@@ -0,0 +1,125 @@
"""M7.2b 因子组合测试:repo CRUD(幂等/去重)+ /api/composites(含方向填充/404/删除)。"""
from __future__ import annotations
import pytest
from app.api import deps
from app.domain.entities.composite import CompositeComponent, CompositeDefinition
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import (
SqlAlchemyCompositeRepository,
)
from app.main import app
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
@pytest.fixture()
def session(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'c.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
with Session() as s:
yield s
def _comp(name="质量动量", **kw) -> CompositeDefinition:
base = dict(
name=name,
description="测试组合",
components=[
CompositeComponent(name="momentum_60", weight=0.7, direction="higher_is_better"),
CompositeComponent(name="volatility_60", weight=0.3, direction="lower_is_better"),
],
)
base.update(kw)
return CompositeDefinition(**base)
class TestCompositeRepository:
def test_save_get_list(self, session) -> None:
repo = SqlAlchemyCompositeRepository(session)
saved = repo.save(_comp().model_copy(update={"id": "CF-TEST-1"}))
session.commit()
assert saved.id == "CF-TEST-1"
got = repo.get("CF-TEST-1")
assert got is not None and got.name == "质量动量"
assert len(got.components) == 2
assert repo.list()[0].components[0].direction == "higher_is_better"
def test_duplicate_name_rejected(self, session) -> None:
repo = SqlAlchemyCompositeRepository(session)
repo.save(_comp().model_copy(update={"id": "CF-A"}))
session.commit()
with pytest.raises(ValueError):
repo.save(_comp().model_copy(update={"id": "CF-B"})) # 同名不同 id
def test_delete(self, session) -> None:
repo = SqlAlchemyCompositeRepository(session)
repo.save(_comp().model_copy(update={"id": "CF-X"}))
session.commit()
assert repo.delete("CF-X") is True
session.commit()
assert repo.get("CF-X") is None
assert repo.delete("CF-X") is False
@pytest.fixture()
def client(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
def _session_override():
with Session() as s:
yield s
app.dependency_overrides[deps.get_session] = _session_override
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
class TestCompositesApi:
def test_crud(self, client) -> None:
body = {
"name": "质量动量组合",
"description": "动量+低波",
"components": [
{"name": "momentum_60", "weight": 0.7},
{"name": "volatility_60", "weight": 0.3},
],
}
created = client.post("/api/composites", json=body)
assert created.status_code == 200
cid = created.json()["id"]
assert cid.startswith("CF-")
# 方向由注册表自动填充
dirs = {c["name"]: c["direction"] for c in created.json()["components"]}
assert dirs == {"momentum_60": "higher_is_better", "volatility_60": "lower_is_better"}
rows = client.get("/api/composites").json()
assert len(rows) == 1
detail = client.get(f"/api/composites/{cid}").json()
assert detail["name"] == "质量动量组合"
assert client.delete(f"/api/composites/{cid}").status_code == 200
assert client.get(f"/api/composites/{cid}").status_code == 404
assert client.delete(f"/api/composites/{cid}").status_code == 404
def test_unknown_factor_rejected(self, client) -> None:
resp = client.post(
"/api/composites",
json={
"name": "坏组合",
"components": [{"name": "no_such_factor", "weight": 1}],
},
)
assert resp.status_code == 400
assert "no_such_factor" in resp.json()["detail"]
def test_duplicate_name_400(self, client) -> None:
body = {"name": "同名", "components": [{"name": "momentum_60", "weight": 1}]}
assert client.post("/api/composites", json=body).status_code == 200
assert client.post("/api/composites", json=body).status_code == 400
+63 -14
View File
@@ -1,9 +1,10 @@
"""配置层测试:默认 SQLite 相对路径解析、config.yaml / .env 加载约定。""" """配置层测试:默认 MySQL URL 组装(config.yaml database.mysql)、.env 加载约定、
SQLite 兜底逻辑。conftest 已强制 DATABASE_URL=tmp sqlite,本文件内用
monkeypatch/cache_clear 单独验证「无 env 时走 mysql 段」的分支(只组装不连接)。"""
from __future__ import annotations from __future__ import annotations
import os import os
from pathlib import Path
from app.core.config import PROJECT_ROOT, _load_env_file, get_settings from app.core.config import PROJECT_ROOT, _load_env_file, get_settings
@@ -14,14 +15,62 @@ def test_project_root_points_to_repo_root() -> None:
assert (PROJECT_ROOT / "backend").is_dir() assert (PROJECT_ROOT / "backend").is_dir()
def test_default_database_url_resolves_to_project_data_dir() -> None: def test_test_runner_isolation_uses_tmp_sqlite() -> None:
"""conftest 强制每进程唯一 /tmp sqlite,测试绝不触碰 MySQL 开发库。"""
settings = get_settings() settings = get_settings()
assert settings.database_url.startswith("sqlite:///") assert settings.database_url.startswith("sqlite:////tmp/qlib-pytest-")
# 相对路径应解析到 <项目根>/data/quant.db
assert settings.database_url.endswith("/data/quant.db")
db_path = Path(settings.database_url.removeprefix("sqlite:///")) def test_mysql_default_from_config_yaml(monkeypatch) -> None:
assert db_path.is_absolute() """未设 DATABASE_URL 时,config.yaml database.mysql 段组装 MySQL URL(不连接)。"""
assert db_path == PROJECT_ROOT / "data" / "quant.db" monkeypatch.delenv("DATABASE_URL", raising=False)
get_settings.cache_clear()
try:
url = get_settings().database_url
finally:
get_settings.cache_clear()
assert url.startswith("mysql+pymysql://qlib:")
assert "192.168.1.10:3306/qlib?charset=utf8mb4" in url
def test_build_mysql_url(monkeypatch) -> None:
from app.core.config import _build_mysql_url
monkeypatch.setenv("MYSQL_PASSWORD", "p@ss:word")
cfg = {
"enabled": True,
"host": "10.0.0.2",
"port": 3307,
"db": "q",
"user": "u",
"password_env": "MYSQL_PASSWORD",
}
url = _build_mysql_url(cfg)
assert url == "mysql+pymysql://u:p%40ss%3Aword@10.0.0.2:3307/q?charset=utf8mb4"
# 未启用 / 缺字段 → None(回退 sqlite 兜底)
assert _build_mysql_url({**cfg, "enabled": False}) is None
assert _build_mysql_url({**cfg, "host": None}) is None
assert _build_mysql_url(None) is None
def test_mysql_url_password_optional(monkeypatch) -> None:
"""密码留空(未设 env)也可组装 URL —— 方便仅内网/免密场景。"""
from app.core.config import _build_mysql_url
monkeypatch.delenv("MYSQL_PASSWORD", raising=False)
cfg = {"enabled": True, "host": "h", "db": "d", "user": "u"}
assert _build_mysql_url(cfg) == "mysql+pymysql://u@h:3306/d?charset=utf8mb4"
def test_sqlite_url_normalization_and_fallback() -> None:
"""仅 sqlite 相对路径被解析为项目根绝对路径;其它 URL 原样透传。"""
from app.core.config import _normalize_sqlite_url
assert _normalize_sqlite_url("mysql+pymysql://u:p@h/d") == "mysql+pymysql://u:p@h/d"
abs_url = _normalize_sqlite_url("sqlite:///./data/quant.db")
assert abs_url.startswith("sqlite:///")
# 绝对路径 sqlite 不再重复解析
assert _normalize_sqlite_url("sqlite:////abs/x.db") == "sqlite:////abs/x.db"
def test_settings_loaded_from_config_yaml() -> None: def test_settings_loaded_from_config_yaml() -> None:
@@ -32,7 +81,7 @@ def test_settings_loaded_from_config_yaml() -> None:
assert settings.data_source_fallback == "sina" assert settings.data_source_fallback == "sina"
def test_env_file_loading(tmp_path: Path, monkeypatch) -> None: def test_env_file_loading(tmp_path, monkeypatch) -> None:
monkeypatch.delenv("TUSHARE_TOKEN", raising=False) monkeypatch.delenv("TUSHARE_TOKEN", raising=False)
env_file = tmp_path / ".env" env_file = tmp_path / ".env"
env_file.write_text( env_file.write_text(
@@ -57,14 +106,14 @@ def test_llm_from_config_yaml() -> None:
def test_llm_env_overrides_yaml(monkeypatch) -> None: def test_llm_env_overrides_yaml(monkeypatch) -> None:
from app.core.config import get_settings
monkeypatch.setenv("LLM_MODEL", "env-model") monkeypatch.setenv("LLM_MODEL", "env-model")
monkeypatch.setenv("LLM_BASE_URL", "https://example.com/v1") monkeypatch.setenv("LLM_BASE_URL", "https://example.com/v1")
monkeypatch.setenv("LLM_API_KEY", "sk-test") monkeypatch.setenv("LLM_API_KEY", "sk-test")
get_settings.cache_clear() # Settings 有 lru_cache,刷新以读取新环境变量 get_settings.cache_clear() # Settings 有 lru_cache,刷新以读取新环境变量
s = get_settings() try:
s = get_settings()
finally:
get_settings.cache_clear()
assert s.llm_model == "env-model" assert s.llm_model == "env-model"
assert s.llm_base_url == "https://example.com/v1" assert s.llm_base_url == "https://example.com/v1"
assert s.llm_api_key == "sk-test" assert s.llm_api_key == "sk-test"
get_settings.cache_clear()
+97
View File
@@ -0,0 +1,97 @@
"""M7.1 因子目录测试:factor_definition 落库(幂等 upsert)+ /api/factors 读库 + seed。
repo 测试走 tmp SQLite;API 测试 override get_session 到 tmp sqlite 种子库。
"""
from __future__ import annotations
import pytest
from app.api import deps
from app.domain.entities.factor import FactorDefinition
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
SqlAlchemyFactorRepository,
)
from app.main import app
from app.quant.factors import list_factors
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
@pytest.fixture()
def session(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'factor.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
with Session() as s:
yield s
class TestFactorRepository:
def test_upsert_idempotent_and_roundtrip(self, session) -> None:
repo = SqlAlchemyFactorRepository(session)
d = FactorDefinition(
name="test_momentum", description="测试", formula="x", brief="b",
lookback=10, requires=["close", "high"],
)
assert repo.upsert_many([d]) == 1
session.commit()
assert len(repo.list()) == 1
got = repo.get("test_momentum")
assert got is not None and got.requires == ["close", "high"]
# 幂等更新
repo.upsert_many([d.model_copy(update={"description": "更新"})])
session.commit()
assert repo.get("test_momentum").description == "更新"
assert len(repo.list()) == 1
def test_seed_from_registry(self, session) -> None:
repo = SqlAlchemyFactorRepository(session)
defs = [FactorDefinition.from_registry_def(d) for d in list_factors()]
assert repo.upsert_many(defs) == len(defs)
session.commit()
names = {f.name for f in repo.list()}
assert len(names) == len(defs)
# 与注册表一致
assert names == {d.name for d in list_factors()}
assert repo.get("momentum_60").direction == "higher_is_better"
@pytest.fixture()
def client(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
def _session_override():
with Session() as s:
yield s
app.dependency_overrides[deps.get_session] = _session_override
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
class TestFactorsApi:
def test_list_seeds_and_reads_db(self, client) -> None:
resp = client.get("/api/factors")
assert resp.status_code == 200
rows = resp.json()
assert isinstance(rows, list) and len(rows) >= 9
first = next(r for r in rows if r["name"] == "momentum_60")
# DB 契约源字段齐全(与前端 FactorMeta 匹配 + version)
assert set(first.keys()) >= {
"name", "description", "brief", "formula", "frequency",
"lookback", "direction", "requires", "version",
}
assert "close" in first["requires"]
def test_list_matches_registry_after_seed(self, client) -> None:
"""seed 后目录 == 代码注册表集合(无额外未知项)。"""
client.get("/api/factors") # 首次访问触发 seed
client.get("/api/factors") # 幂等:二次访问不报错、不重复
resp = client.get("/api/factors")
names = {r["name"] for r in resp.json()}
assert names == {d.name for d in list_factors()}
+1 -1
View File
@@ -59,7 +59,7 @@ class _FakeDailyRepo:
def __init__(self, bars): def __init__(self, bars):
self._bars = bars self._bars = bars
def get_range_many(self, symbols, start, end): def get_range_many(self, symbols, start, end, adjust="none"):
return [b for b in self._bars if b.symbol in symbols and start <= b.trade_date <= end] return [b for b in self._bars if b.symbol in symbols and start <= b.trade_date <= end]
def get_range(self, symbol, start, end): def get_range(self, symbol, start, end):
+106
View File
@@ -0,0 +1,106 @@
"""M7.3 行情口径测试:Repository 读路径按 adjust 过滤(不复权主口径),spec 记录口径。
混合行场景:同 symbol/date 存在 tushare/none 与 sina/qfq 行时,研究读取
(get_range_many / stream)默认只取 adjust=none —— 消除「混合口径污染因子」风险。
"""
from __future__ import annotations
from datetime import date
from decimal import Decimal
from app.domain.entities.market import DailyBar
from app.domain.entities.research import ResearchSpec
from app.domain.entities.selection import SelectionQuery
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
)
from sqlalchemy import create_engine
from sqlalchemy.orm import Session
_D = date(2024, 6, 3)
_D2 = date(2024, 6, 4)
_D3 = date(2024, 6, 5)
def _bar(adjust: str, close: str, source: str = "tushare", day=None) -> DailyBar:
return DailyBar(
symbol="600519.SH",
trade_date=day or _D,
source=source,
adjust=adjust,
open=Decimal("100"), high=Decimal("101"), low=Decimal("99"),
close=Decimal(close), volume=Decimal("1000"), amount=Decimal("100000"),
)
class TestAdjustFilter:
def _session(self, tmp_path) -> Session:
engine = create_engine(f"sqlite:///{tmp_path / 'adj.db'}", future=True)
Base.metadata.create_all(engine)
return Session(engine)
def test_get_range_many_filters_adjust(self, tmp_path) -> None:
with self._session(tmp_path) as session:
repo = SqlAlchemyDailyBarRepository(session)
# 唯一键 (symbol, trade_date):同键共存不可能 —— 用连续三天模拟
# none 主口径两天 + sina/qfq 兜底一天
repo.upsert_many(
[
_bar("none", "1700", day=_D),
_bar("none", "1710", day=_D2),
_bar("qfq", "1680", source="sina", day=_D3),
]
)
session.commit()
none_rows = repo.get_range_many(["600519.SH"], _D, _D3, adjust="none")
assert len(none_rows) == 2 and {float(r.close) for r in none_rows} == {1700, 1710}
qfq_rows = repo.get_range_many(["600519.SH"], _D, _D3, adjust="qfq")
assert len(qfq_rows) == 1 and float(qfq_rows[0].close) == 1680
def test_stream_filters_adjust(self, tmp_path) -> None:
with self._session(tmp_path) as session:
repo = SqlAlchemyDailyBarRepository(session)
repo.upsert_many(
[
_bar("none", "1700", day=_D),
_bar("none", "1710", day=_D2),
_bar("qfq", "1680", source="sina", day=_D3),
]
)
session.commit()
rows = list(
repo.stream_range_many_columns(
["600519.SH"], _D, _D3, ["close"], adjust="none"
)
)
assert len(rows) == 2 and {float(r[-1]) for r in rows} == {1700, 1710}
qrows = list(
repo.stream_range_many_columns(
["600519.SH"], _D, _D3, ["close"], adjust="qfq"
)
)
assert len(qrows) == 1 and float(qrows[0][-1]) == 1680
class TestSpecRecordsAdjustment:
def test_research_spec_default_and_field(self) -> None:
spec = ResearchSpec(
type="backtest",
factors=[{"name": "momentum_60", "weight": 1}],
period=(date(2024, 1, 1), date(2024, 6, 1)),
)
assert spec.price_adjustment == "none"
snap = spec.model_dump(mode="json")
assert snap["price_adjustment"] == "none" # 结果 config_snapshot 可溯源
def test_selection_query_adjustment_in_snapshot(self) -> None:
q = SelectionQuery(factors=[{"name": "momentum_60", "weight": 1}], top_n=5)
assert q.price_adjustment == "none"
assert q.model_dump()["price_adjustment"] == "none"
q2 = SelectionQuery(
factors=[{"name": "momentum_60", "weight": 1}], top_n=5, price_adjustment="qfq"
)
assert q2.price_adjustment == "qfq"
+3 -2
View File
@@ -14,9 +14,10 @@ from app.domain.entities.research import (
SelectionSpec, SelectionSpec,
UniverseSpec, UniverseSpec,
) )
from app.quant.composite import composite_score, cross_sectional_zscore
from app.quant.evaluation import run_factor_test from app.quant.evaluation import run_factor_test
from app.quant.factors import compute_factor from app.quant.factors import compute_factor
from app.quant.local_engine import composite_score, cross_sectional_zscore, rebalance_dates from app.quant.local_engine import rebalance_dates
from pydantic import ValidationError from pydantic import ValidationError
from conftest_quant import synthetic_daily from conftest_quant import synthetic_daily
@@ -115,7 +116,7 @@ class TestEvaluation:
class TestSingleStockDegradation: class TestSingleStockDegradation:
def test_zscore_single_stock_keeps_candidate(self) -> None: def test_zscore_single_stock_keeps_candidate(self) -> None:
from app.quant.local_engine import cross_sectional_zscore from app.quant.composite import cross_sectional_zscore
daily = synthetic_daily({"ONLY": 0.001}, n=80) daily = synthetic_daily({"ONLY": 0.001}, n=80)
_d, panel = compute_factor("momentum_20", daily) _d, panel = compute_factor("momentum_20", daily)
+223
View File
@@ -0,0 +1,223 @@
"""M6.0 选股契约与服务测试:因子评分 TopN 选股、as_of 未来函数防护、universe 过滤。
使用 Fake(内存)Repository + 合成/自定义行情,不触数据库。
"""
from __future__ import annotations
from datetime import date
from decimal import Decimal
import pandas as pd
import pytest
from app.application.services.selection_service import SelectionService
from app.domain.entities.market import DailyBar, Stock
from app.domain.entities.selection import SelectionQuery
from pydantic import ValidationError
from conftest_quant import synthetic_daily
_SYMS = ["600000.SH", "600001.SH", "600002.SH", "600003.SH", "600004.SH"]
def _mem_stocks() -> list[Stock]:
return [
Stock(symbol=s, name=f"测试股份{i}", list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)
]
class _MemStockRepo:
def __init__(self, stocks: list[Stock]) -> None:
self._stocks = stocks
def list(self) -> list[Stock]:
return self._stocks
def get_by_symbol(self, symbol: str) -> Stock | None:
return next((s for s in self._stocks if s.symbol == symbol), None)
class _MemDailyRepo:
"""内存日线仓库:支持 get_range / get_range_many(无流式 → 走回退路径)。"""
def __init__(self, df: pd.DataFrame) -> None:
self._df = df
def _bars(self, symbols, start, end) -> list[DailyBar]:
sub = self._df[
self._df["symbol"].isin(symbols)
& (self._df["trade_date"] >= start)
& (self._df["trade_date"] <= end)
]
out: list[DailyBar] = []
for r in sub.itertuples():
out.append(
DailyBar(
symbol=r.symbol,
trade_date=r.trade_date,
open=Decimal(str(r.open)),
high=Decimal(str(r.high)),
low=Decimal(str(r.low)),
close=Decimal(str(r.close)),
volume=Decimal(str(r.volume)),
amount=Decimal(str(r.amount)),
)
)
return out
def get_range(self, symbol, start, end) -> list[DailyBar]:
return self._bars([symbol], start, end)
def get_range_many(self, symbols, start, end, adjust="none") -> list[DailyBar]:
return self._bars(list(symbols), start, end)
def latest_date(self, symbol: str) -> date | None:
sub = self._df[self._df["symbol"] == symbol]
return None if sub.empty else sub["trade_date"].max()
@pytest.fixture()
def svc():
drifts = {sym: 0.006 - 0.0015 * i for i, sym in enumerate(_SYMS)}
df = synthetic_daily(drifts, n=320)
return SelectionService(_MemStockRepo(_mem_stocks()), _MemDailyRepo(df))
def _score_query(**kw) -> SelectionQuery:
base = dict(
factors=[{"name": "momentum_60", "weight": 1.0}],
top_n=3,
as_of=date(2024, 12, 31),
)
base.update(kw)
return SelectionQuery(**base)
class TestScoreSelection:
def test_returns_top_n_with_rank_reason(self, svc) -> None:
res = svc.select(_score_query())
assert res.method == "score"
assert len(res.candidates) == 3
assert [c.rank for c in res.candidates] == [1, 2, 3]
# 分数降序
scores = [c.score for c in res.candidates]
assert scores == sorted(scores, reverse=True)
# 每个候选带因子值与理由(可解释)
for c in res.candidates:
assert c.symbol in _SYMS
assert "momentum_60" in c.factor_values
assert any("momentum_60" in r for r in c.selection_reason)
# 统计:评估数 > 0
assert res.statistics.evaluated >= 3
assert res.statistics.selected == 3
assert res.as_of_date <= date(2024, 12, 31)
def test_highest_drift_ranked_high(self, svc) -> None:
res = svc.select(_score_query(top_n=1))
# 最高漂移股票(600000.SH)应位列前三(动量因子对强趋势敏感)
assert res.candidates[0].symbol in _SYMS[:3]
def test_top_pct(self, svc) -> None:
res = svc.select(_score_query(top_n=None, top_pct=0.4))
assert 1 <= len(res.candidates) <= 3 # 5 只的 40% ≈ 2
assert res.candidates[0].rank == 1
def test_min_score_filters(self, svc) -> None:
res = svc.select(_score_query(min_score=10_000))
# 合成数据最高动量约 +60%,min_score=10000 应无候选
assert res.candidates == []
class TestAsOfNoFutureFunction:
def _two_stage_df(self) -> tuple[pd.DataFrame, date, date]:
"""X 前段强、Y 前段平后段暴涨:as_of=split 时 Y 的 60 日动量应为 0/无。"""
dates = pd.bdate_range("2024-01-01", periods=200)
split = dates[100].date() # d101:Y 从这天开始暴涨
rows: list[dict] = []
for d in dates:
# X:恒定 +0.8%/日
rows.append({"symbol": "600000.SH", "trade_date": d.date(),
"close": 100 * 1.008 ** ((d - dates[0]).days)})
y_price = 100.0
for d in dates:
if d.date() > split:
y_price = y_price * 1.05
rows.append({"symbol": "600001.SH", "trade_date": d.date(), "close": y_price})
df = pd.DataFrame(rows)
for col in ("open", "high", "low", "volume", "amount"):
df[col] = df["close"] * 1.001 if col != "volume" else 1_000_000
return df, split, dates[-1].date()
def test_as_of_cutoff_excludes_future(self) -> None:
df, split, _end = self._two_stage_df()
stocks = [
Stock(symbol="600000.SH", name="A", list_date=date(1999, 1, 1)),
Stock(symbol="600001.SH", name="B", list_date=date(1999, 1, 1)),
]
svc = SelectionService(_MemStockRepo(stocks), _MemDailyRepo(df))
# 历史选股:as_of=split(Y 尚未暴涨)→ Y 不应排第一(其 60 日动量≈0/缺失)
res = svc.select(_score_query(factors=[{"name": "momentum_60", "weight": 1.0}], as_of=split))
assert res.candidates
assert res.candidates[0].symbol == "600000.SH"
# Y 若进入候选,其理由值应接近 0(而非泄漏未来暴涨的巨幅动量)
y = next((c for c in res.candidates if c.symbol == "600001.SH"), None)
if y is not None:
assert y.factor_values["momentum_60"] < 0.1
class TestUniverseFilter:
def test_exclude_st(self, svc) -> None:
stocks = _mem_stocks()[:2]
stocks[0] = stocks[0].model_copy(update={"name": "ST 风险股份"})
s = SelectionService(_MemStockRepo(stocks), svc._daily_repo)
res = s.select(_score_query())
assert all(c.symbol != "600000.SH" for c in res.candidates)
def test_min_listing_days(self) -> None:
stocks = [
Stock(symbol=_SYMS[0], name="新上市", list_date=date(2024, 12, 1)),
Stock(symbol=_SYMS[1], name="老股", list_date=date(1999, 1, 1)),
]
df = synthetic_daily({_SYMS[0]: 0.01, _SYMS[1]: 0.001}, n=320)
s = SelectionService(_MemStockRepo(stocks), _MemDailyRepo(df))
res = s.select(_score_query(as_of=date(2024, 12, 31)))
assert all(c.symbol != _SYMS[0] for c in res.candidates) # 上市不足 250 自然日被滤除
def test_delisted_before_as_of(self) -> None:
stocks = [
Stock(symbol=_SYMS[0], name="已退市", list_date=date(1990, 1, 1),
delist_date=date(2023, 6, 30)),
Stock(symbol=_SYMS[1], name="正常", list_date=date(1999, 1, 1)),
]
df = synthetic_daily({_SYMS[0]: 0.01, _SYMS[1]: 0.001}, n=320)
s = SelectionService(_MemStockRepo(stocks), _MemDailyRepo(df))
res = s.select(_score_query(as_of=date(2024, 6, 30)))
assert all(c.symbol != _SYMS[0] for c in res.candidates)
class TestEmptyAndValidation:
def test_no_stocks_returns_empty(self) -> None:
s = SelectionService(_MemStockRepo([]), _MemDailyRepo(pd.DataFrame()))
res = s.select(_score_query())
assert res.candidates == []
assert res.statistics.selected == 0
def test_query_validation(self) -> None:
with pytest.raises(ValidationError):
SelectionQuery(method="score", top_n=5) # score 缺 factors
with pytest.raises(ValidationError):
SelectionQuery(method="condition", conditions=[], top_n=5) # condition 缺 conditions
with pytest.raises(ValidationError):
SelectionQuery(factors=[{"name": "a", "weight": 1}], top_n=None, top_pct=None)
with pytest.raises(ValidationError):
SelectionQuery(
factors=[{"name": "a", "weight": 1}, {"name": "a", "weight": 2}], top_n=5
)
# 合法
SelectionQuery(factors=[{"name": "a", "weight": 1}], top_n=5)
def test_query_top_pct_valid(self) -> None:
q = SelectionQuery(factors=[{"name": "momentum_60", "weight": 1}], top_pct=0.5)
assert q.top_n is None and q.top_pct == 0.5
@@ -0,0 +1,148 @@
"""M6.4 一致性回归:回测(TopKBacktestRunner)与独立 select(as_of) 使用同一评分引擎。
v2 §25/§27 红线验证:对任意调仓日 d,SelectionService.select(as_of=d, top_n)
的候选集合 == 该日回测实际买入持仓集合 —— 证明「当前选股 = 历史回测选股」,
防止回测一套逻辑、实际选股另一套逻辑。
"""
from __future__ import annotations
from datetime import date
import pandas as pd
import pytest
from app.application.services.selection_service import SelectionService
from app.domain.entities.market import Stock
from app.domain.entities.research import ResearchSpec
from app.domain.entities.selection import SelectionQuery
from app.quant.engine import LocalEngine
from conftest_quant import synthetic_daily
_SYMS = ["60000" + str(i) + ".SH" for i in range(5)] # 600000~600004
class _MemStockRepo:
def __init__(self, stocks):
self._stocks = stocks
def list(self):
return self._stocks
def get_by_symbol(self, symbol):
return next((s for s in self._stocks if s.symbol == symbol), None)
class _MemDailyRepo:
def __init__(self, df: pd.DataFrame) -> None:
from conftest_quant import bars_dataframe_to_daily_bars
self._bars = bars_dataframe_to_daily_bars(df)
def get_range(self, symbol, start, end):
return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end]
def get_range_many(self, symbols, start, end, adjust="none"):
syms = set(symbols)
return [b for b in self._bars if b.symbol in syms and start <= b.trade_date <= end]
def latest_date(self, symbol):
rows = [b.trade_date for b in self._bars if b.symbol == symbol]
return max(rows) if rows else None
@pytest.fixture()
def daily_df() -> pd.DataFrame:
drifts = {s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)}
return synthetic_daily(drifts, n=320) # 2024-01-01 起 ~320 交易日
def _spec(**kw) -> ResearchSpec:
base = dict(
type="backtest",
universe={"exclude_st": False, "min_listing_days": 0},
factors=[{"name": "momentum_60", "weight": 1.0}],
selection={"top_n": 2},
rebalance="monthly",
period=(date(2024, 5, 1), date(2024, 12, 31)),
)
base.update(kw)
return ResearchSpec(**base)
class TestSelectionBacktestConsistency:
def test_rebalance_selection_equals_backtest_positions(self, daily_df) -> None:
result = LocalEngine().run_backtest(daily_df, _spec())
stocks = [
Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1)) for i, s in enumerate(_SYMS)
]
svc = SelectionService(_MemStockRepo(stocks), _MemDailyRepo(daily_df))
# 回测每个调仓日的实际持仓 → 与 select(as_of=该日) 的 TopN 候选一致
by_date: dict[date, set[str]] = {}
for p in result.positions:
by_date.setdefault(p.date, set()).add(p.symbol)
assert len(by_date) >= 5 # 月调仓多个时点
for d, held in sorted(by_date.items()):
res = svc.select(
SelectionQuery(
universe=_spec().universe,
factors=[{"name": "momentum_60", "weight": 1.0}],
top_n=2,
as_of=d,
)
)
picked = {c.symbol for c in res.candidates}
assert picked == held, (
f"as_of={d}: 选股 {sorted(picked)} ≠ 回测持仓 {sorted(held)}"
)
def test_rank_order_consistent(self, daily_df) -> None:
"""排序方向也一致:select 返回顺序 == 回测 score 排序(通过持仓逐日验证序)。"""
spec = _spec()
result = LocalEngine().run_backtest(daily_df, spec)
stocks = [
Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1)) for i, s in enumerate(_SYMS)
]
svc = SelectionService(_MemStockRepo(stocks), _MemDailyRepo(daily_df))
by_date: dict[date, list[str]] = {}
for p in result.positions:
by_date.setdefault(p.date, []).append(p.symbol)
# 只验证任一日的一致性集合(顺序由 TopK 权重决定,与评分排序一一对应)
d, held = next(iter(by_date.items()))
res = svc.select(
SelectionQuery(
universe=spec.universe,
factors=[{"name": "momentum_60", "weight": 1.0}],
top_n=len(held),
as_of=d,
)
)
assert [c.symbol for c in res.candidates] == sorted(
held, key=lambda s: res.candidates[[x.symbol for x in res.candidates].index(s)].score,
reverse=True,
)
class TestPortfolioEngine:
def test_equal_weight_default_unchanged(self, daily_df) -> None:
"""新增 PortfolioSpec 后默认配置回测结果与未设置前一致(回归由本文件首测已锁数值)。"""
from app.domain.entities.research import PortfolioSpec
from app.quant.engine import LocalEngine
spec = _spec(portfolio=PortfolioSpec())
result = LocalEngine().run_backtest(daily_df, spec)
assert result.summary.total_trades >= 0
# 未设约束 → 无组合约束说明
assert not any("约束未建模" in u for u in result.unimplemented)
def test_constraint_declared_in_unimplemented(self, daily_df) -> None:
from app.domain.entities.research import PortfolioSpec
from app.quant.engine import LocalEngine
spec = _spec(portfolio=PortfolioSpec(max_position_pct=0.1))
result = LocalEngine().run_backtest(daily_df, spec)
assert any("最大单股权重" in u for u in result.unimplemented)
# config_snapshot 记录组合配置
assert result.config_snapshot["portfolio"]["max_position_pct"] == 0.1
+228
View File
@@ -0,0 +1,228 @@
"""M6.2 条件选股测试:结构化条件(static.*/tech 字段与因子/fundamental.*)与未来函数防护。
Fake(内存)Repository;财务行显式携带 announce_date 验证 as_of 可见性。
"""
from __future__ import annotations
from datetime import date
from decimal import Decimal
import pandas as pd
import pytest
from app.application.services.selection_service import SelectionService
from app.domain.entities.market import FinancialIndicator, Stock
from app.domain.entities.research import UniverseSpec
from app.domain.entities.selection import SelectionQuery
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
_SYMS = ["600000.SH", "600001.SH", "600002.SH"]
def _stocks(industries: dict[str, str] | None = None) -> list[Stock]:
industries = industries or {}
return [
Stock(
symbol=s, name=f"股{i}", industry=industries.get(s),
list_date=date(1999, 1, 1),
)
for i, s in enumerate(_SYMS)
]
class _MemDailyRepo:
def __init__(self, df: pd.DataFrame) -> None:
self._bars_all = bars_dataframe_to_daily_bars(df)
def get_range(self, symbol, start, end):
return [b for b in self._bars_all if b.symbol == symbol and start <= b.trade_date <= end]
def get_range_many(self, symbols, start, end, adjust="none"):
syms = set(symbols)
return [b for b in self._bars_all if b.symbol in syms and start <= b.trade_date <= end]
def latest_date(self, symbol):
rows = [b.trade_date for b in self._bars_all if b.symbol == symbol]
return max(rows) if rows else None
class _MemFinancialRepo:
def __init__(self, rows: list[FinancialIndicator]) -> None:
self._rows = rows
def list_announced(self, symbol, as_of_date, report_start=None):
return [
r for r in self._rows
if r.symbol == symbol and r.announce_date <= as_of_date
and (report_start is None or r.report_date >= report_start)
]
def list_announced_many(self, symbols, as_of_date):
syms = set(symbols)
return [
r for r in self._rows
if r.symbol in syms and r.announce_date <= as_of_date
]
def _svc(df, stocks=None, fin_rows=None) -> SelectionService:
return SelectionService(
_MemStockRepo(stocks or _stocks()),
_MemDailyRepo(df),
_MemFinancialRepo(fin_rows or []),
)
class _MemStockRepo:
def __init__(self, stocks) -> None:
self._stocks = stocks
def list(self):
return self._stocks
def get_by_symbol(self, symbol):
return next((s for s in self._stocks if s.symbol == symbol), None)
def _cond_query(conditions, **kw) -> SelectionQuery:
base = dict(method="condition", conditions=conditions, as_of=date(2024, 12, 31))
base.update(kw)
return SelectionQuery(**base)
class TestStaticCondition:
def test_industry_filter(self) -> None:
df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320)
industries = {_SYMS[0]: "白酒", _SYMS[1]: "银行", _SYMS[2]: "白酒"}
svc = _svc(df, stocks=_stocks(industries))
res = svc.select(
_cond_query(
[{"field": "static.industry", "op": "in", "value": ["白酒"]}],
universe=UniverseSpec(min_listing_days=0),
)
)
got = {c.symbol for c in res.candidates}
assert got == {_SYMS[0], _SYMS[2]}
assert res.statistics.selected == 2
# 每候选带条件状态与理由
assert all(c.filter_status for c in res.candidates)
assert all(c.selection_reason for c in res.candidates)
def test_static_ne(self) -> None:
df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320)
industries = {_SYMS[0]: "白酒", _SYMS[1]: "银行", _SYMS[2]: "白酒"}
svc = _svc(df, stocks=_stocks(industries))
res = svc.select(
_cond_query(
[{"field": "static.industry", "op": "ne", "value": "白酒"}],
universe=UniverseSpec(min_listing_days=0),
)
)
assert {c.symbol for c in res.candidates} == {_SYMS[1]}
class TestTechCondition:
def test_momentum_gt_zero(self) -> None:
df = synthetic_daily({_SYMS[0]: 0.008, _SYMS[1]: -0.008, _SYMS[2]: 0.002}, n=320)
svc = _svc(df)
res = svc.select(
_cond_query(
[{"field": "momentum_60", "op": "gt", "value": 0}],
universe=UniverseSpec(min_listing_days=0),
)
)
got = {c.symbol for c in res.candidates}
assert _SYMS[1] not in got # 下跌股 60 日动量为负
assert _SYMS[0] in got and _SYMS[2] in got
def test_close_above_ma60_ref(self) -> None:
df = synthetic_daily({_SYMS[0]: 0.006, _SYMS[1]: -0.006, _SYMS[2]: 0.0005}, n=320)
svc = _svc(df)
res = svc.select(
_cond_query(
[{"field": "close", "op": "gt", "ref": "ma60"}],
universe=UniverseSpec(min_listing_days=0),
)
)
got = {c.symbol for c in res.candidates}
assert _SYMS[1] not in got # 下跌股收盘在 MA60 之下
assert _SYMS[0] in got
def test_lte_threshold(self) -> None:
df = synthetic_daily({_SYMS[0]: 0.01, _SYMS[1]: -0.01, _SYMS[2]: 0.0001}, n=320)
svc = _svc(df)
res = svc.select(
_cond_query(
[{"field": "momentum_60", "op": "lte", "value": 0}],
universe=UniverseSpec(min_listing_days=0),
)
)
got = {c.symbol for c in res.candidates}
assert _SYMS[1] in got
assert _SYMS[0] not in got
class TestFundamentalCondition:
def _fin_rows(self) -> list[FinancialIndicator]:
return [
FinancialIndicator(
symbol=_SYMS[0], report_date=date(2024, 9, 30), announce_date=date(2024, 10, 25),
eps=Decimal("3.5"), roe=Decimal("20.0"), source="tushare",
),
# B:只在 as_of 之后才公告(未来数据)→ as_of 时不可见
FinancialIndicator(
symbol=_SYMS[1], report_date=date(2024, 9, 30), announce_date=date(2025, 3, 30),
roe=Decimal("99.0"), source="tushare",
),
# C:roe 低于阈值
FinancialIndicator(
symbol=_SYMS[2], report_date=date(2024, 6, 30), announce_date=date(2024, 8, 20),
roe=Decimal("5.0"), source="tushare",
),
]
def test_roe_filter_no_future_leak(self) -> None:
df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320)
svc = _svc(df, fin_rows=self._fin_rows())
res = svc.select(
_cond_query(
[{"field": "fundamental.roe", "op": "gte", "value": 15}],
universe=UniverseSpec(min_listing_days=0),
),
# 上面 helper 已带 as_of
)
got = {c.symbol for c in res.candidates}
# A 可见且 roe=20 → 入选;B 公告在未来(防未来函数)→ 不入选;C roe=5 → 不入选
assert got == {_SYMS[0]}
def test_missing_financial_repo_raises(self) -> None:
df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320)
svc = SelectionService(_MemStockRepo(_stocks()), _MemDailyRepo(df)) # 无财务 repo
with pytest.raises(ValueError):
svc.select(
_cond_query(
[{"field": "fundamental.roe", "op": "gte", "value": 15}],
universe=UniverseSpec(min_listing_days=0),
)
)
def test_as_of_earlier_excludes_announced_after(self) -> None:
"""更早 as_of:A 的 roe=20 若在 as_of 之后才公告也不可见。"""
df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320)
rows = [
FinancialIndicator(
symbol=_SYMS[0], report_date=date(2024, 6, 30),
announce_date=date(2024, 10, 1), roe=Decimal("99.0"), source="tushare",
)
]
svc = _svc(df, fin_rows=rows)
res = svc.select(
SelectionQuery(
method="condition",
conditions=[{"field": "fundamental.roe", "op": "gte", "value": 50}],
as_of=date(2024, 9, 1), # announce(10-01) 尚未来
universe=UniverseSpec(min_listing_days=0),
)
)
assert {c.symbol for c in res.candidates} == set() # 无人可见 roe → 全部不通过
+130
View File
@@ -0,0 +1,130 @@
"""M6.3 选股 API 集成测试:POST /api/selections 落库 → GET 读回 → 历史列表。
使用 tmp SQLite + 真实 SQLAlchemy Repository(override get_session),
验证「提交→落库→读回一致」闭环与 v2 §8 历史选股查询。
"""
from __future__ import annotations
from datetime import date
import pytest
from app.api import deps
from app.domain.entities.market import Stock
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
SqlAlchemyStockRepository,
)
from app.main import app
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
_SYMS = ["60000" + str(i) + ".SH" for i in range(5)] # 600000~600004
_AS_OF = date(2024, 12, 31)
@pytest.fixture()
def client(tmp_path) -> TestClient:
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
SessionFactory = sessionmaker(bind=engine, expire_on_commit=False)
drifts = {s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)}
daily_df = synthetic_daily(drifts, n=320)
with SessionFactory() as session:
stocks = [
Stock(symbol=s, name=f"测试股份{i}", industry="白酒",
list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)
]
SqlAlchemyStockRepository(session).upsert_many(stocks)
bars = bars_dataframe_to_daily_bars(daily_df)
SqlAlchemyDailyBarRepository(session).upsert_many(bars)
session.commit()
def _session_override():
with SessionFactory() as session:
yield session
app.dependency_overrides[deps.get_session] = _session_override
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
_SCORE_BODY = {
"universe": {"min_listing_days": 0},
"method": "score",
"factors": [{"name": "momentum_60", "weight": 1.0}],
"top_n": 3,
"as_of": _AS_OF.isoformat(),
}
class TestSelectionsApi:
def test_submit_then_read_back(self, client: TestClient) -> None:
resp = client.post("/api/selections", json=_SCORE_BODY)
assert resp.status_code == 200
body = resp.json()
sel_id = body["selection_id"]
assert sel_id.startswith("SEL-")
result = body["result"]
assert len(result["candidates"]) == 3
assert [c["rank"] for c in result["candidates"]] == [1, 2, 3]
assert result["candidates"][0]["factor_values"] # 有因子值
assert result["candidates"][0]["selection_reason"] # 可解释
assert result["as_of_date"] <= _AS_OF.isoformat()
# 读回一致
got = client.get(f"/api/selections/{sel_id}")
assert got.status_code == 200
g = got.json()
assert g["as_of_date"] == result["as_of_date"]
assert [c["symbol"] for c in g["candidates"]] == [
c["symbol"] for c in result["candidates"]
]
assert g["statistics"]["selected"] == 3
def test_missing_id_404(self, client: TestClient) -> None:
assert client.get("/api/selections/SEL-NOPE").status_code == 404
def test_list_and_filter(self, client: TestClient) -> None:
client.post("/api/selections", json=_SCORE_BODY)
client.post(
"/api/selections",
json={**_SCORE_BODY, "top_n": 2, "method": "score",
"factors": [{"name": "momentum_20", "weight": 1.0}]},
)
rows = client.get("/api/selections").json()
assert len(rows) >= 2
assert all(r["id"].startswith("SEL-") for r in rows)
assert all(r["method"] == "score" for r in rows)
assert all(r["selected"] > 0 for r in rows)
# as_of 过滤
by_date = client.get(f"/api/selections?as_of={_AS_OF.isoformat()}").json()
assert len(by_date) == len(rows)
empty = client.get("/api/selections?as_of=2020-01-01").json()
assert empty == []
def test_condition_submit(self, client: TestClient) -> None:
resp = client.post(
"/api/selections",
json={
"universe": {"min_listing_days": 0},
"method": "condition",
"conditions": [
{"field": "static.industry", "op": "eq", "value": "白酒"},
{"field": "momentum_60", "op": "gt", "value": 0},
],
"as_of": _AS_OF.isoformat(),
},
)
assert resp.status_code == 200
body = resp.json()
assert len(body["result"]["candidates"]) == 5 # 全部上涨 → 5 只全过
assert all(c["filter_status"] for c in body["result"]["candidates"])
+165
View File
@@ -0,0 +1,165 @@
"""M8.1 信号引擎测试:规则判定(BUY/WATCH/SELL)、可解释 reason、engine/service/API 落库回读。
使用合成行情:漂移差异决定 rank;构造「强趋势股(BUY)」与「高位破位股」验证类型。
"""
from __future__ import annotations
from datetime import date
import pandas as pd
import pytest
from app.api import deps
from app.application.services.signal_service import SignalService
from app.domain.entities.market import Stock
from app.domain.entities.research import UniverseSpec
from app.domain.entities.selection import SelectionQuery
from app.domain.entities.signal import SignalRules
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.main import app
from app.quant.signal import generate_signals
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
_SYMS = ["600000.SH", "600001.SH", "600002.SH", "600003.SH", "600004.SH"]
class _MemDailyRepo:
def __init__(self, df: pd.DataFrame) -> None:
self._bars = bars_dataframe_to_daily_bars(df)
def get_range(self, symbol, start, end):
return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end]
def get_range_many(self, symbols, start, end, adjust="none"):
syms = set(symbols)
return [b for b in self._bars if b.symbol in syms and start <= b.trade_date <= end]
class _MemStockRepo:
def __init__(self, stocks):
self._stocks = stocks
def list(self):
return self._stocks
def get_by_symbol(self, symbol):
return next((s for s in self._stocks if s.symbol == symbol), None)
def _query(as_of: date) -> SelectionQuery:
return SelectionQuery(
universe=UniverseSpec(exclude_st=False, min_listing_days=0),
factors=[{"name": "momentum_60", "weight": 1.0}],
top_n=50,
as_of=as_of,
)
class TestSignalEngine:
def test_buy_watch_sell_classification(self) -> None:
# 强势股数量足够时:顶部(高动量)BUY;构造一只「曾强现破位」→ SELL 警示
df = synthetic_daily({s: 0.008 - 0.002 * i for i, s in enumerate(_SYMS)}, n=320)
rules = SignalRules(buy_rank_threshold=2, sell_rank_threshold=4, max_output_rank=5)
res = generate_signals(df, _query(date(2024, 12, 31)), rules, date(2024, 12, 31))
assert res.statistics.buy == 2 # rank1/2 均为强趋势(上涨)→ BUY
types = {e.signal_type for e in res.events}
assert "BUY" in types and "WATCH" in types and "SELL" in types
# 每个事件可解释 + 价格
for e in res.events[:3]:
assert e.trigger_reason and (e.price or e.price == 0 or e.price is None)
bu = [e for e in res.events if e.signal_type == "BUY"]
assert all(any("排名" in r for r in e.trigger_reason) for e in bu)
def test_selected_rank_consistent_with_score(self) -> None:
df = synthetic_daily({s: 0.008 - 0.002 * i for i, s in enumerate(_SYMS)}, n=320)
res = generate_signals(df, _query(date(2024, 12, 31)), SignalRules(), date(2024, 12, 31))
events = sorted(res.events, key=lambda e: -e.score)
assert events == res.events # 按分数降序
assert res.statistics.universe_size == 5
def test_sell_on_trend_break(self) -> None:
"""下跌股在 BUY 区间内(如仅 3 只有分时 rank1 是跌股)→ WATCH/SELL 而非 BUY。"""
df = synthetic_daily({_SYMS[0]: -0.004, _SYMS[1]: 0.006, _SYMS[2]: 0.006}, n=320)
rules = SignalRules(buy_rank_threshold=1, trend_ma=60, max_output_rank=3)
res = generate_signals(df, _query(date(2024, 12, 31)), rules, date(2024, 12, 31))
# rank1 是负动量股(无正动量不构成 BUY 趋势要求?momentum_60 负但 close>ma60 可能仍成立
# 用趋势条件验证:若 rank1 close<ma60 → 不得 BUY
top = res.events[0]
if top.symbol == _SYMS[0]: # 负漂移:多半破位
assert top.signal_type != "BUY"
class TestSignalServiceAndApi:
def _svc(self, df):
stocks = [
Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1)) for i, s in enumerate(_SYMS)
]
return SignalService(_MemStockRepo(stocks), _MemDailyRepo(df))
def test_service_runs(self) -> None:
df = synthetic_daily({s: 0.008 - 0.002 * i for i, s in enumerate(_SYMS)}, n=320)
res = self._svc(df).signal(_query(date(2024, 12, 31)), SignalRules(buy_rank_threshold=2))
assert res.statistics.buy == 2
assert res.config_snapshot["as_of"] == "2024-12-31"
@pytest.fixture()
def seeded_client(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
df = synthetic_daily({s: 0.008 - 0.002 * i for i, s in enumerate(_SYMS)}, n=320)
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
SqlAlchemyStockRepository,
)
with Session() as session:
SqlAlchemyStockRepository(session).upsert_many(
[
Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)
]
)
SqlAlchemyDailyBarRepository(session).upsert_many(bars_dataframe_to_daily_bars(df))
session.commit()
def _session_override():
with Session() as s:
yield s
app.dependency_overrides[deps.get_session] = _session_override
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
class TestSignalApi:
def test_submit_readback_list(self, seeded_client) -> None:
body = {
"query": {
"universe": {"exclude_st": False, "min_listing_days": 0},
"factors": [{"name": "momentum_60", "weight": 1}],
"top_n": 50,
"as_of": "2024-12-31",
},
"rules": {"buy_rank_threshold": 2, "sell_rank_threshold": 4, "max_output_rank": 5},
}
resp = seeded_client.post("/api/signals", json=body)
assert resp.status_code == 200
sig_id = resp.json()["signal_id"]
assert sig_id.startswith("SIG-")
result = resp.json()["result"]
assert result["statistics"]["buy"] == 2
assert result["events"][0]["trigger_reason"]
got = seeded_client.get(f"/api/signals/{sig_id}").json()
assert got["as_of_date"] == "2024-12-31"
assert got["events"][0]["symbol"] == result["events"][0]["symbol"]
rows = seeded_client.get("/api/signals").json()
assert len(rows) >= 1 and rows[0]["id"] == sig_id
assert seeded_client.get("/api/signals/SIG-NOPE").status_code == 404
+123
View File
@@ -0,0 +1,123 @@
"""M8.3 策略测试:repo CRUD + /api/strategies + expand 为 ResearchSpec。"""
from __future__ import annotations
from datetime import date
import pytest
from app.api import deps
from app.domain.entities.research import SelectionSpec
from app.domain.entities.strategy import StrategyDefinition
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
SqlAlchemyStrategyRepository,
)
from app.main import app
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
def _st() -> StrategyDefinition:
return StrategyDefinition(
name="质量成长动量",
description="ROE+动量(演示)",
factors=[{"name": "momentum_60", "weight": 1.0}],
selection=SelectionSpec(top_n=10),
)
@pytest.fixture()
def session(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'st.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
with Session() as s:
yield s
class TestStrategyRepository:
def test_save_get_list_delete(self, session) -> None:
repo = SqlAlchemyStrategyRepository(session)
repo.save(_st().model_copy(update={"id": "STG-T1"}))
session.commit()
got = repo.get("STG-T1")
assert got is not None and got.name == "质量成长动量"
assert got.selection.top_n == 10
assert len(repo.list()) == 1
assert repo.get_by_name("质量成长动量") is not None
assert repo.delete("STG-T1") is True
session.commit()
assert repo.get("STG-T1") is None
def test_duplicate_name(self, session) -> None:
repo = SqlAlchemyStrategyRepository(session)
repo.save(_st().model_copy(update={"id": "STG-A"}))
session.commit()
with pytest.raises(ValueError):
repo.save(_st().model_copy(update={"id": "STG-B"}))
def test_expand_to_research_spec(self, session) -> None:
st = _st().model_copy(update={"id": "STG-E"})
spec = st.to_research_spec((date(2024, 1, 1), date(2024, 6, 1)))
assert spec.type == "backtest"
assert spec.period == (date(2024, 1, 1), date(2024, 6, 1))
assert spec.factors[0].name == "momentum_60"
assert spec.price_adjustment == "none"
@pytest.fixture()
def client(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
def _session_override():
with Session() as s:
yield s
app.dependency_overrides[deps.get_session] = _session_override
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
class TestStrategiesApi:
def test_crud_and_expand(self, client) -> None:
body = {
"name": "演示策略",
"description": "动量",
"factors": [{"name": "momentum_60", "weight": 1}],
"selection": {"top_n": 10},
}
created = client.post("/api/strategies", json=body)
assert created.status_code == 200
sid = created.json()["id"]
assert sid.startswith("STG-")
assert len(client.get("/api/strategies").json()) == 1
detail = client.get(f"/api/strategies/{sid}").json()
assert detail["name"] == "演示策略"
resp = client.post(
f"/api/strategies/{sid}/expand",
json={"period": ["2024-01-01", "2024-06-01"]},
)
assert resp.status_code == 200
spec = resp.json()
assert spec["type"] == "backtest"
assert spec["factors"][0]["name"] == "momentum_60"
assert client.delete(f"/api/strategies/{sid}").status_code == 200
assert client.get(f"/api/strategies/{sid}").status_code == 404
def test_duplicate_and_bad_period(self, client) -> None:
body = {"name": "A", "factors": [{"name": "momentum_60", "weight": 1}]}
assert client.post("/api/strategies", json=body).status_code == 200
assert client.post("/api/strategies", json=body).status_code == 400
sid = client.get("/api/strategies").json()[0]["id"]
bad = client.post(
f"/api/strategies/{sid}/expand",
json={"period": ["2024-06-01", "2024-01-01"]},
)
assert bad.status_code == 400
+70
View File
@@ -0,0 +1,70 @@
"""M6.1 UniverseFilter 测试:as_of 当前/历史日语义、ST、上市天数、退市、symbols 白名单。
filter_stocks 从 quant.universe 引入(原 quant.service 语义,规则化集中)。
"""
from __future__ import annotations
from datetime import date
from app.domain.entities.market import Stock
from app.domain.entities.research import UniverseSpec
from app.quant.universe import filter_stocks
def _stocks() -> list[Stock]:
return [
Stock(symbol="600000.SH", name="正常股份", list_date=date(2000, 1, 1)),
Stock(symbol="600001.SH", name="ST 风险股份", list_date=date(2000, 1, 1)),
Stock(symbol="600002.SH", name="次新股", list_date=date(2024, 10, 1)),
Stock(symbol="600003.SH", name="已退市股", list_date=date(1995, 1, 1),
delist_date=date(2023, 6, 30)),
Stock(symbol="600004.SH", name="老股", list_date=date(1999, 1, 1)),
]
def _sym(rows: list[Stock]) -> set[str]:
return {s.symbol for s in rows}
class TestUniverseFilter:
def test_current_day(self) -> None:
rows = filter_stocks(_stocks(), UniverseSpec(), as_of=date(2025, 1, 1))
# ST、退市被剔除;次新股(上市<250 自然日)被剔除
assert _sym(rows) == {"600000.SH", "600004.SH"}
def test_historical_as_of_keeps_not_yet_delisted(self) -> None:
rows = filter_stocks(_stocks(), UniverseSpec(), as_of=date(2023, 1, 1))
# 2023-01 时 600003 尚未退市(2023-06 退市)→ 应纳入
assert "600003.SH" in _sym(rows)
assert "600002.SH" not in _sym(rows) # 2024-10 才上市,2023-01 尚不存在(上市天数不足被滤)
def test_delisted_before_as_of_excluded(self) -> None:
rows = filter_stocks(_stocks(), UniverseSpec(exclude_st=False), as_of=date(2024, 1, 1))
assert "600003.SH" not in _sym(rows) # 2023-06 已退市
def test_exclude_st_flag(self) -> None:
rows = filter_stocks(_stocks(), UniverseSpec(exclude_st=False), as_of=date(2025, 1, 1))
assert "600001.SH" in _sym(rows)
rows2 = filter_stocks(_stocks(), UniverseSpec(exclude_st=True), as_of=date(2025, 1, 1))
assert "600001.SH" not in _sym(rows2)
def test_min_listing_days_zero_disables(self) -> None:
rows = filter_stocks(
_stocks(), UniverseSpec(min_listing_days=0, exclude_st=True), as_of=date(2025, 1, 1)
)
assert "600002.SH" in _sym(rows)
def test_symbols_whitelist(self) -> None:
rows = filter_stocks(
_stocks(),
UniverseSpec(symbols=["600000.SH", "600003.SH"]),
as_of=date(2025, 1, 1),
)
# 白名单内的 ST/退市过滤仍然生效:600003 已退市被滤,仅剩 600000
assert _sym(rows) == {"600000.SH"}
def test_empty_whitelist_means_all(self) -> None:
assert UniverseSpec().symbols == []
rows = filter_stocks(_stocks(), UniverseSpec(symbols=[]), as_of=date(2025, 1, 1))
assert "600004.SH" in _sym(rows)
+11
View File
@@ -3262,6 +3262,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/64/02/b2606ddf52fa615d4ff3edb2aa05fb064337fa871f99ebb184a39e051cc8/pymongo-4.18.0-cp314-cp314t-win_arm64.whl", hash = "sha256:4a81d166a43e8af1e5152b6854a263ba0a8831f7dd2ca1badc716f219f4f1bc0", size = 817605, upload-time = "2026-09-03T16:00:41.547Z" }, { url = "https://files.pythonhosted.org/packages/64/02/b2606ddf52fa615d4ff3edb2aa05fb064337fa871f99ebb184a39e051cc8/pymongo-4.18.0-cp314-cp314t-win_arm64.whl", hash = "sha256:4a81d166a43e8af1e5152b6854a263ba0a8831f7dd2ca1badc716f219f4f1bc0", size = 817605, upload-time = "2026-09-03T16:00:41.547Z" },
] ]
[[package]]
name = "pymysql"
version = "1.2.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/c9/bc/1c6a92f385940f727daeecf3bacaf186e03875dff57197801046c583bcf0/pymysql-1.2.0.tar.gz", hash = "sha256:6c7b17ca686988104d7426c27895b455cdeea3e9d3ceb1270f0c3704fead8c33", size = 49021, upload-time = "2026-05-19T08:26:22.302Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/c4/bd/2534e130295c8cfd4f0a2e31623baab7502278f1e97bcfe61db75656a77f/pymysql-1.2.0-py3-none-any.whl", hash = "sha256:62169ce6d5510f08e140c5e7990ee884a9764024e4a9a27b2cc11f1099322ae0", size = 45716, upload-time = "2026-05-19T08:26:20.974Z" },
]
[[package]] [[package]]
name = "pyparsing" name = "pyparsing"
version = "3.3.2" version = "3.3.2"
@@ -3540,6 +3549,7 @@ dependencies = [
{ name = "alembic" }, { name = "alembic" },
{ name = "fastapi" }, { name = "fastapi" },
{ name = "httpx" }, { name = "httpx" },
{ name = "pymysql" },
{ name = "pyqlib" }, { name = "pyqlib" },
{ name = "pyyaml" }, { name = "pyyaml" },
{ name = "sqlalchemy" }, { name = "sqlalchemy" },
@@ -3563,6 +3573,7 @@ requires-dist = [
{ name = "alembic", specifier = ">=1.13" }, { name = "alembic", specifier = ">=1.13" },
{ name = "fastapi", specifier = ">=0.115" }, { name = "fastapi", specifier = ">=0.115" },
{ name = "httpx", specifier = ">=0.28.1" }, { name = "httpx", specifier = ">=0.28.1" },
{ name = "pymysql", specifier = ">=1.1" },
{ name = "pyqlib", git = "https://github.com/microsoft/qlib.git?rev=79633dd" }, { name = "pyqlib", git = "https://github.com/microsoft/qlib.git?rev=79633dd" },
{ name = "pyyaml", specifier = ">=6.0" }, { name = "pyyaml", specifier = ">=6.0" },
{ name = "sqlalchemy", specifier = ">=2.0" }, { name = "sqlalchemy", specifier = ">=2.0" },
+15 -1
View File
@@ -17,11 +17,25 @@ api:
prefix: "/api" prefix: "/api"
database: database:
# 留空 → 默认 sqlite:///./data/quant.db(相对路径自动解析到项目根 data/) # 数据库 URL 选择优先级:
# 1) 环境变量 database.url_env(如 DATABASE_URL,可在 .env 覆盖)
# 2) database.mysql(enabled: true 时由 host/port/db/user 组装 mysql+pymysql URL)
# 3) 兜底 sqlite:///./data/quant.db(相对路径自动解析到项目根 data/)
url_env: "DATABASE_URL" url_env: "DATABASE_URL"
echo: false echo: false
# Alembic 迁移脚本目录(相对 backend/) # Alembic 迁移脚本目录(相对 backend/)
migrations_dir: "app/infrastructure/persistence/migrations" migrations_dir: "app/infrastructure/persistence/migrations"
# ---- MySQL(当前默认库;2026-09 由 SQLite data/quant.db 全量迁移而来)----
# host/port/db/user 明文可提交;密码经 password_env 引用 .env 变量
# (AGENT.md §33:密钥只放根目录 .env,本文件不写明文密码)。
mysql:
enabled: true
host: "192.168.1.10"
port: 3306
db: "qlib"
user: "qlib"
password_env: "MYSQL_PASSWORD"
charset: "utf8mb4"
data_source: data_source:
primary: "tushare" primary: "tushare"
+255
View File
@@ -0,0 +1,255 @@
# 下阶段开发计划(架构 v2 落地 · 2026-09 · 选股系统为主线)
> 依据:[ARCHITECTURE_v2.md](./ARCHITECTURE_v2.md)(新架构 v2)、[AGENT.md](../AGENT.md)(开发约束)、
> 两份 2026-09 只读调研(后端研究/量化层、前端/Agent/数据层)与 docs/ROADMAP.md 的 M0–M5 记录。
> 本文件回答:**M5(AI Agent Phase 5)之后,下一阶段做什么、按什么顺序做。**
> 2026-09 更新(用户定调):**下一阶段以「选股系统」为独立主线并提前执行**
> (原"研究层落地"后移为 M7 支撑),先交付可用的选股 MVP,再补因子层增强。
---
## 0. 当前状态与总判断
**已完成(M0–M5)**:数据层(Tushare 首选 + 新浪兜底 + 审计)→ ResearchSpec → 本地研究引擎
(9 个行情因子 / IC·分层评估 / 低频 TopN 等权回测)→ Job/Experiment 异步归档 → Web 六页 →
AI Agent(6 个受控 Tool)。**数据库已于 2026-09 从 SQLite 全量迁移至 MySQL**(见 §1)。
**对照 v2 的核心差距**(调研结论):
1. **选股没有独立成系统**:v2 强调 Factor Research / Composite Factor / **Selection** / Signal /
Portfolio 分层;当前"选股"只是回测器内部 TopN 逻辑,无法单独回答 v2 §8 的两个核心问题:
「2023-08-15 为什么选这只股票」「2026-09-08 当前有哪些股票满足策略」。
2. **因子定义 / 策略等研究元数据未入库**:DB 只有 8 张表(stock / stock_daily / adjust_factor /
financial_indicator / trading_calendar / sync_log / job / experiment);v2 §8 期望的
factor_definition / universe / selection_rule / selection_result / signal_rule / strategy 等全缺。
3. **口径与数据完整性风险待修**:研究链路不消费复权因子(stock_daily 混合 tushare 不复权与新浪前复权行);
universe 用"当前名称含 ST"近似、无历史成分/停牌;Job.stage 从未赋值、SSE 是查库轮询。
4. **前端/Agent 待做实**:Selection/Signal 页面不存在;Agent 工具仅 6 个,
v2 §24 期望的 screen_stocks / explain_selection 等依赖尚未存在的选股引擎。
**阶段命名**:**M6 选股系统(主线)→ M7 因子层落地(支撑)→ M8 引擎延伸与平台化**。
数据库专项单独记 M-DB。
---
## 1. 执行状态(2026-09:M6–M8 主干已交付)
> M6 选股系统、M7 因子层、M8(Signal/Portfolio/Strategy/Agent 工具/Web)均已完成并提交
> (ROADMAP 里程碑行含 commit;每阶段全量 pytest + ruff + tsc 验证通过;Alembic 迁移已应用 MySQL)。
> M8.4 Qlib 模型选股(Selection 模式 C)为「按需」项,未阻塞主线 —— 当前延后,
> 前置:qlib_adapter 特征/模型链路(Alpha158 + LightGBM walk-forward),启用时按 §5.4 推进。
## 2. M-DB:数据库 MySQL 迁移(已完成,本文件记录收尾项)
**已交付(2026-09)**:
- MySQL(192.168.1.10:3306,MariaDB 10.11,库 `qlib`,utf8mb4)建 8 张业务表 + alembic_version,
全部经 `alembic upgrade head`(迁移 e4d18825…53113c80…91c4e27a…d3f6c9a2),结构由既有 ORM 模型保证。
- `scripts/migrate_sqlite_to_mysql.py`:一次性工具 —— sqlite 直连 + pymysql chunk 多值 INSERT +
keyset 分页续传 + 幂等(已迁则跳过)+ `--verify-only` 逐表 COUNT 与抽样比对;支持 `--only <表>` 并行大表、
`--throttle-sec` 低配服务器节流。
- 配置:`config.yaml → database.mysql`(host/port/db/user/charset 明文),密码经 `.env` 的
`MYSQL_PASSWORD`(AGENT.md §33:密钥不进 config/git);URL 优先级
`DATABASE_URL 环境变量 > config.yaml database.mysql > sqlite:///./data/quant.db 兜底`;
实现集中在 `backend/app/core/config.py::_build_mysql_url`,业务层零改动(DAO/Repository 隔离兑现)。
- 测试隔离:`backend/tests/conftest.py` 强制每进程 `/tmp/qlib-pytest-<pid>.db` SQLite(测试绝不触 MySQL)。
- 驱动:`pymysql>=1.1` 已入 `backend/pyproject.toml`。迁移数据量 16,057,265 行,`--verify-only` 校验全部一致。
**遗留收尾(小任务,可随时做)**:
| 项 | 说明 | 建议 |
|---|---|---|
| `job/experiment.result_json` 用 MySQL `TEXT`(≤65,535B) | SQLite TEXT 无上限、测试测不出 | 日级回测结果变长后改 `MEDIUMTEXT`(新 Alembic 迁移 + Model 类型 variant) |
| Repository 批量 upsert 在 MySQL 的压测 | `tuple_.in_()` 行值、`add_all` 单事务无分块 | 全市场日线增量同步跑一次,观察锁等待/包大小 |
| MySQL 路径无自动化测试 | 测试全 SQLite 方言 | 视需要加一个「连 MySQL 的集成冒烟」开关(默认关,显式 env 开启) |
| `scripts/migrate_sqlite_to_mysql.py` 属一次性工具 | 保留供重建目标库 | 文档标注不可再对生产执行(除 --verify-only) |
| 前端/文档残留 "SQLite" 文案 | app-shell / page.tsx / sync.py docstring | 随各页面改动顺手清理 |
---
## 3. 下阶段总路线
```text
M6 选股系统(主线) Universe 选股范围 + Selection Engine(A条件/B评分)
│ + 结果落库/API + 回测共用引擎 + Web 选股页
↓
M7 因子层落地 因子定义入库 + Composite 模块化 + 研究口径修复(支撑选股增强)
↓
M8 引擎延伸与平台化 Signal Engine → Portfolio Engine → Strategy / Agent / Qlib 模型选股
```
设计红线(贯穿):v2 §25「历史回测与当前选股用同一套 Engine」、v2 §9 时间与未来函数防护、
v2 §26 Experiment 可复现、AGENT.md「简单优先 / 每阶段可运行 / 一个功能一个 commit /
先 Domain 后 UI / 带测试」。
---
## 4. M6:选股系统(下一阶段主线)
> 目标:交付一个**独立、可解释、支持当前与历史日期**的选股系统(v2 §14 Selection Engine),
> 让回测器与"现在该选什么"共用同一套引擎(v2 §25 一致性红线)。
> MVP 全部复用**现有数据与因子能力**(stock / financial_indicator / stock_daily + 9 个内置行情因子 +
> composite_score),不依赖 M7 因子层,因此可最先做、尽快见到"选股结果"。
### 3.0 选股输入输出契约(先定 Domain / DTO)
- Domain `SelectionSpec`(v2 §14.2):`universe 范围 + filters(条件) + scoring(因子/复合因子/权重) +
ranking(TopN / Top%) + as_of_date`;沿用并扩展现有 `ResearchSpec` 的 universe/factors/selection 结构,
新增 `as_of_date`(缺省 = 最近交易日)。
- Domain `SelectionResult`(v2 §14.3/§21.1):`as_of_date, universe, candidates[{symbol, rank, score,
factor_values, filter_status, selection_reason}], statistics` —— 作为前后端/Agent 统一 DTO。
- 核心用例:`selection_service.select(spec, as_of=...) -> SelectionResult`(历史/当前日期都可执行)。
- 验收:SelectionResult 进入 `frontend/web/lib/types.ts` 与后端 entity 一一对应。
### 3.1 Universe:选股范围(MVP 规则化)
- 把 `ResearchService.filter_stocks` 的过滤逻辑(market / exclude_st / min_listing_days /
delist_date < as_of / 可选 symbol 白名单)抽成可复用的 `UniverseFilter` 执行器,
输入结构化条件(Pydantic),输出 as_of 时点下的股票范围。
- MVP 不建 universe 大表:条件随 `SelectionSpec` 一起存入 `selection_snapshot`(快照 JSON,
可复现);`universe / universe_rule / index_constituent_history` 表推迟到 M8 需要历史成分时再落。
- 明确近似并写进结果:ST 按当前名称快照判定;停牌/退市边界显式标注(unimplemented 或近似说明)。
- 验收:同一组条件在"当前日"与"历史日"分别得到正确范围(含历史日上市/退市过滤);
现有回测 universe 语义不变(回归)。
### 3.2 Selection Engine:A 条件选股 + B 因子评分
- **A. 条件选股**:结构化条件过滤(示例:`roe > 15 且 pe… 且 close > ma60`),
条件作用于现有字段/派生指标:
- 静态/财务:`market / industry / eps / roe / total_revenue / net_profit / gross_margin…`
(财务值按 `announce_date ≤ as_of` 可见,防未来函数);
- 行情/技术:`close > ma20 / ma60、momentum_20>0、volume_ratio…`(复用 factors.py 的 `compute_factor`)。
- MVP 条件用**结构化 JSON**(字段 + 比较符 + 值,可 and/or 嵌套),不做自由表达式 parser(避免提前造轮子)。
- **B. 因子评分选股**:`score = composite_score(截面 zscore × 方向 × 权重)`(抽取 local_engine 现逻辑),
按 score 排序取 `top_n` / `top_pct`;同时支持阈值过滤。
- **C. 模型预测选股**(Qlib/LightGBM)留 M8.3,引擎接口先留 mode 字段占位。
- 每条候选输出 `selection_reason`(命中的条件 / 分数来源),支撑 v2「为什么选这只股票」。
- 验收:同 spec 下 A/B 模式输出与人工核对一致;条件/评分模式均可对历史 as_of 执行。
### 3.3 结果持久化与 API
- Model(→ Alembic → Repo,全部走 Model→Migration→Test):
- `selection_result`(v2 §8:selection_id / strategy_id 占位 / as_of_date / symbol / rank / score /
factor_values / filter_status / selection_reason / created_at);
- `selection_snapshot`(spec 与 universe/条件快照 JSON,供复现与"历史某日选了什么"查询)。
- `selection_rule`(可选的命名规则持久化,MVP 允许后置:snapshot 已含完整 spec)。
- API(业务对象导向,AGENT.md §17):`POST /api/selections`(提交 spec+as_of,同步或 Job 计算并落库)、
`GET /api/selections/{id}`、`GET /api/selections?as_of=&strategy_id=`(历史选股查询)。
- 验收:提交一次选股 → 结果落库可查;重建同 spec+as_of 得到相同结果(可复现)。
### 3.4 回测一致性改造(v2 §25 红线)
- `TopKBacktestRunner._rebalance` 改为**调用 Selection Engine**(同 spec、同 as_of)得到候选
再进入成交/仓位逻辑——保证"当前选股 = 历史回测每期选股"用同一套代码。
- 默认配置下回测数值与改造前**逐项一致**(回归测试锁住)。
- 验收(v2 §27):`select(strategy, 历史调仓日)` 与回测当日持仓选股逐 symbol 一致。
### 3.5 Web 选股页(可后端就绪后并行)
- 新路由 `/selection`(v2 §19/§21.1):选股条件面板(范围 + A 条件 + B 评分/权重/TopN)+
as_of 日期(当前/历史)→ SelectionResult 表格(rank/score/因子值/理由)→「一键以该 spec 发起回测」。
- `lib/types.ts` 增加 SelectionSpec / SelectionResult;nav(app-shell)加入"股票筛选"入口。
- 验收:页面完成"选范围 → 配条件/评分 → 当前/历史 TopN 结果 → 跳回测"闭环;tsc --noEmit 通过。
### 3.6 M6 验收(里程碑完成判定)
- [ ] `select(spec, as_of=历史日)` 与回测该日调仓选股逐 symbol 一致(自动化测试)
- [ ] A/B 两模式都可对"当前日 / 历史日"执行并落库,`/api/selections` 可查
- [ ] 选股结果带 factor_values 与 selection_reason(可解释)
- [ ] Web 选股页闭环可用;SelectionResult DTO 前后端一致
- [ ] 全量 pytest 通过;相关迁移测试覆盖新表
---
## 5. M7:因子层落地(支撑选股增强,紧随 M6)
> 目标:把"因子"从代码注册表变成**可入库、可版本化、可组合**的研究资产,并修掉混合口径风险,
> 为选股 B 模式(更多/自定义因子)、因子研究做实提供基础。
- **4.1 因子定义入库 + 版本化目录**:Model `factor_definition` + seed 现有 9 因子元数据;
`GET /api/factors` 改读库;自定义因子元数据 CRUD(计算仍走代码注册表)。
- **4.2 Composite Factor Engine 模块化**:抽出 `quant/composite.py`
(`CompositeSpec{components[name,weight,direction], method}` → score 面板);
`factor_composite(+component)` 表;选股 B 模式与回测共用同一实现。验收:默认 spec 回测结果与现一致。
- **4.3 研究读取口径修复(横切)**:Repository 读路径按 `source/adjust` 过滤或提供 qfq 查询;
ResearchSpec/SelectionSpec 增加 `price_adjustment` 字段并写入结果 config_snapshot。
验收:含新浪补缺行的股票回测/选股口径一致且可溯源。
- **4.4 因子测试结果结构化(可选)**:Model `factor_test`,多因子 spec 不再静默只测第一个。
---
## 6. M8:引擎延伸与平台化(Signal / Portfolio / Strategy / Agent / 模型选股)
- **5.1 Signal Engine(规则型)**:`signal_rule / signal_event{signal_date, symbol, signal_type,
trigger_reason, score, price}`;回测与选股都能产事件并落库;API `POST /api/signals`、历史查询。
验收:能回答「某策略某日为什么对这只股票给 BUY/SELL」。
- **5.2 Portfolio Engine**:把回测器 `_rebalance` 资金/仓位逻辑抽为 portfolio 模块
(等权起步 + 单股/行业上限/现金可选约束,约束默认关闭并如实标注 unimplemented)。
验收:默认配置回测数值与现一致。
- **5.3 Strategy 模型与 /api/strategies**:聚合 universe + selection(spec) + signal + portfolio +
rebalance;Experiment 归档 strategy_version;Agent 增加 `create_strategy`。
- **5.4 Qlib 模型链路 = 选股模式 C**(按需):Alpha158 特征 + LightGBM walk-forward →
预测分 → Selection 模式 C。
- **5.5 AI Agent 工具补齐**(依赖上述引擎):`inspect_factor / create_composite_factor /
screen_stocks / explain_selection / generate_signals / create_strategy / get_backtest_result /
create_experiment`,全部经白名单 + Job/Experiment 链路;前端加最小 AI 对话入口。
- **5.6 Web 页面做实**(前端主线,随后端里程碑推进):回测页做实(trades/年度收益/持仓演化/成本参数/
config_snapshot,零新 API)、因子研究做实(IC 时序/对比)、Signal 页、Job SSE 进度 UI
(需先补 executor stage 上报,见 §6)。
---
## 7. 基础设施专项(含本机 Redis 的定位)
**本机已探测:Redis 服务在 127.0.0.1:6379 正常响应(无 redis-cli,连接经 socket 验证)。**
v2 §22 与 AGENT.md §19/§38 的原则是"第一版不引复杂队列;复杂后再引入 Redis + Celery/RQ"。
据此把 Redis 放入**明确的触发点**而非一开始就用:
| 触发场景 | Redis 用途 | 何时启用 |
|---|---|---|
| Job 并发/排队:subprocess 双并发上限打满、丢任务需可恢复队列 | RQ/普通队列替代 FastAPI BackgroundTasks 调度 | M8 之后 / 出现排队丢失问题 |
| Job 细粒度 stage + SSE 长连:executor 每阶段写 stage 后广播 | Redis pub/sub 做 SSE 事件总线(替代 0.4s 查库轮询) | M8.6 前端切 SSE 时(可先行) |
| 行情/因子查询热点:全市场选股重复预热 | 因子结果 / 最新交易日缓存(key 带 as_of 失效) | M6 选股高频查询后按 profile 决定 |
近期只需做一件准备工作:`config.yaml` 预留 `redis.url_env: "REDIS_URL"` + 可选
`redis.url: "redis://127.0.0.1:6379/0"` 配置段,代码不改;真正启用时再落基础设施代码。
**其余基础设施待办**(低优先):Job `cancelled` 取消 API 与 `/api/jobs` 列表端点;executor 增加 stage
上报(data_loading/factor_calculation/backtesting…)——这是前端做 SSE 进度条的前提。
---
## 8. 建议的执行顺序(前 6 个 commit 序列)
按 AGENT.md「最小修改、每阶段可跑、一功能一 commit」,以选股 MVP 为起点的推荐顺序:
1. **M6.0 选股契约**:Domain `SelectionSpec / SelectionResult` + `selection_service.select(spec, as_of)`
骨架(先同步执行、复用 ResearchService 数据装配)。验收:对最近交易日返回 TopN 候选 + 理由。
2. **M6.1 Universe 过滤执行器**:抽 `UniverseFilter`(复用 filter_stocks 语义)。
验收:当前/历史日范围正确、回测回归不变。
3. **M6.2 Selection Engine A + B**:条件选股(结构化条件 JSON)+ 因子评分 TopN(复用因子与
composite_score)。验收:A/B 模式可执行、结果含 factor_values/reason。
4. **M6.3 落库与 API**:`selection_result / selection_snapshot` 表 + Repo + `POST/GET /api/selections`。
验收:迁移测试 + 提交可查、重建一致。
5. **M6.4 回测共用 Selection**:`TopKBacktestRunner` 改调 Selection Engine。
验收:默认 spec 回测结果逐项一致 + 历史日一致性测试。
6. **M6.5 Web 选股页**:`/selection` 路由 + SelectionResult 表 + 一键跳回测。
之后进入 M7(因子层落地)与 M8(Signal/Portfolio/Strategy/Agent/模型选股)与 §6 基础设施专项。
每步都遵循:Domain → Repository Protocol → Infra(Model→Migration→Repo) → Service → API → 前端 →
测试(相关模块)→ 更新文档;新增表一律 Model→Alembic→测试。
---
## 9. 风险清单(开发中持续检查)
- **历史一致性**:回测改用 Selection Engine 后默认结果必须与改造前一致(回归测试锁住,别让重构悄悄改策略语义)。
- **未来函数 / 幸存者偏差**:财务条件按 `announce_date ≤ as_of` 过滤;universe 的 ST/行业/成分历史
口径补齐前,结果须明示近似(ST 当前名称快照等)。
- **数据口径**:选股/回测所用行情口径(不复权 vs 前复权)必须一致并可溯源(M7.3 收口)。
- **范围控制**:M6 阶段不做自由表达式 parser / 不做模型选股(C 模式)/ 不建 universe 大表;
均留到后续里程碑,避免一次性跨层堆量。
- **MySQL 行为差异**:TEXT 上限、批量 upsert 锁等待、pymysql 流式读取 —— 见 §1 遗留收尾表。
- **前端契约**:SelectionResult/SignalResult 先定 DTO(后端 entity)再画页面,避免前后端口径漂移。
+9 -2
View File
@@ -2,7 +2,10 @@
> 依据:[AGENT.md](./AGENT.md)(开发约束)、[docs/ARCHITECTURE.md](./ARCHITECTURE.md)(架构) > 依据:[AGENT.md](./AGENT.md)(开发约束)、[docs/ARCHITECTURE.md](./ARCHITECTURE.md)(架构)
> 引擎:Qlib 只做计算引擎 · 数据源:Tushare 第一、Sina 备用 · 数据库:SQLite 起步、SQLAlchemy 保证 MySQL 可切换 > 引擎:Qlib 只做计算引擎 · 数据源:Tushare 第一、Sina 备用 · 数据库:SQLite 起步、SQLAlchemy 保证 MySQL 可切换
> 现处:**M0 已完成(工程骨架 + 可运行后端)**,下一步进入 **Phase 1 数据层** > 现处:**M5 已完成(AI Agent Phase 5)**;**M6 起的下一阶段见 [docs/DEV_PLAN_v2.md](./DEV_PLAN_v2.md)**
> (选股系统为主线:Universe 选股范围 + Selection Engine A条件/B评分 + 结果落库/API +
> 回测共用引擎 + Web 选股页 → M7 因子层 → M8 Signal/Portfolio/Strategy/Agent)
> 2026-09 增补:数据库已切 MySQL(见 M-DB 行),本文件 M0–M5 为历史记录,新计划以 DEV_PLAN_v2 为准。
--- ---
@@ -15,7 +18,11 @@
| M2 ✅ | Phase 2 研究引擎(Spec→因子→评估→低频回测→标准结果;Qlib 桥接占位见 §2 备注) | commit e9f59d3 · 60 tests | | M2 ✅ | Phase 2 研究引擎(Spec→因子→评估→低频回测→标准结果;Qlib 桥接占位见 §2 备注) | commit e9f59d3 · 60 tests |
| M3 ✅ | Phase 3 Web(业务 API + Next.js 前端闭环) | commit 92627f5 · 68 tests | | M3 ✅ | Phase 3 Web(业务 API + Next.js 前端闭环) | commit 92627f5 · 68 tests |
| M4 ✅ | Phase 4 Experiment 归档 + Job/SSE 异步(一键复跑) | commit 0ea229d · 73 tests | | M4 ✅ | Phase 4 Experiment 归档 + Job/SSE 异步(一键复跑) | commit 0ea229d · 73 tests |
| M5 ✅ | Phase 5 AI Agent(受控 Tool 白名单 + LLM 编排;Key 配置见 .env) | commit(本轮)· 79 tests | | M5 ✅ | Phase 5 AI Agent(受控 Tool 白名单 + LLM 编排;Key 配置见 .env) | commit(M5)· 79 tests |
| M-DB ✅ | 数据库 SQLite → MySQL(config.yaml database.mysql 配置化 + 迁移脚本 + 一致性校验) | 2026-09 · 见 DEV_PLAN_v2 §1 |
| M6 ✅ | 选股系统主线(Universe + Selection A条件/B评分 + 落库/API + 回测共用 + Web 页) | 2026-09 · f3586ad…0ab9038 |
| M7 ✅ | 因子层(factor_definition 入库 + Composite 模块化/落库 + 口径显式化) | 2026-09 · 8f47b5b…ef09d5b |
| M8 ✅ | Signal/Portfolio/Strategy + Agent 工具(10) + Web 做实 | 2026-09 · ba52edc…e5a23f1(M8.4 模型选股:按需延后) |
每阶段结束时同步更新:README / AGENT.md 相关清单 / 文档;**禁止跨阶段提前堆量**(AGENT.md §38)。 每阶段结束时同步更新:README / AGENT.md 相关清单 / 文档;**禁止跨阶段提前堆量**(AGENT.md §38)。
+74 -32
View File
@@ -1,8 +1,8 @@
# qlib-platform 使用说明 # qlib-platform 使用说明
> 适用代码版本:M1–M5 全部完成,含 QlibEngine v1(Qlib 数据管线)与 Parquet 导出 > 适用代码版本:M0–M5 + M-DB + **M6 选股系统 / M7 因子层 / M8(Signal·Portfolio·Strategy·Agent·Web)**
> (HEAD ≈ `fc12ba8`);若文档与代码不一致,以代码与 > 全部完成(数据库已切换 MySQL);若文档与代码不一致,以代码与
> [ARCHITECTURE.md](./ARCHITECTURE.md) / [ROADMAP.md](./ROADMAP.md) 为准。 > [ROADMAP.md](./ROADMAP.md) / [DEV_PLAN_v2.md](./DEV_PLAN_v2.md) 为准。
> >
> 本文覆盖:安装配置、数据同步、启动前后端、研究 API、研究引擎与 Qlib 接入、 > 本文覆盖:安装配置、数据同步、启动前后端、研究 API、研究引擎与 Qlib 接入、
> AI Agent、测试门禁、已知限制。 > AI Agent、测试门禁、已知限制。
@@ -15,14 +15,17 @@
| 模块 | 位置 / 说明 | | 模块 | 位置 / 说明 |
|---|---| |---|---|
| 数据层 | `backend/app/domain`、`infrastructure/data_sources`:Tushare 首选 + 新浪备用(Failover 审计);SQLite 存储 + 日线 Parquet 导出(未来 MySQL 可切换) | | 数据层 | `backend/app/domain`、`infrastructure/data_sources`:Tushare 首选 + 新浪备用(Failover 审计);**默认 MySQL**(`config.yaml database.mysql`)+ 日线 Parquet 导出 |
| 研究引擎 | `backend/app/quant`:ResearchSpec → 因子 → IC/RankIC/分层 → TopK 低频回测 → 标准化 `BacktestResult`;**双引擎**:LocalEngine(默认,纯 pandas)与 QlibEngine v1(QLibDataset 数据管线) | | 选股系统 | `quant/selection.py` + `application/services/selection_service.py`:条件选股(A)/ 因子评分 TopN(B),`as_of` 当前/历史一致;结果落库可复现、可解释 |
| Qlib 接入 | `quant/qlib_adapter/`:SQLite 行情按官方二进制格式落盘 → `D.features` 读取 → 回测(详见 §6.1 与 docs/QLIB_VERIFICATION.md) | | 因子层 | `factor_definition` 入库(`/api/factors` 读库)+ Composite Engine(`quant/composite.py`,组合落库 `/api/composites`)+ 行情口径显式化(`price_adjustment`) |
| Web | `frontend/web`:Next.js(TS)+ECharts —— 总览 / 股票池 / 因子研究 / 回测 / 实验 | | 交易信号 | `quant/signal.py`:评分排名 + 趋势规则 → BUY/WATCH/SELL + 理由,落库 `/api/signals` |
| 研究/回测 | `backend/app/quant`:ResearchSpec → 因子 → IC/RankIC/分层 → TopK 低频回测 → 标准化 `BacktestResult`;**双引擎**:LocalEngine(默认)与 QlibEngine v1 |
| 策略 | `strategy` 落库 + `/api/strategies`(命名策略,可展开为回测 spec) |
| Web | `frontend/web`:总览 / 股票池 / **股票筛选** / 因子研究 / 因子组合 / **交易信号** / 选股回测 / 实验 |
| 异步与归档 | Job 状态机 + Experiment 自动归档 + 一键复跑 + SSE 进度 | | 异步与归档 | Job 状态机 + Experiment 自动归档 + 一键复跑 + SSE 进度 |
| AI Agent | 受控工具白名单 + LLM 编排(需配置模型 Key) | | AI Agent | 10 个受控工具 + LLM 编排(需配置模型 Key) |
**技术栈**:Python 3.12(uv) · FastAPI · SQLAlchemy 2 · Alembic · SQLite · pandas · pyarrow(Parquet) · pyqlib 0.9.8.dev32(源码安装) · LightGBM · Next.js 15 · ECharts **技术栈**:Python 3.12(uv) · FastAPI · SQLAlchemy 2 · Alembic · **MySQL(pymysql)**(SQLite 兜底)· pandas · pyarrow(Parquet) · pyqlib 0.9.8.dev32(源码安装) · LightGBM · Next.js 15 · ECharts · Redis(本机 127.0.0.1:6379 已就绪,触发时接入)
--- ---
@@ -64,7 +67,8 @@ cp .env.example .env
|---|---|---| |---|---|---|
| `TUSHARE_TOKEN` | 同步数据时必填 | Tushare Pro token | | `TUSHARE_TOKEN` | 同步数据时必填 | Tushare Pro token |
| `LLM_API_KEY` | 使用 Agent 时必填 | 大模型 API Key(URL/模型名在 config.yaml) | | `LLM_API_KEY` | 使用 Agent 时必填 | 大模型 API Key(URL/模型名在 config.yaml) |
| `DATABASE_URL` | 否 | 留空使用 SQLite `<项目根>/data/quant.db` | | `DATABASE_URL` | 否 | 默认库 = `config.yaml → database.mysql`(MySQL `192.168.1.10/qlib`);设此项可覆盖(如切回 SQLite) |
| `MYSQL_PASSWORD` | 使用 MySQL 默认库时必填 | MySQL 密码(`config.yaml database.mysql.password_env` 引用;host/db/user 在 config.yaml) |
| `APP_SECRET_KEY` | 否 | 应用密钥(未接登录,可暂不改) | | `APP_SECRET_KEY` | 否 | 应用密钥(未接登录,可暂不改) |
> 约定:**密钥只放 `.env`**;URL、模型名等可配置项放 `config.yaml`。`.env` 已被 gitignore,严禁提交。 > 约定:**密钥只放 `.env`**;URL、模型名等可配置项放 `config.yaml`。`.env` 已被 gitignore,严禁提交。
@@ -76,7 +80,12 @@ cp .env.example .env
```yaml ```yaml
app: {name, version, debug, secret_key_env} app: {name, version, debug, secret_key_env}
api: {prefix: "/api"} api: {prefix: "/api"}
database: {url_env: "DATABASE_URL", echo: false, migrations_dir: ...} database:
url_env: "DATABASE_URL"
migrations_dir: ...
mysql: {enabled: true, host: "192.168.1.10", port: 3306,
db: "qlib", user: "qlib", password_env: "MYSQL_PASSWORD", charset: "utf8mb4"}
# URL 优先级:DATABASE_URL 环境变量 > database.mysql 组装 > sqlite:///./data/quant.db 兜底
data_source: {primary: "tushare", fallback: "sina", tushare_token_env: "TUSHARE_TOKEN"} data_source: {primary: "tushare", fallback: "sina", tushare_token_env: "TUSHARE_TOKEN"}
storage: {parquet_dir: "data/parquet", qlib_dir: "data/qlib", ...} # 相对项目根 storage: {parquet_dir: "data/parquet", qlib_dir: "data/qlib", ...} # 相对项目根
agent: agent:
@@ -94,12 +103,17 @@ agent:
```bash ```bash
cd backend cd backend
uv run alembic upgrade head # 建表(首次会自动建 data/quant.db) uv run alembic upgrade head # 在 config 指定库上建表(默认 MySQL qlib;SQLite 兜底库为 data/quant.db)
# 开发中改 Model 后: # 开发中改 Model 后:
uv run alembic revision --autogenerate -m "desc" uv run alembic revision --autogenerate -m "desc"
uv run alembic upgrade head uv run alembic upgrade head
# 对指定库执行(不随默认配置):
DATABASE_URL='mysql+pymysql://user:pass@host/db' uv run alembic upgrade head
``` ```
> 2026-09 已把数据从 SQLite(`data/quant.db`)全量迁移至 MySQL(`192.168.1.10:3306/qlib`),
> 迁移工具与校验见 `scripts/migrate_sqlite_to_mysql.py`(`--verify-only` 可复查一致性)。
--- ---
## 4. 数据同步与导出 ## 4. 数据同步与导出
@@ -171,6 +185,18 @@ pnpm dev
## 6. 研究 API 与研究引擎 ## 6. 研究 API 与研究引擎
### 6.0 引擎分层(M6–M8)
研究/选股链路自 M6 起分层:`Composite Engine`(因子复合分)→ `Selection Engine`
(条件/评分选股)→ `Signal Engine`(买卖信号)→ `Portfolio`(组合权重)→ 回测。
**当前选股与历史回测共用同一评分引擎**(`quant/composite.build_score_panel`),
保证 v2 §25「历史回测与当前选股用同一套逻辑」(一致性由测试锁定)。
- 选股:`POST /api/selections`,body 为 `SelectionQuery`(universe / factors / top_n /
as_of / method);结果候选含 `factor_values` 与 `selection_reason`(为什么选它)。
- 信号:`POST /api/signals`,body `{query, rules}`;rules 含买入排名阈值/趋势 MA/卖出区间。
- 策略:`POST /api/strategies` 保存命名策略 → `POST /{id}/expand` 补 period 展开为 spec 提交回测。
### 6.1 研究引擎:LocalEngine(默认)与 QlibEngine v1 ### 6.1 研究引擎:LocalEngine(默认)与 QlibEngine v1
服务通过 `QuantEngine` Protocol 注入引擎(`backend/app/quant/engine.py`): 服务通过 `QuantEngine` Protocol 注入引擎(`backend/app/quant/engine.py`):
@@ -217,7 +243,14 @@ curl -X POST http://127.0.0.1:8000/api/backtests \
| `GET /api/health` | 健康检查 | | `GET /api/health` | 健康检查 |
| `GET /api/stocks?q=600519&limit=20` | 股票列表/搜索 | | `GET /api/stocks?q=600519&limit=20` | 股票列表/搜索 |
| `GET /api/stocks/{symbol}` | 单只股票详情 | | `GET /api/stocks/{symbol}` | 单只股票详情 |
| `GET /api/factors` | 因子目录(公式/lookback/方向) | | `GET /api/factors` | 因子目录(读 `factor_definition` 表;空表自动 seed) |
| `POST /api/composites` / `GET /api/composites` | 因子组合保存 / 列表(方向由注册表填充) |
| `POST /api/selections` | 执行选股(`method=score` 评分 TopN / `condition` 条件)→ `SelectionResult` 并落库 |
| `GET /api/selections/{id}` / `GET /api/selections` | 读回 / 历史选股(`as_of`、`method` 过滤) |
| `POST /api/signals` | 生成 BUY/WATCH/SELL 信号(评分+规则)→ 落库 |
| `GET /api/signals/{id}` / `GET /api/signals` | 读回 / 历史信号 |
| `POST /api/strategies` / `GET /api/strategies` | 保存 / 列表命名策略 |
| `POST /api/strategies/{id}/expand` | 展开为回测 ResearchSpec(period + 初始资金) |
| `POST /api/factor-tests` | 同步单因子测试 → `FactorTestReport` | | `POST /api/factor-tests` | 同步单因子测试 → `FactorTestReport` |
| `POST /api/backtests` | 同步回测 → `BacktestResult`(默认 LocalEngine) | | `POST /api/backtests` | 同步回测 → `BacktestResult`(默认 LocalEngine) |
| `GET /api/backtests/last` | 最近一次同步回测 | | `GET /api/backtests/last` | 最近一次同步回测 |
@@ -258,10 +291,13 @@ curl -X POST http://127.0.0.1:8000/api/agent/chat \
} }
``` ```
Agent 能力边界(内置受控工具,只读): Agent 能力边界(10 个内置受控工具,只读 + 受控写库):
- `search_stocks` / `get_market_data`:查询 - `search_stocks` / `get_market_data`:查询
- `test_factor` / `run_backtest`:研究并**自动归档 Experiment** - `test_factor` / `run_backtest`:研究并**自动归档 Experiment**
- `screen_stocks` / `explain_selection`:按因子评分选股(传 `symbols` 白名单避免全市场长任务)/ 解释某次选股理由
- `generate_signals`:生成 BUY/WATCH/SELL 信号
- `create_strategy`:保存命名策略
- `get_experiment` / `compare_experiments`:读取/对比归档 - `get_experiment` / `compare_experiments`:读取/对比归档
**不提供** shell / 任意代码执行 / 删改数据 / 修改配置与凭证。未配置 Key 时接口返回 **不提供** shell / 任意代码执行 / 删改数据 / 修改配置与凭证。未配置 Key 时接口返回
@@ -275,7 +311,7 @@ Agent 能力边界(内置受控工具,只读):
```bash ```bash
cd backend cd backend
uv run ruff check app tests && uv run ruff format --check app tests uv run ruff check app tests && uv run ruff format --check app tests
uv run pytest # 当前 86 passed uv run pytest # 全量测试(每个里程碑提交前均须通过)
cd frontend/web cd frontend/web
pnpm run typecheck pnpm run typecheck
@@ -291,18 +327,21 @@ pnpm run build
## 9. 已知限制与说明 ## 9. 已知限制与说明
1. **QlibEngine v1 边界**:数据管线(落盘→D.features→回测)已打通,见 1. **Qlib 模型选股(Selection 模式 C)未实现**:当前选股为条件(A)/因子评分(B);
docs/QLIB_VERIFICATION.md;**Alpha158 全特征 + LightGBM walk-forward 预测信号仍未实现**, 模型预测选股(Alpha158 + LightGBM walk-forward,M8.4)为「按需」延后项
当前 QlibEngine 使用共享因子分回测(ROADMAP §2 备注为该增强 TODO)。 (见 docs/DEV_PLAN_v2.md §6.4)。
2. **数据规模**:仓库自带示例数据为 20 只权重股 2023–2024 日线(`sync --all` 可扩展 2. **全市场选股为同步请求**:无白名单的全市场选股约需 60s+;Agent 工具要求传 `symbols`
全市场,注意耗时与 Tushare 积分限制)。 白名单,全市场请在 Web 页执行(后续可迁异步 Job)。
3. **回测为近似建模**:涨跌停按收盘相对上一有效收盘判定、成交假设调仓日收盘,未建模 3. **数据规模**:当前 MySQL 已同步全市场(stock 5556 只、stock_daily 约 780 万行、
开盘一字 / 集合竞价 / 盘中路径(详见结果 `unimplemented`)。 adjust_factor 约 790 万行,2026-09 由 SQLite 迁移并校验一致)。
4. **Agent 结论质量取决于模型**:编排只保证「经受控工具 + 归档留痕」,研究有效性判断 4. **回测为近似建模**:涨跌停按收盘相对上一有效收盘判定、成交假设调仓日收盘,
Portfolio v1 仅等权(单股/行业上限约束字段已预留但未建模,设置后会在结果
`unimplemented` 如实标注)(详见结果 `unimplemented`)。
5. **Agent 结论质量取决于模型**:编排只保证「经受控工具 + 归档留痕」,研究有效性判断
需要人复核;接不同厂商模型请核对 `config.yaml` 的 `base_url` 与 `model` 命名。 需要人复核;接不同厂商模型请核对 `config.yaml` 的 `base_url` 与 `model` 命名。
5. **SSE 与 Job 为单进程本地执行**:重启进程后未完成 Job 需重新提交(第一阶段刻意 6. **SSE 与 Job 为单进程本地执行**:重启进程后未完成 Job 需重新提交(第一阶段刻意
不引入 Redis/Celery)。 不引入 Redis/Celery;本机 Redis 已就绪,出现排队/长任务需求时按 DEV_PLAN §7 接入)。
6. **Qlib 数据为进程内单例初始化**:`qlib.init` 以 provider_uri 为键幂等,切换不同 7. **Qlib 数据为进程内单例初始化**:`qlib.init` 以 provider_uri 为键幂等,切换不同
qlib_dir 会重新初始化;测试环境建议注入独立临时目录。 qlib_dir 会重新初始化;测试环境建议注入独立临时目录。
--- ---
@@ -320,9 +359,12 @@ pnpm run build
- **前端连不上后端 / NetworkError**:前端默认经 Next **同源代理**访问 `/api/*` - **前端连不上后端 / NetworkError**:前端默认经 Next **同源代理**访问 `/api/*`
(next.config.ts rewrites → `127.0.0.1:8000`,可用 `BACKEND_API_URL` 覆盖), (next.config.ts rewrites → `127.0.0.1:8000`,可用 `BACKEND_API_URL` 覆盖),
任意 IP 访问 `:3000` 都不需要 CORS 或硬编码后端地址;后端 CORS 开发期为 `*`。 任意 IP 访问 `:3000` 都不需要 CORS 或硬编码后端地址;后端 CORS 开发期为 `*`。
若 API 返回 500 且日志出现 `database is locked`,多半是正在跑全市场数据同步 若使用 SQLite 兜底库且日志出现 `database is locked`,多半是正在跑全市场数据
(长写事务),同步结束后自动恢复(引擎已加 busy_timeout 等待)。 同步(长写事务),同步结束后自动恢复(SQLite 连接已加 busy_timeout);
- **数据库被改动想重置**:删除 `data/quant.db` 后 `uv run alembic upgrade head` 重建 默认 MySQL 库无此问题。
表结构(行情需重新同步)。 - **数据库被改动想重置**:默认 MySQL(qlib@192.168.1.10)时在远端重建后
- **想切换 MySQL**:`.env` 设 `DATABASE_URL=mysql+pymysql://user:pass@host/db`, `uv run alembic upgrade head`(行情需重新同步或从备份恢复);若切回 SQLite 兜底库则删除
业务层无需改动(Repository 已隔离)。 `data/quant.db` 后重建。
- **默认库已是 MySQL**(config.yaml `database.mysql`,密码在 `.env` 的 `MYSQL_PASSWORD`);
**想临时切回 SQLite**:`.env` 设 `DATABASE_URL=sqlite:///./data/quant.db`。
业务层均无需改动(Repository / SQLAlchemy 已隔离方言差异)。
+78
View File
@@ -28,6 +28,10 @@ export default function BacktestPage() {
const [start, setStart] = useState(""); const [start, setStart] = useState("");
const [end, setEnd] = useState(""); const [end, setEnd] = useState("");
const [excludeSt, setExcludeSt] = useState(true); const [excludeSt, setExcludeSt] = useState(true);
const [commission, setCommission] = useState(0.03); // %
const [stamp, setStamp] = useState(0.05);
const [slippage, setSlippage] = useState(0.1);
const [capital, setCapital] = useState(1_000_000);
const [result, setResult] = useState<BacktestResult | null>(null); const [result, setResult] = useState<BacktestResult | null>(null);
const [running, setRunning] = useState(false); const [running, setRunning] = useState(false);
const [jobId, setJobId] = useState(""); const [jobId, setJobId] = useState("");
@@ -65,6 +69,12 @@ export default function BacktestPage() {
factors: [{ name: factor, weight: 1 }], factors: [{ name: factor, weight: 1 }],
selection: { top_n: topN }, selection: { top_n: topN },
rebalance, rebalance,
costs: {
commission_rate: commission / 100,
stamp_tax_rate: stamp / 100,
slippage_rate: slippage / 100,
},
initial_capital: capital,
period: [start, end], period: [start, end],
}; };
// 异步 Job:后台执行(全市场可能数十秒),轮询到终态 // 异步 Job:后台执行(全市场可能数十秒),轮询到终态
@@ -133,6 +143,22 @@ export default function BacktestPage() {
<option value="weekly">周度</option> <option value="weekly">周度</option>
</select> </select>
</Field> </Field>
<Field label="手续费率 %">
<input className="input" type="number" step="0.01" min={0} value={commission}
onChange={(e) => setCommission(Number(e.target.value))} />
</Field>
<Field label="印花税率 %">
<input className="input" type="number" step="0.01" min={0} value={stamp}
onChange={(e) => setStamp(Number(e.target.value))} />
</Field>
<Field label="滑点率 %">
<input className="input" type="number" step="0.01" min={0} value={slippage}
onChange={(e) => setSlippage(Number(e.target.value))} />
</Field>
<Field label="初始资金(元)">
<input className="input" type="number" step={100000} min={10000} value={capital}
onChange={(e) => setCapital(Number(e.target.value))} />
</Field>
<Field label="开始日期"> <Field label="开始日期">
<input type="date" className="input" value={start} onChange={(e) => setStart(e.target.value)} /> <input type="date" className="input" value={start} onChange={(e) => setStart(e.target.value)} />
</Field> </Field>
@@ -228,6 +254,58 @@ function ResultView({ result }: { result: BacktestResult }) {
)} )}
</Card> </Card>
{result.yearly_returns.length ? (
<Card icon="calendar" title="年度收益(%)" tools={<Pill>均值 {(result.yearly_returns.reduce((a, y) => a + y.return_pct, 0) / result.yearly_returns.length).toFixed(2)}%</Pill>}>
<div className="chips">
{result.yearly_returns.map((y) => (
<span className="chip" key={y.year}>
<b>{y.year}</b>
<span className={y.return_pct >= 0 ? "tone-pos" : "tone-neg"}>
{y.return_pct.toFixed(2)}%
</span>
</span>
))}
</div>
</Card>
) : null}
{result.trades.length ? (
<Card
icon="scale"
title={`成交明细 · ${result.trades.length} 笔`}
tools={<Pill>{result.turnover_pct.toFixed(2)}% 累计换手</Pill>}
>
<div className="table-wrap">
<table className="tbl">
<thead>
<tr>
<th>买入日</th>
<th>卖出日</th>
<th>代码</th>
<th>买价</th>
<th>卖价</th>
<th>收益</th>
</tr>
</thead>
<tbody>
{result.trades.map((t) => (
<tr key={`${t.symbol}-${t.entry_date}-${t.exit_date}`}>
<td className="mono dim">{t.entry_date}</td>
<td className="mono dim">{t.exit_date}</td>
<td className="mono"><b>{t.symbol}</b></td>
<td className="mono">{t.entry_price?.toFixed(2) ?? "-"}</td>
<td className="mono">{t.exit_price?.toFixed(2) ?? "-"}</td>
<td className={t.return_pct >= 0 ? "tone-pos" : "tone-neg"}>
{t.return_pct.toFixed(2)}%
</td>
</tr>
))}
</tbody>
</table>
</div>
</Card>
) : null}
<UnimplementedNote items={result.unimplemented} /> <UnimplementedNote items={result.unimplemented} />
</> </>
); );
+15 -1
View File
@@ -36,6 +36,20 @@ const LINKS = [
icon: "gauge" as const, icon: "gauge" as const,
tone: "quick-card--warn", tone: "quick-card--warn",
}, },
{
href: "/selection",
title: "股票筛选",
desc: "因子评分 TopN 或条件选股,支持当前/历史时点,结果可解释",
icon: "target" as const,
tone: "quick-card--violet",
},
{
href: "/signals",
title: "交易信号",
desc: "按评分排名 + 趋势规则输出 BUY / WATCH / SELL 及理由",
icon: "scale" as const,
tone: "quick-card--pos",
},
{ {
href: "/stocks", href: "/stocks",
title: "股票池", title: "股票池",
@@ -135,7 +149,7 @@ export default function DashboardPage() {
因子/回测请求经 <b>异步 Job</b> 后台执行并自动归档为实验(Experiment),可在「实验」页查看与复跑。 因子/回测请求经 <b>异步 Job</b> 后台执行并自动归档为实验(Experiment),可在「实验」页查看与复跑。
</li> </li>
<li> <li>
数据底座为<b>本地 SQLite + Parquet</b>:行情经 <code>sync daily --resume</code> 增量同步; 数据底座为<b>MySQL + Parquet</b>:行情经 <code>sync daily --resume</code> 增量同步;
新浪数据仅在「两边一致」校验通过后兜底入库(<code>source=sina</code>)。 新浪数据仅在「两边一致」校验通过后兜底入库(<code>source=sina</code>)。
</li> </li>
<li> <li>
+387
View File
@@ -0,0 +1,387 @@
"use client";
/** 股票筛选(选股系统 M6,v2 §14/§21.1):
* 评分模式(多因子加权 TopN)与条件模式(结构化条件 AND);
* 支持当前/历史 as_of;结果落库可查(selection_id),右侧展示最近选股历史。
*/
import { useEffect, useState } from "react";
import { apiGet, apiPost } from "@/lib/api";
import type {
FactorMeta,
SelectionCondition,
SelectionMeta,
SelectionQuery,
SelectionResult,
SelectionRun,
} from "@/lib/types";
import {
PageHeader,
Card,
Pill,
Field,
Btn,
Banner,
Empty,
SkeletonLines,
} from "@/components/ui";
const FIELD_OPTIONS: { value: string; label: string; kind: "num" | "str" | "ref" }[] = [
{ value: "static.industry", label: "行业 industry", kind: "str" },
{ value: "static.market", label: "市场 market", kind: "str" },
{ value: "momentum_20", label: "动量 20 日", kind: "num" },
{ value: "momentum_60", label: "动量 60 日", kind: "num" },
{ value: "close", label: "收盘价 close", kind: "num" },
{ value: "close_ma60_ref", label: "收盘 > MA60", kind: "ref" },
{ value: "volume_ratio_5_60", label: "量比", kind: "num" },
{ value: "fundamental.roe", label: "ROE(%)", kind: "num" },
{ value: "fundamental.eps", label: "EPS", kind: "num" },
];
const OPS = ["gt", "gte", "lt", "lte", "eq", "ne", "in"];
function emptyCondition(): SelectionCondition {
return { field: "static.industry", op: "eq", value: "白酒" };
}
export default function SelectionPage() {
const [factors, setFactors] = useState<FactorMeta[]>([]);
const [mode, setMode] = useState<"score" | "condition">("score");
const [asOf, setAsOf] = useState(""); // 空 = 最近交易日
const [excludeSt, setExcludeSt] = useState(true);
const [topN, setTopN] = useState(20);
const [factorRows, setFactorRows] = useState([{ name: "momentum_60", weight: 1 }]);
const [conds, setConds] = useState<SelectionCondition[]>([emptyCondition()]);
const [result, setResult] = useState<SelectionResult | null>(null);
const [selectionId, setSelectionId] = useState("");
const [running, setRunning] = useState(false);
const [error, setError] = useState("");
const [history, setHistory] = useState<SelectionMeta[] | null>(null);
useEffect(() => {
let alive = true;
apiGet<FactorMeta[]>("/factors")
.then((list) => alive && list.length > 0 && setFactors(list))
.catch((e: Error) => alive && setError(e.message));
apiGet<SelectionMeta[]>("/selections?limit=8")
.then((rows) => alive && setHistory(rows))
.catch(() => alive && setHistory([]));
return () => {
alive = false;
};
}, []);
function buildQuery(): SelectionQuery {
if (mode === "score") {
return {
universe: { exclude_st: excludeSt, min_listing_days: 0 },
as_of: asOf || null,
method: "score",
factors: factorRows.filter((r) => r.name),
conditions: [],
top_n: topN,
};
}
return {
universe: { exclude_st: excludeSt, min_listing_days: 0 },
as_of: asOf || null,
method: "condition",
factors: [],
conditions: conds.map((c) => ({
...c,
...(c.field === "close_ma60_ref"
? { field: "close", op: c.op, ref: "ma60", value: null }
: {}),
})),
top_n: null,
};
}
async function run() {
setRunning(true);
setError("");
setResult(null);
try {
const run = await apiPost<SelectionRun>("/selections", buildQuery());
setSelectionId(run.selection_id);
setResult(run.result);
apiGet<SelectionMeta[]>("/selections?limit=8").then(setHistory).catch(() => undefined);
} catch (e) {
setError((e as Error).message);
} finally {
setRunning(false);
}
}
async function openHistory(id: string) {
setError("");
try {
const detail = await apiGet<SelectionResult>(`/selections/${id}`);
setSelectionId(id);
setResult(detail);
} catch (e) {
setError((e as Error).message);
}
}
return (
<>
<PageHeader
title="股票筛选"
sub="以因子评分 TopN 或结构化条件筛选股票,支持当前 / 历史时点;每次选股落库可复现、可解释(为什么选它)。"
actions={
selectionId ? (
<Pill tone="violet" icon="check">
{selectionId}
</Pill>
) : null
}
/>
<Card title="选股条件" icon="filter">
<div className="form-grid">
<Field label="方式">
<select
className="input"
value={mode}
onChange={(e) => setMode(e.target.value as "score" | "condition")}
>
<option value="score">因子评分(TopN)</option>
<option value="condition">条件筛选(AND)</option>
</select>
</Field>
<Field label="选股时点 as_of" hint="留空 = 最近可用交易日">
<input type="date" className="input" value={asOf} onChange={(e) => setAsOf(e.target.value)} />
</Field>
<Field label="标的范围">
<label className="row" style={{ gap: 6, color: "var(--text-2)", fontSize: 13, cursor: "pointer" }}>
<input type="checkbox" checked={excludeSt} onChange={(e) => setExcludeSt(e.target.checked)} />
剔除 ST
</label>
</Field>
</div>
{mode === "score" ? (
<>
<div style={{ fontSize: 13, color: "var(--text-2)", margin: "10px 0 6px" }}>评分因子与权重</div>
{factorRows.map((r, i) => (
<div className="row" style={{ gap: 8, marginBottom: 8 }} key={i}>
<select
className="input"
style={{ flex: 2 }}
value={r.name}
onChange={(e) => {
const next = [...factorRows];
next[i] = { ...next[i], name: e.target.value };
setFactorRows(next);
}}
>
{factors.map((f) => (
<option key={f.name} value={f.name}>{f.name}</option>
))}
</select>
<input
className="input"
style={{ flex: 1 }}
type="number"
step="0.1"
value={r.weight}
onChange={(e) => {
const next = [...factorRows];
next[i] = { ...next[i], weight: Number(e.target.value) };
setFactorRows(next);
}}
/>
<button
className="btn"
disabled={factorRows.length <= 1}
onClick={() => setFactorRows(factorRows.filter((_, j) => j !== i))}
>
删除
</button>
</div>
))}
<div className="row" style={{ gap: 10 }}>
<button className="btn" onClick={() => setFactorRows([...factorRows, { name: "momentum_60", weight: 1 }])}>
+ 因子
</button>
<Field label="Top N">
<input className="input" type="number" min={1} value={topN} onChange={(e) => setTopN(Number(e.target.value))} />
</Field>
</div>
</>
) : (
<>
<div style={{ fontSize: 13, color: "var(--text-2)", margin: "10px 0 6px" }}>条件(全部满足才入选)</div>
{conds.map((c, i) => {
const meta = FIELD_OPTIONS.find((f) => f.value === (c.field === "close" && c.ref ? "close_ma60_ref" : c.field));
const isRef = c.field === "close" && c.ref === "ma60";
return (
<div className="row" style={{ gap: 8, marginBottom: 8, flexWrap: "wrap" }} key={i}>
<select
className="input"
style={{ flex: 2 }}
value={isRef ? "close_ma60_ref" : c.field}
onChange={(e) => {
const next = [...conds];
const v = e.target.value;
next[i] = v === "close_ma60_ref"
? { field: "close", op: "gt", ref: "ma60", value: null }
: { field: v, op: c.op, value: v.startsWith("static.") ? "" : 0 };
setConds(next);
}}
>
{FIELD_OPTIONS.map((f) => (
<option key={f.value} value={f.value}>{f.label}</option>
))}
</select>
<select
className="input"
style={{ flex: 1 }}
value={c.op}
onChange={(e) => {
const next = [...conds];
next[i] = { ...next[i], op: e.target.value as SelectionCondition["op"] };
setConds(next);
}}
>
{OPS.map((o) => <option key={o} value={o}>{o}</option>)}
</select>
{!isRef ? (
<input
className="input"
style={{ flex: 1 }}
placeholder={c.op === "in" ? "逗号分隔多个值" : "值"}
value={String(c.value ?? "")}
onChange={(e) => {
const next = [...conds];
const raw = e.target.value;
const val = c.op === "in"
? raw.split(/[,,]/).map((s) => s.trim()).filter(Boolean)
: (meta?.kind === "num" && !c.field.startsWith("static.")
? Number(raw)
: raw);
next[i] = { ...next[i], value: val as never };
setConds(next);
}}
/>
) : null}
<button
className="btn"
disabled={conds.length <= 1}
onClick={() => setConds(conds.filter((_, j) => j !== i))}
>
删除
</button>
</div>
);
})}
<button className="btn" onClick={() => setConds([...conds, emptyCondition()])}>+ 条件</button>
</>
)}
<div style={{ marginTop: 14 }}>
<Btn variant="primary" icon="play" loading={running} disabled={running} onClick={run}>
{running ? "筛选中…" : "执行选股"}
</Btn>
</div>
{error ? (
<div style={{ marginTop: 12 }}>
<Banner tone="error">{error}</Banner>
</div>
) : null}
</Card>
{!result && !running && history === null ? <SkeletonLines n={4} /> : null}
{result ? <ResultView result={result} /> : null}
{history && history.length > 0 ? (
<Card title="最近选股记录" icon="archive">
<div className="chips">
{history.map((m) => (
<button
key={m.id}
className="chip"
style={{ cursor: "pointer" }}
onClick={() => openHistory(m.id)}
title={`${m.id} · ${m.method} · ${m.selected} 只`}
>
<b className="mono">{m.as_of}</b>
{m.method} · {m.selected} 只
</button>
))}
</div>
</Card>
) : null}
</>
);
}
function ResultView({ result }: { result: SelectionResult }) {
const s = result.statistics;
if (result.candidates.length === 0) {
return (
<Card>
<Empty
icon="filter"
title="无符合条件/可评分的股票"
hint={`范围 ${s.universe_size} 只,可评分 ${s.evaluated} 只。可放宽条件或选择更早的 as_of。`}
/>
</Card>
);
}
return (
<Card
icon="target"
title={`选股结果 · ${result.as_of_date} · ${result.method}`}
tools={
<>
<Pill>范围 {s.universe_size}</Pill>
<Pill>评分 {s.evaluated}</Pill>
<Pill tone="pos">选出 {s.selected}</Pill>
</>
}
>
<div className="table-wrap">
<table className="tbl">
<thead>
<tr>
<th>#</th>
<th>代码</th>
<th>得分</th>
<th>因子值</th>
<th>入选理由</th>
</tr>
</thead>
<tbody>
{result.candidates.map((c) => (
<tr key={c.symbol}>
<td className="mono dim">{c.rank}</td>
<td className="mono">
<b>{c.symbol}</b>
</td>
<td className="mono">{c.score.toFixed(4)}</td>
<td style={{ fontSize: 12 }}>
{Object.entries(c.factor_values)
.map(([k, v]) => `${k}=${typeof v === "number" ? v.toFixed(4) : v}`)
.join(" ")}
</td>
<td>
<div className="chips">
{c.selection_reason.map((r) => (
<span className="chip" key={r}>{r.slice(0, 40)}</span>
))}
</div>
</td>
</tr>
))}
</tbody>
</table>
</div>
{result.unimplemented.length ? (
<div style={{ marginTop: 10 }}>
<Banner tone="info">{result.unimplemented.join(";")}</Banner>
</div>
) : null}
</Card>
);
}
+200
View File
@@ -0,0 +1,200 @@
"use client";
/** 交易信号(v2 §15):基于选股评分排名 + 趋势规则生成 BUY / WATCH / SELL,
* 每条带理由(为什么);支持当前/历史 as_of,历史记录点击回看。
*/
import { useEffect, useState } from "react";
import { apiGet, apiPost } from "@/lib/api";
import type {
FactorMeta,
SelectionQuery,
SignalMeta,
SignalResult,
SignalRun,
} from "@/lib/types";
import {
PageHeader,
Card,
Pill,
Field,
Btn,
Banner,
Empty,
SkeletonLines,
} from "@/components/ui";
export default function SignalsPage() {
const [factors, setFactors] = useState<FactorMeta[]>([]);
const [factor, setFactor] = useState("momentum_60");
const [asOf, setAsOf] = useState("");
const [symbols, setSymbols] = useState("600519.SH,000001.SZ,300750.SZ,601318.SH,000858.SZ");
const [buyRank, setBuyRank] = useState(2);
const [sellRank, setSellRank] = useState(5);
const [result, setResult] = useState<SignalResult | null>(null);
const [signalId, setSignalId] = useState("");
const [running, setRunning] = useState(false);
const [error, setError] = useState("");
const [history, setHistory] = useState<SignalMeta[] | null>(null);
useEffect(() => {
let alive = true;
apiGet<FactorMeta[]>("/factors")
.then((list) => alive && list.length > 0 && setFactors(list))
.catch((e: Error) => alive && setError(e.message));
apiGet<SignalMeta[]>("/signals?limit=8").then(setHistory).catch(() => alive && setHistory([]));
return () => {
alive = false;
};
}, []);
function buildQuery(): SelectionQuery {
return {
universe: { exclude_st: true, min_listing_days: 0, symbols: symbols.split(/[,,]/).map((s) => s.trim()).filter(Boolean) },
as_of: asOf || null,
method: "score",
factors: [{ name: factor, weight: 1 }],
conditions: [],
top_n: sellRank + 10,
};
}
async function run() {
setRunning(true);
setError("");
setResult(null);
try {
const run = await apiPost<SignalRun>("/signals", {
query: buildQuery(),
rules: { buy_rank_threshold: buyRank, sell_rank_threshold: sellRank, max_output_rank: 60 },
});
setSignalId(run.signal_id);
setResult(run.result);
apiGet<SignalMeta[]>("/signals?limit=8").then(setHistory).catch(() => undefined);
} catch (e) {
setError((e as Error).message);
} finally {
setRunning(false);
}
}
async function openHistory(id: string) {
try {
const detail = await apiGet<SignalResult>(`/signals/${id}`);
setSignalId(id);
setResult(detail);
} catch (e) {
setError((e as Error).message);
}
}
return (
<>
<PageHeader
title="交易信号"
sub="基于选股综合分排名与趋势规则输出 BUY / WATCH / SELL,每条带触发理由;与回测共用同一评分引擎。"
actions={signalId ? <Pill tone="violet" icon="check">{signalId}</Pill> : null}
/>
<Card title="信号参数" icon="scale">
<div className="form-grid">
<Field label="评分因子">
<select className="input" value={factor} onChange={(e) => setFactor(e.target.value)}>
{factors.map((f) => (
<option key={f.name} value={f.name}>{f.name}</option>
))}
</select>
</Field>
<Field label="买入排名阈值 (≤)">
<input className="input" type="number" min={1} value={buyRank} onChange={(e) => setBuyRank(Number(e.target.value))} />
</Field>
<Field label="卖出/警示排名阈值 (>)">
<input className="input" type="number" min={1} value={sellRank} onChange={(e) => setSellRank(Number(e.target.value))} />
</Field>
<Field label="时点 as_of" hint="留空 = 最近交易日">
<input type="date" className="input" value={asOf} onChange={(e) => setAsOf(e.target.value)} />
</Field>
<Field label="股票白名单" hint="逗号分隔(≤60)">
<input className="input" value={symbols} onChange={(e) => setSymbols(e.target.value)} />
</Field>
<Btn variant="primary" icon="play" loading={running} disabled={running} onClick={run}>
{running ? "生成中…" : "生成信号"}
</Btn>
</div>
{error ? <div style={{ marginTop: 12 }}><Banner tone="error">{error}</Banner></div> : null}
</Card>
{!result && !running && history === null ? <SkeletonLines n={4} /> : null}
{result ? <ResultView result={result} /> : null}
{history && history.length > 0 ? (
<Card title="最近信号记录" icon="archive">
<div className="chips">
{history.map((m) => (
<button key={m.id} className="chip" style={{ cursor: "pointer" }} onClick={() => openHistory(m.id)} title={m.id}>
<b className="mono">{m.as_of}</b>
买 {m.buy} / 观 {m.watch} / 卖 {m.sell}
</button>
))}
</div>
</Card>
) : null}
</>
);
}
function ResultView({ result }: { result: SignalResult }) {
const st = result.statistics;
return (
<Card
icon="scale"
title={`信号 · ${result.as_of_date}`}
tools={
<>
<Pill tone="pos">BUY {st.buy}</Pill>
<Pill>WATCH {st.watch}</Pill>
<Pill tone="neg">SELL {st.sell}</Pill>
</>
}
>
{result.events.length === 0 ? (
<Empty icon="scale" title="无信号输出" hint="放宽排名阈值或检查白名单数据。" />
) : (
<div className="table-wrap">
<table className="tbl">
<thead>
<tr>
<th>类型</th>
<th>代码</th>
<th>得分</th>
<th>价格</th>
<th>触发理由</th>
</tr>
</thead>
<tbody>
{result.events.map((e) => (
<tr key={e.symbol}>
<td>
<Pill tone={e.signal_type === "BUY" ? "pos" : e.signal_type === "SELL" ? "neg" : undefined}>
{e.signal_type}
</Pill>
</td>
<td className="mono"><b>{e.symbol}</b></td>
<td className="mono">{e.score?.toFixed(4)}</td>
<td className="mono">{e.price ?? "-"}</td>
<td style={{ fontSize: 12 }}>
<div className="chips">
{e.trigger_reason.map((r) => (
<span className="chip" key={r}>{r}</span>
))}
</div>
</td>
</tr>
))}
</tbody>
</table>
</div>
)}
</Card>
);
}
+3 -1
View File
@@ -18,6 +18,8 @@ const GROUPS: { title: string; items: NavItem[] }[] = [
items: [ items: [
{ href: "/", label: "总览", icon: "grid" }, { href: "/", label: "总览", icon: "grid" },
{ href: "/stocks", label: "股票池", icon: "candles" }, { href: "/stocks", label: "股票池", icon: "candles" },
{ href: "/selection", label: "股票筛选", icon: "target" },
{ href: "/signals", label: "交易信号", icon: "scale" },
{ href: "/factors", label: "因子研究", icon: "flask" }, { href: "/factors", label: "因子研究", icon: "flask" },
{ href: "/factors/compose", label: "因子组合", icon: "layers" }, { href: "/factors/compose", label: "因子组合", icon: "layers" },
{ href: "/backtest", label: "选股回测", icon: "gauge" }, { href: "/backtest", label: "选股回测", icon: "gauge" },
@@ -109,7 +111,7 @@ export function AppShell({ children }: { children: React.ReactNode }) {
<div className="rail-foot"> <div className="rail-foot">
数据源 Tushare 首选 · 新浪校验兜底 数据源 Tushare 首选 · 新浪校验兜底
<br /> <br />
研究引擎 Qlib(本地 SQLite / Parquet) 数据库 MySQL · 研究引擎 Qlib
</div> </div>
</aside> </aside>
+115 -2
View File
@@ -22,10 +22,23 @@ export interface FactorMeta {
export interface ResearchSpec { export interface ResearchSpec {
type: "factor_test" | "backtest"; type: "factor_test" | "backtest";
universe: { exclude_st?: boolean; min_listing_days?: number }; universe: {
exclude_st?: boolean;
min_listing_days?: number;
symbols?: string[];
};
price_adjustment?: "none" | "qfq";
factors: { name: string; weight: number }[]; factors: { name: string; weight: number }[];
selection: { top_n: number }; selection: { top_n: number };
rebalance: "weekly" | "monthly"; rebalance: "weekly" | "monthly";
costs?: {
commission_rate?: number;
stamp_tax_rate?: number;
slippage_rate?: number;
benchmark?: string;
};
portfolio?: { weighting?: string; max_position_pct?: number | null; max_industry_weight_pct?: number | null };
initial_capital?: number;
period: [string, string]; period: [string, string];
} }
@@ -56,7 +69,7 @@ export interface BacktestResult {
monthly_returns: { year: number; month: number; return_pct: number }[]; monthly_returns: { year: number; month: number; return_pct: number }[];
yearly_returns: { year: number; return_pct: number }[]; yearly_returns: { year: number; return_pct: number }[];
positions: { date: string; symbol: string; weight: number }[]; positions: { date: string; symbol: string; weight: number }[];
trades: { entry_date: string; exit_date: string; symbol: string; return_pct: number }[]; trades: { entry_date: string; exit_date: string; symbol: string; entry_price?: number; exit_price?: number; return_pct: number }[];
turnover_pct: number; turnover_pct: number;
unimplemented: string[]; unimplemented: string[];
config_snapshot: Record<string, unknown>; config_snapshot: Record<string, unknown>;
@@ -73,3 +86,103 @@ export interface FactorTestReport {
sample_days: number; sample_days: number;
unimplemented: string[]; unimplemented: string[];
} }
/* ---- 选股系统(v2 §14/§21.1,与 domain/entities/selection.py 对应) ---- */
export interface SelectionCondition {
field: string;
op: "gt" | "gte" | "lt" | "lte" | "eq" | "ne" | "in" | "not_in";
value?: number | string | (number | string)[] | null;
ref?: string | null;
}
export interface SelectionQuery {
universe: {
exclude_st?: boolean;
min_listing_days?: number;
symbols?: string[];
};
as_of?: string | null; // null/省略 = 最近交易日
method: "score" | "condition";
factors: { name: string; weight: number }[];
conditions: SelectionCondition[];
top_n?: number | null;
top_pct?: number | null;
min_score?: number | null;
}
export interface SelectionCandidate {
symbol: string;
rank: number;
score: number;
factor_values: Record<string, number>;
filter_status: string[];
selection_reason: string[];
}
export interface SelectionResult {
as_of_date: string;
method: string;
statistics: { universe_size: number; evaluated: number; selected: number };
candidates: SelectionCandidate[];
unimplemented: string[];
config_snapshot: Record<string, unknown>;
}
export interface SelectionMeta {
id: string;
as_of: string;
method: string;
universe_size: number;
selected: number;
created_at?: string | null;
}
export interface SelectionRun {
selection_id: string;
result: SelectionResult;
}
/* ---- 交易信号(v2 §15,与 domain/entities/signal.py 对应) ---- */
export interface SignalRules {
buy_rank_threshold?: number;
buy_require_trend?: boolean;
buy_require_momentum?: boolean;
trend_ma?: number;
sell_rank_threshold?: number;
sell_on_trend_break?: boolean;
max_output_rank?: number;
}
export interface SignalEvent {
symbol: string;
signal_date: string;
signal_type: "BUY" | "WATCH" | "SELL";
score?: number | null;
price?: number | null;
trigger_reason: string[];
}
export interface SignalResult {
as_of_date: string;
rules: SignalRules;
statistics: { universe_size: number; buy: number; watch: number; sell: number };
events: SignalEvent[];
config_snapshot: Record<string, unknown>;
}
export interface SignalMeta {
id: string;
as_of: string;
buy: number;
watch: number;
sell: number;
created_at?: string | null;
}
export interface SignalRun {
signal_id: string;
result: SignalResult;
}
+280
View File
@@ -0,0 +1,280 @@
#!/usr/bin/env python3
"""SQLite → MySQL 全量数据迁移脚本(一次性工具,不入应用代码)。
把 <项目根>/data/quant.db(SQLite)的全部业务表数据迁移到 MySQL 目标库,
供后端以 MySQL 作为默认数据库运行(docs/DEV_PLAN「数据库迁移」专项)。
设计:
- 源:标准库 sqlite3 直连(只读查询)。
- 目标:SQLAlchemy engine(URL 形态任意,mysql+pymysql 由配置/env 提供),
用 raw_connection() 拿到 pymysql 连接做 chunk 多值 INSERT 批量写入。
- 逐表 keyset 分页(int 主键按 id > last),保留原始主键 id。
- 幂等:目标表已有数据时从 max(id)+1 续传;string 主键小表整表一次。
选项 --reset 先 TRUNCATE 再全量重传。
- 完成后逐表 COUNT 校验 + 抽样逐行比对。
用法:
cd backend && .venv/bin/python ../scripts/migrate_sqlite_to_mysql.py [--reset] [--verify-only]
"""
from __future__ import annotations
import argparse
import os
import sqlite3
import sys
import time
from decimal import Decimal
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
BACKEND = ROOT / "backend"
DEFAULT_SOURCE = ROOT / "data" / "quant.db"
# 表 → (列清单, 主键类型)。列顺序即插入顺序;全部显式带 id。
TABLES: dict[str, tuple[list[str], str]] = {
"stock": (
["id", "symbol", "name", "industry", "area", "market", "exchange",
"list_date", "delist_date", "status"],
"int",
),
"trading_calendar": (["id", "calendar_date", "is_open"], "int"),
"sync_log": (
["id", "source", "api", "request_time", "success", "failure_reason",
"row_count", "data_start", "data_end"],
"int",
),
"job": (
["id", "kind", "status", "stage", "spec_json", "error", "result_json",
"experiment_id", "created_at", "started_at", "finished_at"],
"str",
),
"experiment": (
["id", "kind", "spec_json", "result_json", "summary_text",
"code_version", "data_version", "job_id", "created_at"],
"str",
),
"financial_indicator": (
["id", "symbol", "report_date", "announce_date", "source", "eps", "roe",
"total_revenue", "net_profit", "gross_margin"],
"int",
),
"stock_daily": (
["id", "symbol", "trade_date", "source", "adjust", "open", "high",
"low", "close", "volume", "amount"],
"int",
),
"adjust_factor": (["id", "symbol", "trade_date", "factor"], "int"),
}
CHUNK = 5000 # 每批行数:5000 × 11 列占位符 ≈ 5.5 万 < MySQL 65535 上限
PROGRESS_EVERY = 200_000
def resolve_dst_engine():
from sqlalchemy import create_engine
url = os.environ.get("DATABASE_URL")
if not url:
sys.path.insert(0, str(BACKEND))
from app.core.config import get_settings
url = get_settings().database_url
print(f"目标 URL: {url.split('@')[-1]}", flush=True)
return create_engine(url, future=True, pool_pre_ping=True)
def migrate_table(src: sqlite3.Connection, dst, table: str, cols: list[str], pk_kind: str,
*, reset: bool, throttle_sec: float = 0.0) -> int:
col_list = ", ".join(cols) # sqlite 侧
col_q = ", ".join(f"`{c}`" for c in cols) # mysql 侧
src_total = src.execute(f'SELECT COUNT(*) FROM "{table}"').fetchone()[0]
cur = dst.cursor()
try:
if pk_kind == "str":
# 小表(job/experiment):整表读入,幂等靠唯一主键冲突忽略
if reset:
cur.execute(f"TRUNCATE TABLE `{table}`")
dst.commit()
elif count(dst, table):
print(f"[{table}] 目标非空且未 --reset,跳过(校验阶段会核对)")
return 0
rows = src.execute(f'SELECT {col_list} FROM "{table}"').fetchall()
if rows:
_insert_chunk(cur, table, col_q, cols, rows)
dst.commit()
print(f"[{table}] 完成 {len(rows)} 行")
return len(rows)
# int 主键:keyset 分页
if reset:
cur.execute(f"TRUNCATE TABLE `{table}`")
dst.commit()
last = 0
else:
cur.execute(f"SELECT COALESCE(MAX(id), 0) FROM `{table}`")
last = int(cur.fetchone()[0])
if last:
print(f"[{table}] 目标已有数据,从 id={last + 1} 续传")
inserted = 0
t0 = time.time()
commit_every = max(100_000 // CHUNK, 1) * CHUNK # 每 ~10 万行 commit 一次
while True:
rows = src.execute(
f'SELECT {col_list} FROM "{table}" WHERE id > ? ORDER BY id LIMIT ?',
(last, CHUNK),
).fetchall()
if not rows:
break
_insert_chunk(cur, table, col_q, cols, rows) # execute 但不立即 commit
inserted += len(rows)
last = rows[-1][0]
if throttle_sec:
time.sleep(throttle_sec) # 低配 MySQL:每批间限速,降低写入压力
if inserted % commit_every < CHUNK:
dst.commit() # 攒批 commit:显著减少 fsync 次数
if inserted // PROGRESS_EVERY != (inserted - len(rows)) // PROGRESS_EVERY:
el = time.time() - t0
rate = inserted / el if el else 0
print(f"[{table}] {inserted:,}/{src_total:,} 行 ({el:.0f}s, {rate:,.0f} 行/s)",
flush=True)
dst.commit()
el = time.time() - t0
print(f"[{table}] 完成 {inserted:,} 行,耗时 {el:.0f}s", flush=True)
return inserted
finally:
cur.close()
def _insert_chunk(cur, table: str, col_q: str, cols: list[str], rows) -> None:
n = len(rows)
n_col = len(cols)
per = ", ".join(["%s"] * n_col)
values = ", ".join([f"({per})"] * n)
sql = f"INSERT INTO `{table}` ({col_q}) VALUES {values}"
flat = [v for row in rows for v in row]
cur.execute(sql, flat)
def count(dst, table: str) -> int:
cur = dst.cursor()
try:
cur.execute(f"SELECT COUNT(*) FROM `{table}`")
return int(cur.fetchone()[0])
finally:
cur.close()
def _verify_eq(a, b) -> bool:
"""抽样比对:None 严格相等;日期对象与 ISO 字符串互比;数值近似
(sqlite float vs mysql Decimal 舍入差);其余直等。"""
if a is None or b is None:
return a is b
# 数值(sqlite float / mysql Decimal / int):按 6 位小数容差比较
if isinstance(a, (int, float, Decimal)) and isinstance(b, (int, float, Decimal)):
return round(float(a), 6) == round(float(b), 6)
if isinstance(a, str) and hasattr(b, "isoformat"):
return a == b.isoformat()
if isinstance(b, str) and hasattr(a, "isoformat"):
return a.isoformat() == b
return a == b
def verify(src: sqlite3.Connection, dst) -> bool:
ok = True
print("---- 行数校验 ----", flush=True)
for table, (_, _pk) in TABLES.items():
s = src.execute(f'SELECT COUNT(*) FROM "{table}"').fetchone()[0]
d = count(dst, table)
flag = "OK" if s == d else "MISMATCH"
ok = ok and s == d
print(f" {table:<22} sqlite={s:>12,} mysql={d:>12,} {flag}", flush=True)
print("---- 抽样比对 ----", flush=True)
cur = dst.cursor()
try:
cases = [
# 注意:sqlite/mysql 双兼容写法 —— 表名裸名 + 字符串一律单引号
# (mysql 版仅把表名包反引号,见下方 replace("FROM ", "FROM `") 思路外实现)
("stock_daily",
"SELECT symbol, trade_date, open, high, low, close, volume, amount, source, adjust "
'FROM "stock_daily" WHERE symbol IN (\'600519.SH\',\'000001.SZ\') '
"ORDER BY trade_date, symbol LIMIT 300"),
("adjust_factor",
'SELECT symbol, trade_date, factor FROM "adjust_factor" WHERE symbol = \'600519.SH\' '
"ORDER BY trade_date, symbol LIMIT 300"),
("financial_indicator",
'SELECT symbol, report_date, announce_date, eps, roe, total_revenue, net_profit '
'FROM "financial_indicator" WHERE symbol = \'600519.SH\' '
"ORDER BY report_date, symbol LIMIT 300"),
("stock",
'SELECT symbol, name, industry, market, list_date, status FROM "stock" '
"WHERE symbol IN ('600519.SH','000001.SZ','300750.SZ')"),
]
for table, ssql in cases:
s_rows = src.execute(ssql).fetchall()
# mysql 版:双引号仅用于表名(字符串已是单引号),安全替换为反引号
m_sql = ssql.replace('"', "`")
cur.execute(m_sql)
m_rows = cur.fetchall()
bad = 0
for a, b in zip(s_rows, m_rows, strict=False):
if any(not _verify_eq(x, y) for x, y in zip(a, b, strict=False)):
bad += 1
same_len = len(s_rows) == len(m_rows)
flag = "OK" if bad == 0 and same_len else "MISMATCH"
ok = ok and bad == 0 and same_len
print(f" {table:<22} sqlite={len(s_rows):>4} mysql={len(m_rows):>4} 不一致行={bad} {flag}",
flush=True)
finally:
cur.close()
return ok
def main() -> int:
ap = argparse.ArgumentParser(description="SQLite → MySQL 数据迁移")
ap.add_argument("--source", default=str(DEFAULT_SOURCE))
ap.add_argument("--reset", action="store_true", help="目标表已有数据时 TRUNCATE 重传")
ap.add_argument("--verify-only", action="store_true", help="仅校验不写入")
ap.add_argument("--only", nargs="+", default=None,
help="只迁移指定表(可多个);并行迁移大表时使用(自动跳过末尾校验)")
ap.add_argument("--throttle-sec", type=float, default=0.0,
help="每批(5000 行)之间的休眠秒数;低配 MySQL/共享服务器请设 0.3~1.0 降压力")
args = ap.parse_args()
tables = {k: v for k, v in TABLES.items() if args.only is None or k in args.only}
if not tables:
print(f"--only 未匹配任何表,可用: {', '.join(TABLES)}", file=sys.stderr)
return 2
src_path = Path(args.source)
if not src_path.exists():
print(f"源 sqlite 不存在: {src_path}", file=sys.stderr)
return 2
print(f"源 : {src_path} ({src_path.stat().st_size / 1e9:.2f} GB)", flush=True)
print(f"目标库: {resolve_dst_engine().url.database}", flush=True)
src = sqlite3.connect(str(src_path))
engine = resolve_dst_engine()
dst = engine.raw_connection()
try:
if args.verify_only:
ok = verify(src, dst)
return 0 if ok else 1
for table, (cols, pk_kind) in tables.items():
migrate_table(src, dst, table, cols, pk_kind, reset=args.reset,
throttle_sec=args.throttle_sec)
if args.only:
return 0 # 并行分路:由单独 --verify-only 统一校验
ok = verify(src, dst)
print("迁移完成:", "全部一致 ✓" if ok else "存在不一致,请人工核对 ✗")
return 0 if ok else 1
finally:
src.close()
dst.close()
engine.dispose()
if __name__ == "__main__":
raise SystemExit(main())