汇总三轮未提交的开发(每轮均在本机 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 逐类验证归档页)。
287 lines
11 KiB
Python
287 lines
11 KiB
Python
"""Phase 4 测试:Job 状态机 / 执行与 Experiment 自动归档 / API 提交与查询。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import sys
|
||
from datetime import date, datetime
|
||
|
||
import pytest
|
||
from app.application.services import job_executor as je
|
||
from app.application.services.job_executor import execute_job
|
||
from app.domain.entities.research import (
|
||
JobRecord,
|
||
JobStatus,
|
||
ResearchSpec,
|
||
)
|
||
from app.infrastructure.persistence.sqlalchemy import session as sess_mod
|
||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||
from app.infrastructure.persistence.sqlalchemy.models.market import (
|
||
StockDailyModel,
|
||
StockModel,
|
||
)
|
||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||
SqlAlchemyExperimentRepository,
|
||
SqlAlchemyJobRepository,
|
||
)
|
||
from app.main import app
|
||
from app.quant.engine import LocalEngine
|
||
from fastapi.testclient import TestClient
|
||
from sqlalchemy import create_engine
|
||
from sqlalchemy.orm import Session
|
||
|
||
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
|
||
|
||
_SYMS = ["60000" + str(i) + ".SH" for i in range(5)]
|
||
|
||
|
||
def _spec(**kw) -> ResearchSpec:
|
||
base = dict(
|
||
type="backtest",
|
||
universe={"exclude_st": False, "min_listing_days": 0},
|
||
factors=[{"name": "momentum_20", "weight": 1.0}],
|
||
selection={"top_n": 1},
|
||
rebalance="monthly",
|
||
period=["2024-03-01", "2024-10-31"],
|
||
)
|
||
base.update(kw)
|
||
return ResearchSpec.model_validate(base)
|
||
|
||
|
||
class _FakeStockRepo:
|
||
def __init__(self, stocks):
|
||
self._stocks = stocks
|
||
|
||
def list(self):
|
||
return self._stocks
|
||
|
||
|
||
class _FakeDailyRepo:
|
||
def __init__(self, bars):
|
||
self._bars = bars
|
||
|
||
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]
|
||
|
||
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 _fakes():
|
||
from app.domain.entities.market import Stock
|
||
|
||
stocks = [
|
||
Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1)) for i, s in enumerate(_SYMS)
|
||
]
|
||
drifts = {s: 0.003 - 0.0015 * i for i, s in enumerate(_SYMS)}
|
||
bars = bars_dataframe_to_daily_bars(synthetic_daily(drifts, n=300))
|
||
return stocks, bars
|
||
|
||
|
||
@pytest.fixture()
|
||
def sf(tmp_path):
|
||
engine = create_engine(f"sqlite:///{tmp_path / 'jobs.db'}", future=True)
|
||
Base.metadata.create_all(engine)
|
||
return lambda: Session(engine)
|
||
|
||
|
||
class TestJobExecutor:
|
||
def test_success_archives_experiment(self, sf) -> None:
|
||
stocks, bars = _fakes()
|
||
job = JobRecord(
|
||
id="JOB-TEST-1",
|
||
kind="backtest",
|
||
spec_json=_spec().model_dump_json(),
|
||
status=JobStatus.QUEUED,
|
||
created_at=datetime.now(),
|
||
)
|
||
with sf() as session:
|
||
SqlAlchemyJobRepository(session).create(job)
|
||
session.commit()
|
||
|
||
execute_job(
|
||
"JOB-TEST-1",
|
||
session_factory=sf,
|
||
job_repo_factory=lambda s: SqlAlchemyJobRepository(s),
|
||
experiment_repo_factory=lambda s: SqlAlchemyExperimentRepository(s),
|
||
stock_repo_factory=lambda s: _FakeStockRepo(stocks),
|
||
daily_repo_factory=lambda s: _FakeDailyRepo(bars),
|
||
engine=LocalEngine(),
|
||
)
|
||
|
||
with sf() as session:
|
||
done = SqlAlchemyJobRepository(session).get("JOB-TEST-1")
|
||
assert done is not None
|
||
assert done.status == JobStatus.SUCCESS
|
||
# 完整结果只在 experiment 存一份(job 侧不再重复落库)
|
||
assert done.result_json is None
|
||
assert done.experiment_id is not None
|
||
exp = SqlAlchemyExperimentRepository(session).get(done.experiment_id)
|
||
assert exp is not None
|
||
assert exp.kind == "backtest"
|
||
assert exp.result_json # 归档里必须有完整结果
|
||
assert exp.summary_text is not None
|
||
assert "总收益" in (exp.summary_text or "")
|
||
|
||
def test_failure_marks_failed(self, sf) -> None:
|
||
stocks, _bars = _fakes()
|
||
job = JobRecord(
|
||
id="JOB-TEST-2",
|
||
kind="backtest",
|
||
spec_json=_spec(factors=[{"name": "no_such", "weight": 1.0}]).model_dump_json(),
|
||
status=JobStatus.QUEUED,
|
||
created_at=datetime.now(),
|
||
)
|
||
with sf() as session:
|
||
SqlAlchemyJobRepository(session).create(job)
|
||
session.commit()
|
||
|
||
execute_job(
|
||
"JOB-TEST-2",
|
||
session_factory=sf,
|
||
job_repo_factory=lambda s: SqlAlchemyJobRepository(s),
|
||
experiment_repo_factory=lambda s: SqlAlchemyExperimentRepository(s),
|
||
stock_repo_factory=lambda s: _FakeStockRepo(stocks),
|
||
daily_repo_factory=lambda s: _FakeDailyRepo([]),
|
||
engine=LocalEngine(),
|
||
)
|
||
with sf() as session:
|
||
done = SqlAlchemyJobRepository(session).get("JOB-TEST-2")
|
||
assert done is not None
|
||
assert done.status == JobStatus.FAILED
|
||
assert done.error
|
||
|
||
|
||
class TestJobResilience:
|
||
"""内存隔离专项:启动残留清理 + 子进程异常退出兜底。"""
|
||
|
||
def test_mark_stale_jobs_failed(self, sf) -> None:
|
||
repo_f = lambda s: SqlAlchemyJobRepository(s) # noqa: E731
|
||
rows = [
|
||
("JOB-STALE-1", JobStatus.RUNNING),
|
||
("JOB-STALE-2", JobStatus.QUEUED),
|
||
("JOB-OK-3", JobStatus.SUCCESS),
|
||
]
|
||
for jid, status in rows:
|
||
job = JobRecord(
|
||
id=jid,
|
||
kind="backtest",
|
||
spec_json=_spec().model_dump_json(),
|
||
status=status,
|
||
created_at=datetime.now(),
|
||
)
|
||
with sf() as session:
|
||
repo_f(session).create(job)
|
||
session.commit()
|
||
|
||
n = je.mark_stale_jobs_failed(session_factory=sf, job_repo_factory=repo_f)
|
||
assert n == 2
|
||
with sf() as session:
|
||
stale = repo_f(session).get("JOB-STALE-1")
|
||
queued = repo_f(session).get("JOB-STALE-2")
|
||
ok = repo_f(session).get("JOB-OK-3")
|
||
assert stale is not None and stale.status == JobStatus.FAILED
|
||
assert stale.finished_at is not None and "重启" in (stale.error or "")
|
||
assert queued is not None and queued.status == JobStatus.FAILED
|
||
assert ok is not None and ok.status == JobStatus.SUCCESS
|
||
|
||
def test_subprocess_abnormal_exit_marks_failed(self, sf, monkeypatch) -> None:
|
||
repo_f = lambda s: SqlAlchemyJobRepository(s) # noqa: E731
|
||
job = JobRecord(
|
||
id="JOB-SUB-1",
|
||
kind="backtest",
|
||
spec_json=_spec().model_dump_json(),
|
||
status=JobStatus.QUEUED,
|
||
created_at=datetime.now(),
|
||
)
|
||
with sf() as session:
|
||
repo_f(session).create(job)
|
||
session.commit()
|
||
monkeypatch.setattr(je, "_job_mode", lambda: "subprocess")
|
||
monkeypatch.setattr(
|
||
je, "_job_worker_cmd", lambda job_id: [sys.executable, "-c", "import sys; sys.exit(7)"]
|
||
)
|
||
je.run_job_background("JOB-SUB-1", session_factory=sf, job_repo_factory=repo_f)
|
||
|
||
with sf() as session:
|
||
done = repo_f(session).get("JOB-SUB-1")
|
||
assert done is not None
|
||
assert done.status == JobStatus.FAILED
|
||
assert done.error is not None and "code=7" in done.error
|
||
|
||
|
||
@pytest.fixture()
|
||
def seeded_api_db(tmp_path, monkeypatch) -> None:
|
||
"""把全局 SessionLocal 指向 tmp 种子库(5 只股票 × 300 交易日),
|
||
Job 后台执行跑小数据,避免依赖真实全市场库(779 万行)导致过慢。"""
|
||
from sqlalchemy.orm import sessionmaker
|
||
|
||
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
|
||
Base.metadata.create_all(engine)
|
||
sf = sessionmaker(bind=engine, expire_on_commit=False)
|
||
drifts = {s: 0.003 - 0.0015 * i for i, s in enumerate(_SYMS)}
|
||
daily = synthetic_daily(drifts, n=320)
|
||
with sf() as session:
|
||
for i, sym in enumerate(_SYMS):
|
||
session.add(StockModel(symbol=sym, name=f"测试股份{i}", list_date=date(1999, 1, 1)))
|
||
rows = [
|
||
{
|
||
"symbol": r.symbol,
|
||
"trade_date": r.trade_date,
|
||
"source": "tushare",
|
||
"adjust": "none",
|
||
"open": float(r.open),
|
||
"high": float(r.high),
|
||
"low": float(r.low),
|
||
"close": float(r.close),
|
||
"volume": float(r.volume),
|
||
"amount": float(r.amount),
|
||
}
|
||
for r in daily.itertuples()
|
||
]
|
||
session.execute(StockDailyModel.__table__.insert(), rows)
|
||
session.commit()
|
||
monkeypatch.setattr(sess_mod, "SessionLocal", sf)
|
||
|
||
|
||
class TestJobsApi:
|
||
def test_submit_then_query(self, seeded_api_db) -> None:
|
||
with TestClient(app) as client:
|
||
resp = client.post("/api/jobs", json=_spec().model_dump(mode="json"))
|
||
assert resp.status_code == 200
|
||
body = resp.json()
|
||
assert body["status"] == "queued"
|
||
job_id = body["job_id"]
|
||
|
||
# 后台任务随请求执行完成(TestClient 同步运行 BackgroundTasks)
|
||
for _ in range(20):
|
||
state = client.get(f"/api/jobs/{job_id}").json()
|
||
if state["status"] in (JobStatus.SUCCESS, JobStatus.FAILED):
|
||
break
|
||
assert state["status"] == JobStatus.SUCCESS
|
||
assert state["result"] is not None
|
||
assert state["spec"]["factors"][0]["name"] == "momentum_20"
|
||
assert state["experiment_id"]
|
||
|
||
def test_unknown_job_404(self, seeded_api_db) -> None:
|
||
with TestClient(app) as client:
|
||
assert client.get("/api/jobs/JOB-NOPE").status_code == 404
|
||
|
||
def test_experiments_list_and_detail(self, seeded_api_db) -> None:
|
||
with TestClient(app) as client:
|
||
# 先跑一个 job 生成归档,再验证列表与详情
|
||
resp = client.post("/api/jobs", json=_spec().model_dump(mode="json"))
|
||
assert resp.status_code == 200
|
||
job_id = resp.json()["job_id"]
|
||
for _ in range(30):
|
||
st = client.get(f"/api/jobs/{job_id}").json()["status"]
|
||
if st in (JobStatus.SUCCESS, JobStatus.FAILED):
|
||
break
|
||
assert st == JobStatus.SUCCESS
|
||
|
||
exps = client.get("/api/experiments").json()
|
||
assert exps, "job 成功后应有 Experiment 归档"
|
||
first = exps[0]
|
||
detail = client.get(f"/api/experiments/{first['id']}")
|
||
assert detail.status_code == 200
|
||
assert "result" in detail.json()
|