Files
qlib/backend/tests/test_jobs_experiments.py
T
Simon 0ea229d766 feat: Phase 4 — Experiment 自动归档 + 异步 Job(状态机 / SSE / 一键复跑)
- 数据表:job / experiment(spec/result JSON 存档、code_version),Alembic 迁移 53113c80257f
- Job:queued→running→(success|failed) 状态机,BackgroundTasks 本地执行 + 失败兜底标记;结果与 Experiment 关联
- Experiment:每次研究成功自动归档(含 git commit 与收益摘要),支持一键复跑(同 spec 重建 Job)
- API:POST /api/jobs、GET /api/jobs/{id}(内嵌结果)、SSE /api/jobs/{id}/events、/api/experiments 列表/详情/rerun
- 前端:新增「实验」页(列表 / 详情 / 复跑 + Job 轮询);导航更新
- 端到端验证:真实 20 股 job 提交→后台执行→success→EXP 归档(-12.41%);executor 成功/失败路径单测
- 测试 73 passed(新增 5 项 Job/Experiment)/ ruff clean / 前端 tsc + build 通过
2026-09-06 17:18:46 +08:00

178 lines
6.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Phase 4 测试:Job 状态机 / 执行与 Experiment 自动归档 / API 提交与查询。"""
from __future__ import annotations
from datetime import date, datetime
import pytest
from app.application.services.job_executor import execute_job
from app.domain.entities.research import (
JobRecord,
JobStatus,
ResearchSpec,
)
from app.infrastructure.persistence.sqlalchemy.base import Base
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):
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
assert done.result_json is not 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.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 TestJobsApi:
def test_submit_then_query(self) -> 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) -> None:
with TestClient(app) as client:
assert client.get("/api/jobs/JOB-NOPE").status_code == 404
def test_experiments_list_and_detail(self) -> None:
with TestClient(app) as client:
resp = client.get("/api/experiments")
assert resp.status_code == 200
exps = resp.json()
if exps: # 本机真实库中可能有历史实验;detail 可读即通过
first = exps[0]
detail = client.get(f"/api/experiments/{first['id']}")
assert detail.status_code == 200
assert "result" in detail.json()