feat: 股息率案例口径 + 策略库与图表统一 + 回测存档完整化

汇总三轮未提交的开发(每轮均在本机 MariaDB + 真实浏览器上验证):

1) 股息率案例(全市场股息率最高 n 只,默认 20,每 m 月择股)
   - 新增日频估值表 daily_basic + 迁移;股息率因子(dv_ratio / dividend_yield / TTM)
   - 名称历史表 stock_name_history:剔除 ST 按**择股日当时名称**判定,消除
     「曾高股息后 ST」的股息陷阱(实测 3.70pp 偏差)
   - 区间择股/调仓双周期(m 择股 / y 调仓)、指数成分与白名单、停牌近似剔除
   - 复权因子口径核对(4,164,742 行、缺失 0.0%)、收盘价成交与涨跌停拦单
   - 案例实测:2020-01-01~2026-09-04 总收益 +24.86%(年化 3.52%、回撤 -28.58%)

2) 策略库与前端统一
   - strategy 表 + CRUD/PUT 原地更新 + `describe_strategy` 按 spec 真实推导
     「一句话说明 + 计算公式 + 执行步骤 + 注意事项」(与引擎实执行规则同源)
   - 任何出现股票代码处都成对显示名称且可点击进个股页
   - 全站图表基座统一 TradingView Lightweight Charts(ECharts 依赖、
     锁文件、组件与文档标注一并清除),买卖点标记只落在真实交易日上

3) 回测存档完整化(可往复查看)
   - 同步端点(POST /api/backtests、/api/factor-tests)此前完全不落库 → 现在同样归档,
     归档 id 经响应头 X-Experiment-Id 返回(不破坏 response_model)
   - data_version 首次真实写入(数据快照指纹:最新交易日 + 各表规模)
   - 个股收益曲线默认**全量保存**(此前硬截断 60 只);超出体积预算才裁剪,
     并写 archive_meta(机器可读)+ unimplemented(人可读)如实标注
   - 列表 kind/q 过滤 + X-Total-Count(此前 limit=50 静默截断)、DELETE 归档
   - 只读归档页 /experiments/{id}(Server Component,SSR 直出**选股条件**与
     **交易执行依据**);结果视图按 kind 分发(backtest/factor_test/selection),
     非回测归档不套用回测口径
   - 新增 CLI:prune_experiments(保留策略,默认 dry-run)、
     restore_experiment_from_job(从 Job 副本按原 id 重建被删的历史归档,默认 dry-run)

