Compare commits
21
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
314bfc159f | ||
|
|
db147c4232 | ||
|
|
e5a23f176d | ||
|
|
e1ac23fa25 | ||
|
|
8b2f8ac35c | ||
|
|
9d25d466e5 | ||
|
|
692bdb3be5 | ||
|
|
ba52edc2d6 | ||
|
|
ef09d5b419 | ||
|
|
4fa2bb748e | ||
|
|
273aee2772 | ||
|
|
8f47b5b603 | ||
|
|
0ab9038570 | ||
|
|
0d3e123de3 | ||
|
|
c60dc78c88 | ||
|
|
75c5472c31 | ||
|
|
25a1d9531a | ||
|
|
f3586adb25 | ||
|
|
697ffc767b | ||
|
|
0ffd574f30 | ||
|
|
6c2f198261 |
+7
-3
@@ -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 中配置
|
||||||
|
|||||||
@@ -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 模型选股按需延后)
|
||||||
|
|
||||||
## 约定速查
|
## 约定速查
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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}
|
||||||
@@ -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
@@ -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()
|
|
||||||
]
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -0,0 +1,65 @@
|
|||||||
|
"""交易信号 API(M8.1):提交/查询信号(落库可复现)。
|
||||||
|
|
||||||
|
POST /api/signals body: {query: SelectionQuery, rules?: SignalRules}
|
||||||
|
GET /api/signals/{id} 读回某次信号
|
||||||
|
GET /api/signals 历史信号元数据(可过滤 as_of)
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import date
|
||||||
|
from typing import Annotated
|
||||||
|
|
||||||
|
from fastapi import APIRouter, HTTPException, Query
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
from app.api.deps import DbSession, SignalRepoDep, SignalServiceDep
|
||||||
|
from app.application.services.job_executor import new_id
|
||||||
|
from app.domain.entities.selection import SelectionQuery
|
||||||
|
from app.domain.entities.signal import SignalMeta, SignalResult, SignalRules
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/signals", tags=["signals"])
|
||||||
|
|
||||||
|
_AsOfQuery = Annotated[date | None, Query(description="按信号时点过滤")]
|
||||||
|
_LimitQuery = Annotated[int, Query(ge=1, le=200)]
|
||||||
|
|
||||||
|
|
||||||
|
class SignalRequest(BaseModel):
|
||||||
|
query: SelectionQuery
|
||||||
|
rules: SignalRules = SignalRules()
|
||||||
|
|
||||||
|
|
||||||
|
class SignalRun(BaseModel):
|
||||||
|
signal_id: str
|
||||||
|
result: SignalResult
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("", response_model=SignalRun, summary="生成一次交易信号(同步)并落库")
|
||||||
|
def run_signal(
|
||||||
|
req: SignalRequest,
|
||||||
|
service: SignalServiceDep,
|
||||||
|
signal_repo: SignalRepoDep,
|
||||||
|
session: DbSession,
|
||||||
|
) -> SignalRun:
|
||||||
|
result = service.signal(req.query, req.rules)
|
||||||
|
signal_id = new_id("SIG")
|
||||||
|
signal_repo.save(signal_id, result)
|
||||||
|
session.commit()
|
||||||
|
return SignalRun(signal_id=signal_id, result=result)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{signal_id}", response_model=SignalResult, summary="读回一次信号结果")
|
||||||
|
def get_signal(signal_id: str, signal_repo: SignalRepoDep) -> SignalResult:
|
||||||
|
result = signal_repo.get(signal_id)
|
||||||
|
if result is None:
|
||||||
|
raise HTTPException(status_code=404, detail=f"信号记录 {signal_id} 不存在")
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("", response_model=list[SignalMeta], summary="历史信号元数据列表")
|
||||||
|
def list_signals(
|
||||||
|
signal_repo: SignalRepoDep,
|
||||||
|
as_of: _AsOfQuery = None,
|
||||||
|
limit: _LimitQuery = 20,
|
||||||
|
) -> list[SignalMeta]:
|
||||||
|
return signal_repo.list_recent(as_of=as_of, limit=limit)
|
||||||
@@ -0,0 +1,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)
|
||||||
@@ -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"
|
||||||
|
|||||||
@@ -0,0 +1,33 @@
|
|||||||
|
"""因子组合(Composite Factor)领域实体(M7.2b,v2 §13)。
|
||||||
|
|
||||||
|
组合 = 一组 {因子, 权重} + method(MVP fixed:截面 zscore×方向×权重求和;
|
||||||
|
方向在计算时取因子注册表元数据,落库时冗余快照以便列表展示)。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field, model_validator
|
||||||
|
|
||||||
|
|
||||||
|
class CompositeComponent(BaseModel):
|
||||||
|
name: str
|
||||||
|
weight: float = Field(default=1.0, gt=0)
|
||||||
|
direction: str = Field(default="higher_is_better")
|
||||||
|
|
||||||
|
|
||||||
|
class CompositeDefinition(BaseModel):
|
||||||
|
id: str = ""
|
||||||
|
name: str = Field(min_length=1, max_length=64)
|
||||||
|
method: str = Field(default="fixed", pattern="^(fixed)$")
|
||||||
|
description: str = ""
|
||||||
|
components: list[CompositeComponent] = Field(min_length=1)
|
||||||
|
created_at: datetime | None = None
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def _no_duplicate(self) -> CompositeDefinition:
|
||||||
|
names = [c.name for c in self.components]
|
||||||
|
if len(set(names)) != len(names):
|
||||||
|
raise ValueError("components 存在重复因子名")
|
||||||
|
return self
|
||||||
@@ -0,0 +1,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),
|
||||||
|
)
|
||||||
@@ -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")
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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。"""
|
||||||
@@ -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: ...
|
||||||
@@ -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 过滤)。"""
|
||||||
@@ -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: ...
|
||||||
+62
@@ -0,0 +1,62 @@
|
|||||||
|
"""selection_snapshot / selection_result 表(M6.3 选股落库)
|
||||||
|
|
||||||
|
Revision ID: a6c91d4e7f20
|
||||||
|
Revises: d3f6c9a21b04
|
||||||
|
Create Date: 2026-09-08
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision: str = "a6c91d4e7f20"
|
||||||
|
down_revision: str | None = "d3f6c9a21b04"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.create_table(
|
||||||
|
"selection_snapshot",
|
||||||
|
sa.Column("id", sa.String(length=32), nullable=False),
|
||||||
|
sa.Column("as_of", sa.Date(), nullable=False),
|
||||||
|
sa.Column("method", sa.String(length=16), nullable=False),
|
||||||
|
sa.Column("query_json", sa.Text(), nullable=False),
|
||||||
|
sa.Column("statistics_json", sa.Text(), nullable=False),
|
||||||
|
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
|
op.create_index("ix_selection_snapshot_as_of", "selection_snapshot", ["as_of"])
|
||||||
|
op.create_index("ix_selection_snapshot_created_at", "selection_snapshot", ["created_at"])
|
||||||
|
|
||||||
|
op.create_table(
|
||||||
|
"selection_result",
|
||||||
|
sa.Column(
|
||||||
|
"id",
|
||||||
|
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||||
|
autoincrement=True,
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.Column("selection_id", sa.String(length=32), nullable=False),
|
||||||
|
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||||
|
sa.Column("rank", sa.Integer(), nullable=False),
|
||||||
|
sa.Column("score", sa.Numeric(precision=14, scale=6), nullable=False),
|
||||||
|
sa.Column("factor_values_json", sa.Text(), nullable=True),
|
||||||
|
sa.Column("filter_status_json", sa.Text(), nullable=True),
|
||||||
|
sa.Column("reason_json", sa.Text(), nullable=True),
|
||||||
|
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
|
op.create_index("ix_selection_result_selection_id", "selection_result", ["selection_id"])
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_index("ix_selection_result_selection_id", table_name="selection_result")
|
||||||
|
op.drop_table("selection_result")
|
||||||
|
op.drop_index("ix_selection_snapshot_created_at", table_name="selection_snapshot")
|
||||||
|
op.drop_index("ix_selection_snapshot_as_of", table_name="selection_snapshot")
|
||||||
|
op.drop_table("selection_snapshot")
|
||||||
+40
@@ -0,0 +1,40 @@
|
|||||||
|
"""factor_definition 表(M7.1 因子定义入库)
|
||||||
|
|
||||||
|
Revision ID: b7f2a5e81c33
|
||||||
|
Revises: a6c91d4e7f20
|
||||||
|
Create Date: 2026-09-09
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision: str = "b7f2a5e81c33"
|
||||||
|
down_revision: str | None = "a6c91d4e7f20"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.create_table(
|
||||||
|
"factor_definition",
|
||||||
|
sa.Column("name", sa.String(length=64), nullable=False),
|
||||||
|
sa.Column("description", sa.String(length=500), nullable=False),
|
||||||
|
sa.Column("formula", sa.String(length=500), nullable=False),
|
||||||
|
sa.Column("brief", sa.String(length=500), nullable=False),
|
||||||
|
sa.Column("frequency", sa.String(length=16), nullable=False),
|
||||||
|
sa.Column("lookback", sa.Integer(), nullable=False),
|
||||||
|
sa.Column("direction", sa.String(length=32), nullable=False),
|
||||||
|
sa.Column("requires_json", sa.Text(), nullable=False),
|
||||||
|
sa.Column("version", sa.String(length=16), nullable=False),
|
||||||
|
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||||
|
sa.PrimaryKeyConstraint("name"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_table("factor_definition")
|
||||||
+37
@@ -0,0 +1,37 @@
|
|||||||
|
"""factor_composite 表(M7.2b 因子组合保存/复用)
|
||||||
|
|
||||||
|
Revision ID: c3e9a0d1f4b5
|
||||||
|
Revises: b7f2a5e81c33
|
||||||
|
Create Date: 2026-09-09
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision: str = "c3e9a0d1f4b5"
|
||||||
|
down_revision: str | None = "b7f2a5e81c33"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.create_table(
|
||||||
|
"factor_composite",
|
||||||
|
sa.Column("id", sa.String(length=32), nullable=False),
|
||||||
|
sa.Column("name", sa.String(length=64), nullable=False),
|
||||||
|
sa.Column("method", sa.String(length=16), nullable=False),
|
||||||
|
sa.Column("description", sa.String(length=300), nullable=False),
|
||||||
|
sa.Column("components_json", sa.Text(), nullable=False),
|
||||||
|
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
sa.UniqueConstraint("name", name="uq_factor_composite_name"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_table("factor_composite")
|
||||||
+62
@@ -0,0 +1,62 @@
|
|||||||
|
"""signal_snapshot / signal_event 表(M8.1 交易信号)
|
||||||
|
|
||||||
|
Revision ID: d8e0b2f3c4d5
|
||||||
|
Revises: c3e9a0d1f4b5
|
||||||
|
Create Date: 2026-09-09
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision: str = "d8e0b2f3c4d5"
|
||||||
|
down_revision: str | None = "c3e9a0d1f4b5"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.create_table(
|
||||||
|
"signal_snapshot",
|
||||||
|
sa.Column("id", sa.String(length=32), nullable=False),
|
||||||
|
sa.Column("as_of", sa.Date(), nullable=False),
|
||||||
|
sa.Column("query_json", sa.Text(), nullable=False),
|
||||||
|
sa.Column("rules_json", sa.Text(), nullable=False),
|
||||||
|
sa.Column("statistics_json", sa.Text(), nullable=False),
|
||||||
|
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
|
op.create_index("ix_signal_snapshot_as_of", "signal_snapshot", ["as_of"])
|
||||||
|
op.create_index("ix_signal_snapshot_created_at", "signal_snapshot", ["created_at"])
|
||||||
|
|
||||||
|
op.create_table(
|
||||||
|
"signal_event",
|
||||||
|
sa.Column(
|
||||||
|
"id",
|
||||||
|
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||||
|
autoincrement=True,
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.Column("signal_id", sa.String(length=32), nullable=False),
|
||||||
|
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||||
|
sa.Column("signal_date", sa.Date(), nullable=False),
|
||||||
|
sa.Column("signal_type", sa.String(length=8), nullable=False),
|
||||||
|
sa.Column("score", sa.Numeric(precision=14, scale=6), nullable=True),
|
||||||
|
sa.Column("price", sa.Numeric(precision=14, scale=4), nullable=True),
|
||||||
|
sa.Column("reason_json", sa.Text(), nullable=True),
|
||||||
|
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
|
op.create_index("ix_signal_event_signal_id", "signal_event", ["signal_id"])
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_index("ix_signal_event_signal_id", table_name="signal_event")
|
||||||
|
op.drop_table("signal_event")
|
||||||
|
op.drop_index("ix_signal_snapshot_created_at", table_name="signal_snapshot")
|
||||||
|
op.drop_index("ix_signal_snapshot_as_of", table_name="signal_snapshot")
|
||||||
|
op.drop_table("signal_snapshot")
|
||||||
+38
@@ -0,0 +1,38 @@
|
|||||||
|
"""strategy 表(M8.3 策略持久化)
|
||||||
|
|
||||||
|
Revision ID: e1f2a3b4c5d6
|
||||||
|
Revises: d8e0b2f3c4d5
|
||||||
|
Create Date: 2026-09-09
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision: str = "e1f2a3b4c5d6"
|
||||||
|
down_revision: str | None = "d8e0b2f3c4d5"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.create_table(
|
||||||
|
"strategy",
|
||||||
|
sa.Column("id", sa.String(length=32), nullable=False),
|
||||||
|
sa.Column("name", sa.String(length=64), nullable=False),
|
||||||
|
sa.Column("description", sa.String(length=300), nullable=False),
|
||||||
|
sa.Column("spec_type", sa.String(length=16), nullable=False),
|
||||||
|
sa.Column("config_json", sa.Text(), nullable=False),
|
||||||
|
sa.Column("version", sa.String(length=16), nullable=False),
|
||||||
|
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
sa.UniqueConstraint("name", name="uq_strategy_name"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_table("strategy")
|
||||||
@@ -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,
|
||||||
|
)
|
||||||
@@ -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))
|
||||||
@@ -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()
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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)
|
|
||||||
|
|||||||
@@ -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")},
|
||||||
|
)
|
||||||
@@ -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
|
||||||
@@ -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]
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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()
|
|
||||||
|
|||||||
@@ -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()}
|
||||||
@@ -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):
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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 → 全部不通过
|
||||||
@@ -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"])
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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)
|
||||||
Generated
+11
@@ -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
@@ -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"
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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 已隔离方言差异)。
|
||||||
|
|||||||
@@ -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} />
|
||||||
</>
|
</>
|
||||||
);
|
);
|
||||||
|
|||||||
@@ -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>
|
||||||
|
|||||||
@@ -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>
|
||||||
|
);
|
||||||
|
}
|
||||||
@@ -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>
|
||||||
|
);
|
||||||
|
}
|
||||||
@@ -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
@@ -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;
|
||||||
|
}
|
||||||
|
|||||||
@@ -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())
|
||||||
Reference in New Issue
Block a user