Files
qlib/backend/tests/test_jobs_experiments.py
T
Simon 02e42184be test+fix: 新浪 jsonp 稳健解析(容忍注释/括号包裹/尾部杂字符);Job API 测试改 tmp 种子库
- 真实响应含 /*...*/ 注释与 var data=([...]); 前缀 → _extract_jsonp 取首 '[' 至末 ']'
- 新浪真实网络冒烟通过:财务 100 期(含披露日)+ 日K 最近窗口(source=sina/adjust=qfq)
- TestJobsApi / experiments 用例改为 monkeypatch SessionLocal → tmp 种子库(5 股×300 日),
  与真实全市场库(779 万行)解耦,全量稳定 <1min
- pytest 全量 141 passed / ruff clean
2026-09-06 21:34:29 +08:00

225 lines
8.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 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):
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
@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()