门禁:pytest 388 passed、ruff All checks passed、tsc 0 错误、图表单测 7 passed、
next build 成功、契约脚本 verify_strategy_workspace 59/59(含按 kind 逐类验证归档页)。
This commit is contained in:
Simon
2026-09-20 07:31:04 +08:00
parent 7e15b7251e
commit 23972e7063
112 changed files with 17908 additions and 3893 deletions
@@ -0,0 +1,129 @@
"""数据集快照指纹:`experiment.data_version` 的真实取值来源。
归档的用途是「往复查看 / 复现」,因此除了代码版本(code_version)还必须记录
**当时的数据口径**:数据同步到哪一天、各表规模多大。此前该字段从未被写入
(永远 NULL),本模块补上。
诚实性与成本(AGENT.md §7 不静默 / §24 不假装支持):
- 交易日 `MAX(trade_date)` 走索引,实测 ~0.2ms,**真实值**;
- 大表(stock_daily / adjust_factor / daily_basic,各 800 万行量级)的全表
`COUNT(*)` 实测单次 ~1.2s,三次合计 ~3.6s —— 不允许出现在请求路径上;
故 MySQL 下改用 `information_schema.TABLES.TABLE_ROWS`(单次查询实测 ~0.5ms),
它是 InnoDB 统计缓存的**近似值**(实测 stock_daily 7688126 vs 真实 COUNT(*)
8052698,偏差 ~4.6%),因此字符串里用 `≈` 明确标注为近似,绝不冒充精确计数。
- 非 MySQL 方言(测试用的 SQLite 等)数据量小,直接 `COUNT(*)` 得到精确值,
不带 `≈` 标记;
- 任何一段取不到(连接失败 / 表不存在 / 权限不足)都**降级**:要么丢弃该段,
要么整体返回 `unavailable`,绝不编造数字。
字符串格式(≤40 字符 —— experiment.data_version 列是 varchar(40),本文件不新增
迁移,故必须塞得下;段超长时从右往左丢弃低优先级段,再退化到 `d<date>`):
d<YYYYMMDD> stock_daily 的最大交易日(`d-` 表示该表无数据 / 取不到交易日)
n<count> stock_daily 行数
a<count> adjust_factor 行数
b<count> daily_basic 行数
行数段写法:`≈<N>k` = MySQL 近似值(k = 千行,四舍五入);`<N>` = 精确值。
示例:`d20260904;n≈8053k;a≈8189k;b≈7718k`
"""
from __future__ import annotations
import logging
from sqlalchemy import func, select, text
from sqlalchemy.orm import Session
from app.infrastructure.persistence.sqlalchemy.models.market import (
AdjustFactorModel,
DailyBasicModel,
StockDailyModel,
)
logger = logging.getLogger(__name__)
# experiment.data_version 列宽(varchar(40),见 models/jobs.py + 既有迁移)
DATA_VERSION_MAX_LEN = 40
# 段定义:(段名, model),顺序即优先级(越靠前越重要)
_TABLES: tuple[tuple[str, type], ...] = (
("n", StockDailyModel),
("a", AdjustFactorModel),
("b", DailyBasicModel),
)
_FALLBACK = "unavailable"
def _fmt_rows(rows: int | None, *, approx: bool) -> str | None:
"""行数段:近似值用「≈Nk」(千行),精确值用原样数字。"""
if rows is None:
return None
if not approx:
return str(rows)
return f"≈{round(rows / 1000)}k"
def _latest_trade_date(session: Session) -> str:
"""stock_daily 最大交易日 → `YYYYMMDD`;无数据 / 取不到 → `-`。"""
try:
day = session.execute(select(func.max(StockDailyModel.trade_date))).scalar()
except Exception as exc: # noqa: BLE001 —— 指纹取不到必须降级,不影响归档主体
logger.warning("data_version: MAX(trade_date) 取不到:%s: %s", type(exc).__name__, exc)
return "-"
return day.strftime("%Y%m%d") if day is not None else "-"
def _approx_rows_mysql(session: Session) -> dict[str, int]:
"""MySQL:一次 information_schema 查询拿全部表的近似行数(实测 ~0.5ms)。"""
try:
rows = session.execute(
text(
"SELECT TABLE_NAME, TABLE_ROWS FROM information_schema.TABLES "
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME IN :names"
).bindparams(names=tuple(m.__tablename__ for _, m in _TABLES))
).all()
except Exception as exc: # noqa: BLE001
logger.warning("data_version: information_schema 查询失败:%s: %s", type(exc).__name__, exc)
return {}
return {str(name): int(n) for name, n in rows if n is not None}
def _exact_rows(session: Session, model: type) -> int | None:
"""其它方言(SQLite 等,数据量小):精确 COUNT(*)。"""
try:
return int(session.execute(select(func.count()).select_from(model)).scalar() or 0)
except Exception as exc: # noqa: BLE001
logger.warning("data_version: COUNT(*) 失败:%s: %s", type(exc).__name__, exc)
return None
def compute_data_version(session: Session) -> str:
"""计算数据集快照指纹(格式与取舍见模块 docstring)。失败降级为 `unavailable`。"""
try:
dialect = session.get_bind().dialect.name
except Exception: # noqa: BLE001
dialect = ""
approx = dialect == "mysql"
approx_rows = _approx_rows_mysql(session) if approx else {}
# 首段恒为交易日段:`d-` 本身也是真实信息(表为空 / 取不到交易日),保留
segments: list[str] = [f"d{_latest_trade_date(session)}"]
for name, model in _TABLES:
rows = (
approx_rows.get(model.__tablename__) if approx else _exact_rows(session, model)
)
seg = _fmt_rows(rows, approx=approx)
if seg is not None:
segments.append(f"{name}{seg}")
# 列宽护栏:超 40 字符时从右往左丢弃低优先级段(保留的段仍是真实值,绝不截断数字)
while len(segments) > 1 and len(";".join(segments)) > DATA_VERSION_MAX_LEN:
segments.pop()
if len(segments) == 1 and segments[0] == "d-":
# 三张表连行数都读不到、交易日也没有:如实标记「不可用」,不编造
return _FALLBACK
return ";".join(segments)
@@ -10,15 +10,20 @@ from app.infrastructure.persistence.sqlalchemy.models.composite import ( # noqa
from app.infrastructure.persistence.sqlalchemy.models.factor import ( # noqa: F401
FactorDefinitionModel,
)
from app.infrastructure.persistence.sqlalchemy.models.index import ( # noqa: F401
IndexWeightModel,
)
from app.infrastructure.persistence.sqlalchemy.models.jobs import ( # noqa: F401
ExperimentModel,
JobModel,
)
from app.infrastructure.persistence.sqlalchemy.models.market import ( # noqa: F401
AdjustFactorModel,
DailyBasicModel,
FinancialIndicatorModel,
StockDailyModel,
StockModel,
StockNameHistoryModel,
SyncLogModel,
TradingCalendarModel,
)
@@ -1,14 +1,25 @@
"""Phase 4:Job(异步任务)与 Experiment(实验归档)表。"""
"""Phase 4:Job(异步任务)与 Experiment(实验归档)表。
`result_json` 使用 MEDIUMTEXT(MySQL 上限 16MB):全市场多年回测的结果含
净值/回撤曲线、逐笔成交、Signal↔Fill 记录与个股收益曲线,实测可达数 MB,
MySQL `TEXT`(64KB)会直接报 1406 Data too long 导致 Job 归档失败
(2026-09 高股息案例实测:约 1.2MB → 落库失败)。
SQLite 不区分 TEXT 长度,故模型层统一用 MEDIUMTEXT.with_variant 保持跨库可用。
"""
from __future__ import annotations
from datetime import datetime
from sqlalchemy import DateTime, String, Text
from sqlalchemy.dialects.mysql import MEDIUMTEXT
from sqlalchemy.orm import Mapped, mapped_column
from app.infrastructure.persistence.sqlalchemy.base import Base
# 长 JSON 列类型:MySQL 用 MEDIUMTEXT(16MB),其它方言退化为 TEXT(SQLite 无长度限制)
_LONG_JSON = Text().with_variant(MEDIUMTEXT(), "mysql")
class JobModel(Base):
__tablename__ = "job"
@@ -19,7 +30,7 @@ class JobModel(Base):
stage: Mapped[str | None] = mapped_column(String(24), nullable=True)
spec_json: Mapped[str] = mapped_column(Text)
error: Mapped[str | None] = mapped_column(Text, nullable=True)
result_json: Mapped[str | None] = mapped_column(Text, nullable=True)
result_json: Mapped[str | None] = mapped_column(_LONG_JSON, nullable=True)
experiment_id: Mapped[str | None] = mapped_column(String(32), nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime)
started_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
@@ -32,7 +43,7 @@ class ExperimentModel(Base):
id: Mapped[str] = mapped_column(String(32), primary_key=True)
kind: Mapped[str] = mapped_column(String(16))
spec_json: Mapped[str] = mapped_column(Text)
result_json: Mapped[str] = mapped_column(Text)
result_json: Mapped[str] = mapped_column(_LONG_JSON)
summary_text: Mapped[str | None] = mapped_column(String(200), nullable=True)
code_version: Mapped[str | None] = mapped_column(String(40), nullable=True)
data_version: Mapped[str | None] = mapped_column(String(40), nullable=True)
@@ -81,6 +81,37 @@ class AdjustFactorModel(Base):
factor: Mapped[Decimal] = mapped_column(Numeric(20, 6))
class DailyBasicModel(Base):
"""每日指标快照(Tushare daily_basic)—— 估值 / 股息率 / 市值。
幂等键 (symbol, trade_date):同一交易日同一股票唯一一行。
dv_ratio/dv_ttm 为时点值,研究侧按 trade_date <= as_of 取用(无未来函数)。
"""
__tablename__ = "daily_basic"
__table_args__ = (UniqueConstraint("symbol", "trade_date", name="uq_basic_symbol_date"),)
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), index=True)
trade_date: Mapped[date] = mapped_column(Date, index=True)
source: Mapped[str] = mapped_column(String(16), default="tushare", server_default="tushare")
close: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
turnover_rate: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
volume_ratio: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
pe: Mapped[Decimal | None] = mapped_column(Numeric(16, 4), nullable=True)
pe_ttm: Mapped[Decimal | None] = mapped_column(Numeric(16, 4), nullable=True)
pb: Mapped[Decimal | None] = mapped_column(Numeric(16, 4), nullable=True)
ps: Mapped[Decimal | None] = mapped_column(Numeric(16, 4), nullable=True)
ps_ttm: Mapped[Decimal | None] = mapped_column(Numeric(16, 4), nullable=True)
dv_ratio: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
dv_ttm: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
total_share: Mapped[Decimal | None] = mapped_column(Numeric(24, 4), nullable=True)
float_share: Mapped[Decimal | None] = mapped_column(Numeric(24, 4), nullable=True)
free_share: Mapped[Decimal | None] = mapped_column(Numeric(24, 4), nullable=True)
total_mv: Mapped[Decimal | None] = mapped_column(Numeric(24, 4), nullable=True)
circ_mv: Mapped[Decimal | None] = mapped_column(Numeric(24, 4), nullable=True)
class FinancialIndicatorModel(Base):
"""财务指标快照 —— report_date(报告期) 与 announce_date(公告日) 并存。"""
@@ -115,3 +146,26 @@ class SyncLogModel(Base):
row_count: Mapped[int] = mapped_column(default=0)
data_start: Mapped[date | None] = mapped_column(Date, nullable=True)
data_end: Mapped[date | None] = mapped_column(Date, nullable=True)
class StockNameHistoryModel(Base):
"""股票名称变更历史(Tushare namechange)—— 时点 ST / 风险警示判定的依据。
幂等键 (symbol, start_date):同一股票同一名称生效起点唯一一行。
查询语义:`name` 在 [start_date, end_date] 内有效;`end_date` 为空表示至今有效。
时点取值:`start_date <= as_of AND (end_date IS NULL OR end_date >= as_of)`。
"""
__tablename__ = "stock_name_history"
__table_args__ = (
UniqueConstraint("symbol", "start_date", name="uq_name_symbol_start"),
)
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), index=True)
name: Mapped[str] = mapped_column(String(64))
start_date: Mapped[date] = mapped_column(Date, index=True)
end_date: Mapped[date | None] = mapped_column(Date, nullable=True, index=True)
ann_date: Mapped[date | None] = mapped_column(Date, nullable=True)
change_reason: Mapped[str | None] = mapped_column(String(32), nullable=True)
source: Mapped[str] = mapped_column(String(16), default="tushare", server_default="tushare")
@@ -2,10 +2,10 @@
from __future__ import annotations
from sqlalchemy import select
from sqlalchemy import Select, func, select
from sqlalchemy.orm import Session
from app.domain.entities.research import ExperimentRecord, JobRecord
from app.domain.entities.research import ExperimentRecord, ExperimentSummary, JobRecord
from app.infrastructure.persistence.sqlalchemy.models.jobs import ExperimentModel, JobModel
@@ -53,6 +53,12 @@ class SqlAlchemyJobRepository:
class SqlAlchemyExperimentRepository:
"""Experiment 归档仓储。
`list_filtered` 刻意**不加载 `result_json`**(MEDIUMTEXT,完整存档后单条数 MB),
只 SELECT 元数据列 + SQL 侧算出的字符长度,避免列表接口把上百 MB 拉进内存。
"""
def __init__(self, session: Session) -> None:
self._session = session
@@ -61,6 +67,12 @@ class SqlAlchemyExperimentRepository:
self._session.flush()
return experiment
def upsert(self, experiment: ExperimentRecord) -> ExperimentRecord:
"""按 id 插入或覆盖(`merge` 语义):重建归档时保持原 id 不变。"""
self._session.merge(ExperimentModel(**experiment.model_dump()))
self._session.flush()
return experiment
def get(self, experiment_id: str) -> ExperimentRecord | None:
row = self._session.get(ExperimentModel, experiment_id)
return ExperimentRecord.model_validate(row, from_attributes=True) if row else None
@@ -70,3 +82,98 @@ class SqlAlchemyExperimentRepository:
select(ExperimentModel).order_by(ExperimentModel.created_at.desc()).limit(limit)
).all()
return [ExperimentRecord.model_validate(r, from_attributes=True) for r in rows]
# -- 列表 / 过滤(SQL 层完成过滤+分页,参数绑定,不做字符串拼接)-------------
@staticmethod
def _escape_like(value: str) -> str:
"""转义 LIKE 元字符:用户输入里的 % / _ / \\ 不应被当成通配符。"""
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
def _filter_conditions(self, *, kind: str | None, q: str | None):
"""kind 精确匹配 + q 大小写不敏感模糊匹配(id / spec_json=因子名 / summary_text)。"""
conditions = []
if kind:
conditions.append(ExperimentModel.kind == kind)
if q:
pattern = f"%{self._escape_like(q.lower())}%"
conditions.append(
func.lower(ExperimentModel.id).like(pattern, escape="\\")
| func.lower(ExperimentModel.spec_json).like(pattern, escape="\\")
| func.lower(func.coalesce(ExperimentModel.summary_text, "")).like(
pattern, escape="\\"
)
)
return conditions
def _summary_stmt(self, conditions) -> Select:
return select(
ExperimentModel.id,
ExperimentModel.kind,
ExperimentModel.spec_json,
ExperimentModel.summary_text,
ExperimentModel.code_version,
ExperimentModel.data_version,
ExperimentModel.job_id,
ExperimentModel.created_at,
self._char_length().label("result_bytes"),
).where(*conditions)
def _char_length(self):
"""归档 JSON 的**字符数**(不取大字段本身)。
MySQL 的 `LENGTH()` 返回字节数,字符数要用 `CHAR_LENGTH()`;SQLite 只有
`length()`(对 TEXT 返回字符数)。按方言选择,避免中文归档在两个库上口径
不一致。
"""
if self._session.get_bind().dialect.name == "sqlite":
return func.length(ExperimentModel.result_json)
return func.char_length(ExperimentModel.result_json)
def list_filtered(
self,
*,
kind: str | None = None,
q: str | None = None,
limit: int = 200,
offset: int = 0,
) -> list[ExperimentSummary]:
stmt = (
self._summary_stmt(self._filter_conditions(kind=kind, q=q))
# created_at 在 MySQL datetime(0) 下只有秒精度,同秒创建的归档必须有
# 稳定的次级排序键,否则分页会重复/漏项
.order_by(ExperimentModel.created_at.desc(), ExperimentModel.id.desc())
.limit(max(limit, 0))
.offset(max(offset, 0))
)
return [
ExperimentSummary(
id=row.id,
kind=row.kind,
spec_json=row.spec_json,
summary_text=row.summary_text,
code_version=row.code_version,
data_version=row.data_version,
job_id=row.job_id,
created_at=row.created_at,
result_bytes=int(row.result_bytes or 0),
)
for row in self._session.execute(stmt).all()
]
def count_filtered(self, *, kind: str | None = None, q: str | None = None) -> int:
stmt = (
select(func.count())
.select_from(ExperimentModel)
.where(*self._filter_conditions(kind=kind, q=q))
)
return int(self._session.execute(stmt).scalar() or 0)
def delete(self, experiment_id: str) -> bool:
"""删除归档本身;返回 False 表示不存在。不触碰 job 表。"""
row = self._session.get(ExperimentModel, experiment_id)
if row is None:
return False
self._session.delete(row)
self._session.flush()
return True
@@ -12,28 +12,39 @@ from collections.abc import Iterator, Sequence
from datetime import date
from typing import Any
from sqlalchemy import Float, String, cast, select
from sqlalchemy.orm import Session
from sqlalchemy import Float, String, and_, cast, func, insert, or_, select
from sqlalchemy.orm import Session, aliased
from app.domain.entities.market import (
DAILY_BAR_NUMERIC_FIELDS,
DAILY_BASIC_NUMERIC_FIELDS,
AdjustFactor,
DailyBar,
DailyBasic,
FinancialIndicator,
Stock,
StockNameHistory,
SyncLog,
TradingCalendar,
)
from app.infrastructure.persistence.sqlalchemy.models.market import (
AdjustFactorModel,
DailyBasicModel,
FinancialIndicatorModel,
StockDailyModel,
StockModel,
StockNameHistoryModel,
SyncLogModel,
TradingCalendarModel,
)
# 日线数值列白名单(研究面板只需这些;symbol/trade_date 恒返回)
BAR_FLOAT_COLUMNS = ("open", "high", "low", "close", "volume", "amount")
# 数值列白名单(研究面板只需这些;symbol/trade_date 恒返回)
# 单一事实来源在 domain.entities.market —— quant 层装配与本次列裁剪共用同一套列名。
BAR_FLOAT_COLUMNS = DAILY_BAR_NUMERIC_FIELDS
BASIC_FLOAT_COLUMNS = DAILY_BASIC_NUMERIC_FIELDS
# 复权折算只作用于价格列;volume/amount 保持原始口径(股数/金额不因除权而缩放)
_ADJUSTED_PRICE_COLUMNS = frozenset({"open", "high", "low", "close"})
# 实体类型 → (ORM Model, 幂等键列)
_TABLE = {
@@ -41,7 +52,9 @@ _TABLE = {
TradingCalendar: (TradingCalendarModel, ["calendar_date"]),
DailyBar: (StockDailyModel, ["symbol", "trade_date"]),
AdjustFactor: (AdjustFactorModel, ["symbol", "trade_date"]),
DailyBasic: (DailyBasicModel, ["symbol", "trade_date"]),
FinancialIndicator: (FinancialIndicatorModel, ["symbol", "report_date", "announce_date"]),
StockNameHistory: (StockNameHistoryModel, ["symbol", "start_date"]),
SyncLog: (SyncLogModel, ["id"]),
}
@@ -60,9 +73,14 @@ def _upsert_by_business_key(
) -> int:
"""按业务幂等键**批量**查重后 insert/update(AGENT.md §1 性能友好:不做逐行 select)。
- 一次查询取出本批已有的键 → 新行批量 add,已有行就地更新
- 先按业务键分批查出**已存在的键**(只取键列,不物化 ORM 实体)
- 新行走 Core 批量 insert(SQLAlchemy 2.0 insertmanyvalues:多行单语句),
避免「逐行 ORM add + flush」在整表同步下产生数千次往返
- 已存在行仍走 ORM 就地更新(量小,且保留属性级语义)
- 同一批内重复键(数据源偶发):以「后出现者」为准覆盖(避免 UNIQUE 冲突)
- 返回处理实体总数(含新增与更新),与旧逐行实现语义一致
方言无关:只用 select/insert/update,无 SQLite / MySQL 专有语法。
"""
if not entities:
return 0
@@ -74,30 +92,72 @@ def _upsert_by_business_key(
values = _fields_of(ent)
keyed.append((tuple(values[k] for k in key_cols), values))
existing_rows: dict[tuple, Any] = {}
keys = [k for k, _v in keyed]
if keys:
from sqlalchemy import tuple_
existing_rows = _load_existing(session, model_cls, key_cols, key_col_attrs, keyed)
stmt = select(model_cls).where(tuple_(*key_col_attrs).in_(keys))
for row in session.scalars(stmt):
key = tuple(getattr(row, k) for k in key_cols)
existing_rows[key] = row
pending: dict[tuple, Any] = {}
to_insert: list[dict] = []
to_update: list[Any] = []
seen: dict[tuple, Any] = {}
for key, values in keyed:
row = existing_rows.get(key)
if row is None and key in pending:
row = pending[key] # 批内已待插入的同键 → 覆盖为新值
if row is None:
pending[key] = model_cls(**values)
else:
if row is not None:
for col, val in values.items():
setattr(row, col, val)
session.add_all(pending.values())
to_update.append(row)
seen[key] = row
continue
if key in seen:
# 批内重复键:先插入的那条就地改名(同一 session 内 pending 对象可直接改属性)
holder = seen[key]
if isinstance(holder, dict):
holder.clear()
holder.update(values)
else:
for col, val in values.items():
setattr(holder, col, val)
continue
payload = dict(values)
seen[key] = payload
to_insert.append(payload)
if to_insert:
# insertmanyvalues:分批以控制单条 SQL 参数规模(避免超出驱动的占位符上限)
chunk = 500
for i in range(0, len(to_insert), chunk):
session.execute(insert(model_cls), to_insert[i : i + chunk])
if to_update:
session.flush() # ORM 更新落库
return len(keyed)
def _load_existing(
session: Session,
model_cls,
key_cols: Sequence[str],
key_col_attrs: Sequence,
keyed: Sequence[tuple[tuple, dict]],
) -> dict[tuple, Any]:
"""查出已存在的业务键 → ORM 行(键列分批查询,规避超长 IN 子句)。"""
from sqlalchemy import tuple_
keys = [k for k, _v in keyed]
if not keys:
return {}
unique_keys = list(dict.fromkeys(keys))
out: dict[tuple, Any] = {}
chunk = 500
for i in range(0, len(unique_keys), chunk):
part = unique_keys[i : i + chunk]
# 单键(如 symbol)直接用 IN,快且可走索引;复合键用 tuple IN
cond = (
key_col_attrs[0].in_([k[0] for k in part])
if len(key_col_attrs) == 1
else tuple_(*key_col_attrs).in_(part)
)
for row in session.scalars(select(model_cls).where(cond)):
out[tuple(getattr(row, k) for k in key_cols)] = row
return out
class SqlAlchemyStockRepository:
def __init__(self, session: Session) -> None:
self._session = session
@@ -185,28 +245,90 @@ class SqlAlchemyDailyBarRepository:
end: date,
columns: Sequence[str],
adjust: str = "none",
price_adjust: str = "none",
) -> Iterator[tuple]:
"""流式返回 (symbol, trade_date_iso, *float_cols) 元组,分批拉取。
内存优化:与 get_range_many 不同,不实例化 ORM 对象 / Decimal,
只 SELECT 所需列并在 SQL 侧 CAST 为 REAL,适合一次装配几十万~几百万行面板。
两个 adjust 参数**语义不同,不可混用**(v3 §20.5):
- `adjust`:行集口径过滤(stock_daily.adjust = none/qfq 兜底行),默认 none 主口径。
- `price_adjust`:复权折算口径(none/qfq/hfq),在 SQL 侧 JOIN adjust_factor 折算
**价格列**(open/high/low/close;volume/amount 不折算,保持原口径):
hfq → price × factor;qfq → price × factor / 该股最新 factor。
因子缺失时按 1.0 兜底(不折算),缺口由 count_price_adjust_gaps 统计上报。
"""
cols = list(columns)
unknown = [c for c in cols if c not in BAR_FLOAT_COLUMNS]
if unknown:
raise ValueError(f"不支持的行情列: {unknown}(可用: {BAR_FLOAT_COLUMNS})")
numeric_expr = [cast(getattr(StockDailyModel, c), Float) for c in cols]
stmt = (
select(StockDailyModel.symbol, cast(StockDailyModel.trade_date, String), *numeric_expr)
.where(
StockDailyModel.symbol.in_(list(symbols)),
StockDailyModel.trade_date >= start,
StockDailyModel.trade_date <= end,
StockDailyModel.adjust == adjust,
if price_adjust not in ("none", "qfq", "hfq"):
raise ValueError(f"不支持的复权口径: {price_adjust}(可用: none/qfq/hfq)")
day_col = cast(StockDailyModel.trade_date, String)
if price_adjust == "none":
numeric_expr = [cast(getattr(StockDailyModel, c), Float) for c in cols]
stmt = (
select(StockDailyModel.symbol, day_col, *numeric_expr)
.where(
StockDailyModel.symbol.in_(list(symbols)),
StockDailyModel.trade_date >= start,
StockDailyModel.trade_date <= end,
StockDailyModel.adjust == adjust,
)
.order_by(StockDailyModel.symbol, StockDailyModel.trade_date)
.execution_options(yield_per=20000)
)
.order_by(StockDailyModel.symbol, StockDailyModel.trade_date)
.execution_options(yield_per=20000)
)
else:
af = aliased(AdjustFactorModel)
mult = func.coalesce(af.factor, 1.0)
from_clause = None
if price_adjust == "qfq":
# 以该股**最新**因子归一(表内最大值,与 chart_service 的显示口径一致)。
# ⚠️ COALESCE 必须在最外层:若写成 coalesce(factor,1)/max,缺失因子的行会被
# 缩放到 1/max(例:max=5 → 假跌 80%),与「缺失因子按 1.0 兜底不折算」的
# 口径相矛盾,并污染 qfq 的收益/回撤/个股曲线。
latest = (
select(
AdjustFactorModel.symbol.label("symbol"),
func.max(AdjustFactorModel.factor).label("mx"),
)
.group_by(AdjustFactorModel.symbol)
.subquery()
)
mult = func.coalesce(af.factor / func.coalesce(latest.c.mx, 1.0), 1.0)
from_clause = latest
numeric_expr = [
cast(getattr(StockDailyModel, c), Float) * mult
if c in _ADJUSTED_PRICE_COLUMNS
else cast(getattr(StockDailyModel, c), Float)
for c in cols
]
stmt = (
select(StockDailyModel.symbol, day_col, *numeric_expr)
.select_from(StockDailyModel)
.outerjoin(
af,
and_(
af.symbol == StockDailyModel.symbol,
af.trade_date == StockDailyModel.trade_date,
),
)
)
if from_clause is not None:
stmt = stmt.outerjoin(from_clause, from_clause.c.symbol == StockDailyModel.symbol)
stmt = (
stmt.where(
StockDailyModel.symbol.in_(list(symbols)),
StockDailyModel.trade_date >= start,
StockDailyModel.trade_date <= end,
StockDailyModel.adjust == adjust,
)
.order_by(StockDailyModel.symbol, StockDailyModel.trade_date)
.execution_options(yield_per=20000)
)
result = self._session.execute(stmt)
while True:
chunk = result.fetchmany(20000)
@@ -215,6 +337,47 @@ class SqlAlchemyDailyBarRepository:
for row in chunk:
yield tuple(row)
def count_price_adjust_gaps(
self, symbols: Sequence[str], start: date, end: date
) -> tuple[int, int]:
"""统计复权因子覆盖缺口:返回 (行情行数, 无对应 adjust_factor 的行数)。
用途:复权回测必须如实告知「哪些行是按 1.0 兜底未折算的」
(AGENT.md §24 禁止假装支持)。单条聚合查询,代价可忽略。
"""
syms = list(symbols)
if not syms:
return 0, 0
total = self._session.scalar(
select(func.count())
.select_from(StockDailyModel)
.where(
StockDailyModel.symbol.in_(syms),
StockDailyModel.trade_date >= start,
StockDailyModel.trade_date <= end,
StockDailyModel.adjust == "none",
)
)
missing = self._session.scalar(
select(func.count())
.select_from(StockDailyModel)
.outerjoin(
AdjustFactorModel,
and_(
AdjustFactorModel.symbol == StockDailyModel.symbol,
AdjustFactorModel.trade_date == StockDailyModel.trade_date,
),
)
.where(
StockDailyModel.symbol.in_(syms),
StockDailyModel.trade_date >= start,
StockDailyModel.trade_date <= end,
StockDailyModel.adjust == "none",
AdjustFactorModel.id.is_(None),
)
)
return int(total or 0), int(missing or 0)
def latest_date(self, symbol: str) -> date | None:
return self._session.scalar(
select(StockDailyModel.trade_date)
@@ -244,6 +407,179 @@ class SqlAlchemyAdjustFactorRepository:
return [AdjustFactor.model_validate(r, from_attributes=True) for r in rows]
class SqlAlchemyDailyBasicRepository:
"""每日指标(估值 / 股息率 / 市值)仓储实现。"""
def __init__(self, session: Session) -> None:
self._session = session
def upsert_many(self, rows: Sequence[DailyBasic]) -> int:
return _upsert_by_business_key(self._session, DailyBasic, rows)
def get_range(self, symbol: str, start: date, end: date) -> list[DailyBasic]:
rows = self._session.scalars(
select(DailyBasicModel)
.where(
DailyBasicModel.symbol == symbol,
DailyBasicModel.trade_date >= start,
DailyBasicModel.trade_date <= end,
)
.order_by(DailyBasicModel.trade_date)
).all()
return [DailyBasic.model_validate(r, from_attributes=True) for r in rows]
def get_range_many(
self,
symbols: Sequence[str],
start: date,
end: date,
) -> list[DailyBasic]:
if not symbols:
return []
rows = self._session.scalars(
select(DailyBasicModel)
.where(
DailyBasicModel.symbol.in_(list(symbols)),
DailyBasicModel.trade_date >= start,
DailyBasicModel.trade_date <= end,
)
.order_by(DailyBasicModel.symbol, DailyBasicModel.trade_date)
).all()
return [DailyBasic.model_validate(r, from_attributes=True) for r in rows]
def stream_range_many_columns(
self,
symbols: Sequence[str],
start: date,
end: date,
columns: Sequence[str],
) -> Iterator[tuple]:
"""流式 (symbol, trade_date_iso, *float_cols);SQL 侧 CAST 为 REAL。"""
cols = list(columns)
unknown = [c for c in cols if c not in BASIC_FLOAT_COLUMNS]
if unknown:
raise ValueError(f"不支持的指标列: {unknown}(可用: {BASIC_FLOAT_COLUMNS})")
if not symbols:
return
numeric_expr = [cast(getattr(DailyBasicModel, c), Float) for c in cols]
stmt = (
select(DailyBasicModel.symbol, cast(DailyBasicModel.trade_date, String), *numeric_expr)
.where(
DailyBasicModel.symbol.in_(list(symbols)),
DailyBasicModel.trade_date >= start,
DailyBasicModel.trade_date <= end,
)
.order_by(DailyBasicModel.symbol, DailyBasicModel.trade_date)
.execution_options(yield_per=20000)
)
result = self._session.execute(stmt)
while True:
chunk = result.fetchmany(20000)
if not chunk:
break
for row in chunk:
yield tuple(row)
def latest_date(self) -> date | None:
return self._session.scalar(
select(DailyBasicModel.trade_date)
.order_by(DailyBasicModel.trade_date.desc())
.limit(1)
)
def missing_dates(self, start: date, end: date) -> list[date]:
"""开市但表内无任何行的交易日(按交易日历判定)。"""
have = select(DailyBasicModel.trade_date).where(
DailyBasicModel.trade_date >= start,
DailyBasicModel.trade_date <= end,
)
rows = self._session.scalars(
select(TradingCalendarModel.calendar_date)
.where(
TradingCalendarModel.calendar_date >= start,
TradingCalendarModel.calendar_date <= end,
TradingCalendarModel.is_open.is_(True),
TradingCalendarModel.calendar_date.not_in(have),
)
.order_by(TradingCalendarModel.calendar_date)
).all()
return list(rows)
def has_date(self, day: date) -> bool:
exists = self._session.scalar(
select(DailyBasicModel.id).where(DailyBasicModel.trade_date == day).limit(1)
)
return exists is not None
class SqlAlchemyStockNameHistoryRepository:
"""名称变更历史仓储:时点名称查询(exclude_st 的点时口径依据)。"""
def __init__(self, session: Session) -> None:
self._session = session
def upsert_many(self, rows: Sequence[StockNameHistory]) -> int:
return _upsert_by_business_key(self._session, StockNameHistory, rows)
def names_as_of(self, symbols: Sequence[str], as_of: date) -> dict[str, str]:
if not symbols:
return {}
stmt = (
select(
StockNameHistoryModel.symbol,
StockNameHistoryModel.name,
StockNameHistoryModel.start_date,
)
.where(
StockNameHistoryModel.symbol.in_(list(symbols)),
StockNameHistoryModel.start_date <= as_of,
or_(
StockNameHistoryModel.end_date.is_(None),
StockNameHistoryModel.end_date >= as_of,
),
)
.order_by(StockNameHistoryModel.symbol, StockNameHistoryModel.start_date)
)
out: dict[str, str] = {}
for symbol, name, _start in self._session.execute(stmt):
out[symbol] = name # 升序遍历 → 最后写入的是该时点生效的那个区间
return out
def name_spans(
self, symbols: Sequence[str]
) -> dict[str, list[tuple[date, date | None, str]]]:
if not symbols:
return {}
stmt = (
select(
StockNameHistoryModel.symbol,
StockNameHistoryModel.start_date,
StockNameHistoryModel.end_date,
StockNameHistoryModel.name,
)
.where(StockNameHistoryModel.symbol.in_(list(symbols)))
.order_by(StockNameHistoryModel.symbol, StockNameHistoryModel.start_date)
)
spans: dict[str, list[tuple[date, date | None, str]]] = {}
for symbol, start, end, name in self._session.execute(stmt):
spans.setdefault(symbol, []).append((start, end, name))
return spans
def count_rows(self) -> int:
return int(self._session.scalar(select(func.count()).select_from(StockNameHistoryModel)) or 0)
def namechange_dates(self) -> tuple[date | None, date | None]:
row = self._session.execute(
select(
func.min(StockNameHistoryModel.start_date),
func.max(StockNameHistoryModel.start_date),
)
).first()
if not row:
return None, None
return row[0], row[1]
class SqlAlchemyFinancialRepository:
def __init__(self, session: Session) -> None:
self._session = session
@@ -77,10 +77,13 @@ class SqlAlchemyStrategyRepository:
def _to_entity(row: StrategyModel) -> StrategyDefinition:
data = json.loads(row.config_json)
# 列字段由 DB 行回填,避免与 config_json 重复
# 列字段由 DB 行回填,避免与 config_json 重复。
# description 必须一并回填:它是列字段(String(300)),save() 会写入,
# 但这里若只从 config_json 里 pop 掉却不回填,读回的策略说明会恒为空串
# (读写不对称:保存的说明看不到,策略库/编辑页都拿不到)。
for key in ("name", "version", "description", "spec_type"):
data.pop(key, None)
return StrategyDefinition(
id=row.id, name=row.name, version=row.version,
id=row.id, name=row.name, version=row.version, description=row.description,
created_at=row.created_at, **data,
)