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:
+67
@@ -0,0 +1,67 @@
|
||||
"""add daily_basic table
|
||||
|
||||
Revision ID: 22d7380706f7
|
||||
Revises: f5e0d1c2b3a4
|
||||
Create Date: 2026-09-19 15:53:37.029097
|
||||
|
||||
每日指标表(Tushare daily_basic):估值 / 股息率 / 市值,幂等键 (symbol, trade_date)。
|
||||
供高股息等横截面选股因子使用(研究侧按 trade_date <= as_of 取用,无未来函数)。
|
||||
|
||||
说明:本文件由 `alembic revision --autogenerate` 生成后**手工裁剪**。
|
||||
自动生成时同时检出了 `index_weight` 的 drop —— 那是 `IndexWeightModel`
|
||||
未在 `models/__init__.py` 注册导致的假差异(表实际存在),
|
||||
已在同一提交中补上注册,此处只保留 daily_basic 的建表语句。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = '22d7380706f7'
|
||||
down_revision: str | None = 'f5e0d1c2b3a4'
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
'daily_basic',
|
||||
sa.Column(
|
||||
'id',
|
||||
sa.BigInteger().with_variant(sa.Integer(), 'sqlite'),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column('symbol', sa.String(length=12), nullable=False),
|
||||
sa.Column('trade_date', sa.Date(), nullable=False),
|
||||
sa.Column('source', sa.String(length=16), server_default='tushare', nullable=False),
|
||||
sa.Column('close', sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column('turnover_rate', sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column('volume_ratio', sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column('pe', sa.Numeric(precision=16, scale=4), nullable=True),
|
||||
sa.Column('pe_ttm', sa.Numeric(precision=16, scale=4), nullable=True),
|
||||
sa.Column('pb', sa.Numeric(precision=16, scale=4), nullable=True),
|
||||
sa.Column('ps', sa.Numeric(precision=16, scale=4), nullable=True),
|
||||
sa.Column('ps_ttm', sa.Numeric(precision=16, scale=4), nullable=True),
|
||||
sa.Column('dv_ratio', sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column('dv_ttm', sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column('total_share', sa.Numeric(precision=24, scale=4), nullable=True),
|
||||
sa.Column('float_share', sa.Numeric(precision=24, scale=4), nullable=True),
|
||||
sa.Column('free_share', sa.Numeric(precision=24, scale=4), nullable=True),
|
||||
sa.Column('total_mv', sa.Numeric(precision=24, scale=4), nullable=True),
|
||||
sa.Column('circ_mv', sa.Numeric(precision=24, scale=4), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('symbol', 'trade_date', name='uq_basic_symbol_date'),
|
||||
)
|
||||
with op.batch_alter_table('daily_basic', schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f('ix_daily_basic_symbol'), ['symbol'], unique=False)
|
||||
batch_op.create_index(batch_op.f('ix_daily_basic_trade_date'), ['trade_date'], unique=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.batch_alter_table('daily_basic', schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f('ix_daily_basic_trade_date'))
|
||||
batch_op.drop_index(batch_op.f('ix_daily_basic_symbol'))
|
||||
op.drop_table('daily_basic')
|
||||
+47
@@ -0,0 +1,47 @@
|
||||
"""research result_json → MEDIUMTEXT(长回测结果落库)
|
||||
|
||||
Revision ID: 7b1c4e9a52d8
|
||||
Revises: 22d7380706f7
|
||||
Create Date: 2026-09-19
|
||||
|
||||
背景:全市场多年回测结果(净值/回撤曲线 + 逐笔成交 + Signal↔Fill + 个股收益曲线)
|
||||
实测约 1.2MB,MySQL `TEXT`(64KB)会报 1406 Data too long → Job 归档失败
|
||||
(2026-09 高股息案例实测:回测本身成功,落库失败)。
|
||||
|
||||
本迁移把 job.result_json / experiment.result_json 放宽到 MEDIUMTEXT(16MB)。
|
||||
SQLite 等方言不区分 TEXT 长度,迁移里做方言判断,非 MySQL 直接跳过。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from alembic import op
|
||||
from sqlalchemy import text
|
||||
|
||||
revision = "7b1c4e9a52d8"
|
||||
down_revision = "22d7380706f7"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
_COLUMNS = (("job", "result_json", True), ("experiment", "result_json", False))
|
||||
|
||||
|
||||
def _dialect() -> str:
|
||||
return op.get_bind().dialect.name
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if _dialect() != "mysql":
|
||||
return # SQLite/TEXT 无长度限制,无需变更
|
||||
for table, column, nullable in _COLUMNS:
|
||||
null_clause = "NULL" if nullable else "NOT NULL"
|
||||
op.execute(
|
||||
text(f"ALTER TABLE {table} MODIFY COLUMN {column} MEDIUMTEXT {null_clause}")
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if _dialect() != "mysql":
|
||||
return
|
||||
for table, column, nullable in _COLUMNS:
|
||||
null_clause = "NULL" if nullable else "NOT NULL"
|
||||
op.execute(text(f"ALTER TABLE {table} MODIFY COLUMN {column} TEXT {null_clause}"))
|
||||
+64
@@ -0,0 +1,64 @@
|
||||
"""add stock_name_history table
|
||||
|
||||
Revision ID: a3f8c21d9b47
|
||||
Revises: 7b1c4e9a52d8
|
||||
Create Date: 2026-09-19 21:40:12.000000
|
||||
|
||||
股票名称变更历史(Tushare namechange):时点 ST / 风险警示判定的依据。
|
||||
|
||||
**为什么需要**:`stock.name` 是最新名称快照,用它做 `exclude_st` 会把
|
||||
「曾为高股息、后来才变 ST/退市」的标的在整段历史里都排除 —— 而那正是
|
||||
「股息陷阱」样本。实测对照:同一 spec 仅改 exclude_st,收益 +35.71% → +32.01%,
|
||||
即约 3.70pp 的收益被名称快照口径隐藏。
|
||||
|
||||
幂等键 (symbol, start_date)。时点查询:
|
||||
`start_date <= as_of AND (end_date IS NULL OR end_date >= as_of)`。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = 'a3f8c21d9b47'
|
||||
down_revision: str | None = '7b1c4e9a52d8'
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
'stock_name_history',
|
||||
sa.Column(
|
||||
'id',
|
||||
sa.BigInteger().with_variant(sa.Integer(), 'sqlite'),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column('symbol', sa.String(length=12), nullable=False),
|
||||
sa.Column('name', sa.String(length=64), nullable=False),
|
||||
sa.Column('start_date', sa.Date(), nullable=False),
|
||||
sa.Column('end_date', sa.Date(), nullable=True),
|
||||
sa.Column('ann_date', sa.Date(), nullable=True),
|
||||
sa.Column('change_reason', sa.String(length=32), nullable=True),
|
||||
sa.Column('source', sa.String(length=16), server_default='tushare', nullable=False),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('symbol', 'start_date', name='uq_name_symbol_start'),
|
||||
)
|
||||
with op.batch_alter_table('stock_name_history', schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f('ix_stock_name_history_symbol'), ['symbol'], unique=False)
|
||||
batch_op.create_index(
|
||||
batch_op.f('ix_stock_name_history_start_date'), ['start_date'], unique=False
|
||||
)
|
||||
batch_op.create_index(
|
||||
batch_op.f('ix_stock_name_history_end_date'), ['end_date'], unique=False
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.batch_alter_table('stock_name_history', schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f('ix_stock_name_history_end_date'))
|
||||
batch_op.drop_index(batch_op.f('ix_stock_name_history_start_date'))
|
||||
batch_op.drop_index(batch_op.f('ix_stock_name_history_symbol'))
|
||||
op.drop_table('stock_name_history')
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user