汇总三轮未提交的开发(每轮均在本机 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 逐类验证归档页)。
180 lines
7.0 KiB
Python
180 lines
7.0 KiB
Python
"""Job / Experiment 的 SQLAlchemy 实现(Phase 4)。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from sqlalchemy import Select, func, select
|
||
from sqlalchemy.orm import Session
|
||
|
||
from app.domain.entities.research import ExperimentRecord, ExperimentSummary, JobRecord
|
||
from app.infrastructure.persistence.sqlalchemy.models.jobs import ExperimentModel, JobModel
|
||
|
||
|
||
def _job_to_model(job: JobRecord) -> JobModel:
|
||
return JobModel(**job.model_dump())
|
||
|
||
|
||
class SqlAlchemyJobRepository:
|
||
def __init__(self, session: Session) -> None:
|
||
self._session = session
|
||
|
||
def create(self, job: JobRecord) -> JobRecord:
|
||
self._session.add(_job_to_model(job))
|
||
self._session.flush()
|
||
return job
|
||
|
||
def get(self, job_id: str) -> JobRecord | None:
|
||
row = self._session.get(JobModel, job_id)
|
||
return JobRecord.model_validate(row, from_attributes=True) if row else None
|
||
|
||
def update(self, job: JobRecord) -> None:
|
||
row = self._session.get(JobModel, job.id)
|
||
if row is None:
|
||
raise KeyError(f"job {job.id} 不存在")
|
||
for k, v in job.model_dump().items():
|
||
setattr(row, k, v)
|
||
|
||
def list_recent(self, kind: str | None = None, limit: int = 20) -> list[JobRecord]:
|
||
stmt = select(JobModel).order_by(JobModel.created_at.desc()).limit(limit)
|
||
if kind:
|
||
stmt = stmt.where(JobModel.kind == kind)
|
||
return [
|
||
JobRecord.model_validate(r, from_attributes=True)
|
||
for r in self._session.scalars(stmt).all()
|
||
]
|
||
|
||
def list_by_status(self, status: str, limit: int = 100) -> list[JobRecord]:
|
||
rows = self._session.scalars(
|
||
select(JobModel)
|
||
.where(JobModel.status == status)
|
||
.order_by(JobModel.created_at)
|
||
.limit(limit)
|
||
).all()
|
||
return [JobRecord.model_validate(r, from_attributes=True) for r in rows]
|
||
|
||
|
||
class SqlAlchemyExperimentRepository:
|
||
"""Experiment 归档仓储。
|
||
|
||
`list_filtered` 刻意**不加载 `result_json`**(MEDIUMTEXT,完整存档后单条数 MB),
|
||
只 SELECT 元数据列 + SQL 侧算出的字符长度,避免列表接口把上百 MB 拉进内存。
|
||
"""
|
||
|
||
def __init__(self, session: Session) -> None:
|
||
self._session = session
|
||
|
||
def save(self, experiment: ExperimentRecord) -> ExperimentRecord:
|
||
self._session.add(ExperimentModel(**experiment.model_dump()))
|
||
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
|
||
|
||
def list_recent(self, limit: int = 50) -> list[ExperimentRecord]:
|
||
rows = self._session.scalars(
|
||
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
|